From a44930f9df11afda443c2c603182950c4c0234c9 Mon Sep 17 00:00:00 2001 From: Jason Morcos Date: Sun, 9 Aug 2026 10:53:56 -0700 Subject: [PATCH] feat(protocol): bound DTLS handshake deadline --- README.md | 9 + smartthings_local/protocol/dtls_handshake.py | 72 ++++ smartthings_local/protocol/dtls_probe.py | 98 ++---- smartthings_local/protocol/dtls_session.py | 93 ++--- tests/test_endpoint.py | 2 +- tests/test_public_api_contract.py | 6 + tests/test_session_connect_deadline.py | 344 +++++++++++++++++++ 7 files changed, 513 insertions(+), 111 deletions(-) create mode 100644 smartthings_local/protocol/dtls_handshake.py create mode 100644 tests/test_session_connect_deadline.py diff --git a/README.md b/README.md index 52c2f41..1355ef1 100644 --- a/README.md +++ b/README.md @@ -48,6 +48,15 @@ sess.subscribe(["operational", "state", "vs", "0"], # OBSERVE sess.close() ``` +`connect()` uses a 12-second monotonic DTLS handshake deadline by default. A +caller that needs a shorter bounded attempt can pass a positive finite value +without changing later reader timeouts. OpenSSL's DTLS timer schedules flight +retransmissions within that same deadline: + +```python +sess.connect(timeout=4.0) +``` + If the cert/key are minted at runtime and never written to disk (e.g. inside an HA config flow), create the provider from memory instead: diff --git a/smartthings_local/protocol/dtls_handshake.py b/smartthings_local/protocol/dtls_handshake.py new file mode 100644 index 0000000..8280e3b --- /dev/null +++ b/smartthings_local/protocol/dtls_handshake.py @@ -0,0 +1,72 @@ +"""Shared memory-BIO driver for bounded DTLS handshakes.""" + +from __future__ import annotations + +import time +from collections.abc import Callable + +from OpenSSL import SSL + +from .coap import split_dtls + +_HANDSHAKE_POLL_S = 0.5 +_MAX_DATAGRAM_SIZE = 65535 + + +def _drive_dtls_handshake( + connection, + sock, + *, + deadline: float, + retries: int | None = None, + on_datagram: Callable[[bytes], None] | None = None, +) -> bool: + """Drive one memory-BIO DTLS handshake up to a monotonic deadline. + + OpenSSL owns the retransmission schedule. ``retries`` optionally limits + how many expired retransmission timers are serviced; a normal session is + bounded only by its deadline, while the diagnostic probe retains its + explicit retry budget. + + Return ``True`` only when the handshake completes before the deadline. + TLS and socket failures are left to the caller to classify. + """ + retransmits = 0 + while time.monotonic() < deadline: + try: + connection.do_handshake() + return time.monotonic() < deadline + except SSL.WantReadError: + pass + + try: + output = connection.bio_read(_MAX_DATAGRAM_SIZE) + except SSL.WantReadError: + output = None + if output: + for record in split_dtls(output): + if sock.send(record) != len(record): + raise OSError("incomplete UDP send") + + remaining = deadline - time.monotonic() + if remaining <= 0: + break + sock.settimeout(min(_HANDSHAKE_POLL_S, remaining)) + try: + datagram = sock.recv(_MAX_DATAGRAM_SIZE) + except TimeoutError: + timer = connection.DTLSv1_get_timeout() + if timer is not None and timer <= 0: + if retries is not None and retransmits >= retries: + break + connection.DTLSv1_handle_timeout() + retransmits += 1 + continue + + if not datagram: + continue + if on_datagram is not None: + on_datagram(datagram) + connection.bio_write(datagram) + + return False diff --git a/smartthings_local/protocol/dtls_probe.py b/smartthings_local/protocol/dtls_probe.py index 9232912..3fbd5ef 100644 --- a/smartthings_local/protocol/dtls_probe.py +++ b/smartthings_local/protocol/dtls_probe.py @@ -39,8 +39,9 @@ from dataclasses import dataclass from OpenSSL import SSL from ..errors import ProbeError +from .auth import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain from .coap import split_dtls -from .dtls_session import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain +from .dtls_handshake import _drive_dtls_handshake from .endpoint import open_connected_udp_socket # DTLS record content types (RFC 6347 §4.1) @@ -603,73 +604,40 @@ def diagnose_dtls_handshake( started = time.monotonic() deadline = started + timeout seen = set() - retransmits = 0 + + def record_datagram(datagram): + if result.rtt_s is None: + result.rtt_s = time.monotonic() - started + result.datagrams.append(datagram) + for content_type, detail in classify_datagram(datagram): + if content_type == _CT_HANDSHAKE: + if detail not in seen: + seen.add(detail) + result.handshake_msgs.append(detail) + if result.outcome == DEAD: + result.outcome = LIVE + elif content_type == _CT_ALERT and detail is not None: + level, name = detail + result.alert = (level, name) + if level == 2: # fatal + result.outcome = REJECTED + try: - while time.monotonic() < deadline: - try: - conn.do_handshake() - result.outcome = COMPLETED - if result.rtt_s is None: - result.rtt_s = time.monotonic() - started - break - except SSL.WantReadError: - pass - except SSL.Error: - # A fatal Alert lands here; the alert record was already - # captured below, so classification still works. - result.error = ProbeError() - break - - try: - o = conn.bio_read(65535) - if o: - for r in split_dtls(o): - if sock.send(r) != len(r): - raise OSError('short UDP send') - except SSL.WantReadError: - pass - - remaining = deadline - time.monotonic() - if remaining <= 0: - break - sock.settimeout(min(0.5, remaining)) - try: - d = sock.recv(65535) - except TimeoutError: - # No answer to the last flight. Service OpenSSL's DTLS - # retransmit timer: once it has counted down to 0, - # handle_timeout() re-queues the previous flight into the - # write BIO for the next iteration to flush. A live server - # answers within a flight or two; a silent/non-DTLS port - # never does, so we give up only after `retries` - # retransmits — one dropped ClientHello no longer reads as - # a false DEAD. - to = conn.DTLSv1_get_timeout() - if to is not None and to <= 0: - if retransmits >= retries: - break - conn.DTLSv1_handle_timeout() - retransmits += 1 - continue - if not d: - continue - + completed = _drive_dtls_handshake( + conn, + sock, + deadline=deadline, + retries=retries, + on_datagram=record_datagram, + ) + if completed: + result.outcome = COMPLETED if result.rtt_s is None: result.rtt_s = time.monotonic() - started - result.datagrams.append(d) - for ct, detail in classify_datagram(d): - if ct == _CT_HANDSHAKE: - if detail not in seen: - seen.add(detail) - result.handshake_msgs.append(detail) - if result.outcome == DEAD: - result.outcome = LIVE - elif ct == _CT_ALERT and detail is not None: - level, name = detail - result.alert = (level, name) - if level == 2: # fatal - result.outcome = REJECTED - conn.bio_write(d) + except SSL.Error: + # A fatal Alert lands here; record_datagram() has already classified + # the alert record before it is fed back into OpenSSL. + result.error = ProbeError() except OSError: result.error = ProbeError() finally: diff --git a/smartthings_local/protocol/dtls_session.py b/smartthings_local/protocol/dtls_session.py index 3e0034f..910bbe1 100644 --- a/smartthings_local/protocol/dtls_session.py +++ b/smartthings_local/protocol/dtls_session.py @@ -19,6 +19,7 @@ on a per-token Event the reader signals. OBSERVE notifications are delivered via the on_notification callback. """ import errno +import math import os import socket import threading @@ -48,6 +49,7 @@ from .auth import ( _OCF_ROOT_CA, _load_pem_chain, ) +from .dtls_handshake import _HANDSHAKE_POLL_S, _drive_dtls_handshake from .endpoint import open_connected_udp_socket import logging @@ -88,6 +90,20 @@ _ADVISORY_ERRNOS = frozenset( ) if value is not None ) +def _validate_handshake_timeout(timeout, default): + """Return one finite, positive DTLS handshake timeout.""" + value = default if timeout is None else timeout + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise TypeError('timeout must be a number or None') + try: + value = float(value) + except OverflowError: + raise ValueError( + 'timeout must be a positive finite number or None') from None + if not math.isfinite(value) or value <= 0: + raise ValueError('timeout must be a positive finite number or None') + return value + class DtlsCoapSession: """Single sustained DTLS-CoAP session. @@ -199,9 +215,16 @@ class DtlsCoapSession: # ---- lifecycle --------------------------------------------------- - def connect(self): - """DTLS handshake. Blocks up to HANDSHAKE_TIMEOUT_S. Raises - ConnectionError / TimeoutError on failure.""" + def connect(self, *, timeout: float | None = None): + """Perform a DTLS handshake within a monotonic deadline. + + ``timeout`` overrides ``HANDSHAKE_TIMEOUT_S`` for this call. OpenSSL + owns DTLS retransmission timing while every receive is capped by the + remaining budget, so wall-clock adjustments cannot change the bound. + """ + handshake_timeout = _validate_handshake_timeout( + timeout, self.HANDSHAKE_TIMEOUT_S) + deadline = time.monotonic() + handshake_timeout ctx = SSL.Context(SSL.DTLS_METHOD) self.auth.configure_context(ctx) @@ -209,59 +232,39 @@ class DtlsCoapSession: conn.set_connect_state() conn.set_ciphertext_mtu(self.mtu) + remaining = deadline - time.monotonic() + if remaining <= 0: + raise SessionTimeoutError() sock, endpoint = open_connected_udp_socket( self.host, self.port, family=self.family, local_port=self.local_port, - timeout=2.0, + timeout=min(_HANDSHAKE_POLL_S, remaining), ) dest = endpoint.sockaddr - t0 = time.time() backend_failed = False - while time.time() - t0 < self.HANDSHAKE_TIMEOUT_S: - try: - conn.do_handshake() - break - except SSL.WantReadError: - pass - except SSL.Error: - sock.close() - backend_failed = True - break - send_failed = False - try: - o = conn.bio_read(65535) - if o: - for r in _split_dtls(o): - if sock.send(r) != len(r): - raise OSError('incomplete UDP send') - except SSL.WantReadError: - pass - except OSError: - sock.close() - send_failed = True - if send_failed: - raise EndpointError() from OSError('UDP send failed') - receive_failed = False - try: - d = sock.recv(65535) - if d: - conn.bio_write(d) - except socket.timeout: - pass - except OSError: - sock.close() - receive_failed = True - if receive_failed: - raise EndpointError() from OSError('UDP receive failed') - time.sleep(0.05) - else: + io_failed = False + try: + completed = _drive_dtls_handshake( + conn, + sock, + deadline=deadline, + ) + except SSL.Error: + backend_failed = True + except OSError: + io_failed = True + if backend_failed: + sock.close() + raise SessionError() from ConnectionError('DTLS backend failed') + if io_failed: + sock.close() + raise EndpointError() from OSError('UDP handshake I/O failed') + if not completed: sock.close() raise SessionTimeoutError() - if backend_failed: - raise SessionError() from ConnectionError('DTLS backend failed') self.sock = sock self.conn = conn diff --git a/tests/test_endpoint.py b/tests/test_endpoint.py index 6d6ddf3..cd4643b 100644 --- a/tests/test_endpoint.py +++ b/tests/test_endpoint.py @@ -376,7 +376,7 @@ def test_session_uses_connected_socket_send_and_recv(monkeypatch): assert open_calls == [(('device.example', 5684), { 'family': socket.AF_INET6, 'local_port': None, - 'timeout': 2.0, + 'timeout': 0.5, })] session.close() diff --git a/tests/test_public_api_contract.py b/tests/test_public_api_contract.py index f40e5b4..5abf19a 100644 --- a/tests/test_public_api_contract.py +++ b/tests/test_public_api_contract.py @@ -81,6 +81,12 @@ def test_dtls_session_keeps_current_consumer_methods(): "subscribe", } assert expected <= set(dir(DtlsCoapSession)) + _assert_compatible_signature(DtlsCoapSession.connect, ["self"]) + connect_timeout = inspect.signature(DtlsCoapSession.connect).parameters[ + "timeout" + ] + assert connect_timeout.kind is inspect.Parameter.KEYWORD_ONLY + assert connect_timeout.default is None _assert_compatible_signature( DtlsCoapSession.get, [ diff --git a/tests/test_session_connect_deadline.py b/tests/test_session_connect_deadline.py new file mode 100644 index 0000000..f9bd0a4 --- /dev/null +++ b/tests/test_session_connect_deadline.py @@ -0,0 +1,344 @@ +"""Deterministic tests for bounded DTLS handshake timing.""" + +from __future__ import annotations + +import socket + +import pytest +from OpenSSL import SSL + +from smartthings_local.errors import SessionTimeoutError +from smartthings_local.protocol import dtls_session +from smartthings_local.protocol.dtls_session import DtlsCoapSession +from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint + + +class _Clock: + def __init__(self): + self.now = 100.0 + + def monotonic(self): + return self.now + + def advance(self, seconds): + self.now += seconds + + +class _Auth: + def __init__(self, clock=None, configure_delay=0.0): + self.clock = clock + self.configure_delay = configure_delay + + def configure_context(self, _context): + if self.clock is not None: + self.clock.advance(self.configure_delay) + + +class _Connection: + def __init__(self, outcomes=None, outputs=None, timer=None): + self.outcomes = list(outcomes or ()) + self.outputs = list(outputs or ()) + self.timer = timer + self.bio_writes = [] + self.timeout_calls = 0 + + def set_connect_state(self): + return None + + def set_ciphertext_mtu(self, _mtu): + return None + + def do_handshake(self): + outcome = self.outcomes.pop(0) if self.outcomes else "want-read" + if outcome == "want-read": + raise SSL.WantReadError() + if isinstance(outcome, Exception): + raise outcome + + def bio_read(self, _size): + if self.outputs: + return self.outputs.pop(0) + raise SSL.WantReadError() + + def bio_write(self, data): + self.bio_writes.append(data) + + def DTLSv1_get_timeout(self): + return self.timer + + def DTLSv1_handle_timeout(self): + self.timeout_calls += 1 + + +class _Socket: + def __init__(self, clock, inbound=()): + self.clock = clock + self.inbound = list(inbound) + self.timeouts = [] + self.sent = [] + self.closed = False + + def settimeout(self, timeout): + self.timeouts.append(timeout) + + def send(self, data): + self.sent.append(data) + return len(data) + + def recv(self, _size): + if self.inbound: + result = self.inbound.pop(0) + if isinstance(result, Exception): + self.clock.advance(self.timeouts[-1]) + raise result + return result + self.clock.advance(self.timeouts[-1]) + raise TimeoutError() + + def close(self): + self.closed = True + + +def _session(auth=None): + return DtlsCoapSession( + "device.example", + 5684, + auth=auth or _Auth(), + ) + + +def _install_handshake( + monkeypatch, + clock, + *, + outcomes=(), + outputs=(), + inbound=(), + timer=None, +): + connection = _Connection(outcomes, outputs, timer) + sock = _Socket(clock, inbound) + endpoint = ResolvedUdpEndpoint( + socket.AF_INET, + ("192.0.2.10", 5684), + ) + open_calls = [] + + def open_socket(*args, **kwargs): + open_calls.append((args, kwargs)) + sock.settimeout(kwargs["timeout"]) + return sock, endpoint + + monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object()) + monkeypatch.setattr( + dtls_session.SSL, + "Connection", + lambda *_args: connection, + ) + monkeypatch.setattr( + dtls_session, + "open_connected_udp_socket", + open_socket, + ) + monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic) + monkeypatch.setattr(dtls_session.time, "sleep", clock.advance) + monkeypatch.setattr( + dtls_session.time, + "time", + lambda: pytest.fail("wall clock must not control handshake deadlines"), + ) + return connection, sock, endpoint, open_calls + + +@pytest.mark.parametrize("timeout", (True, "1", object())) +def test_connect_timeout_type_is_explicit(timeout): + with pytest.raises(TypeError, match="number or None"): + _session().connect(timeout=timeout) + + +@pytest.mark.parametrize( + "timeout", + ( + 0, + -1, + float("nan"), + float("inf"), + float("-inf"), + 10**1000, + ), +) +def test_connect_timeout_must_be_positive_and_finite(timeout): + with pytest.raises(ValueError, match="positive finite"): + _session().connect(timeout=timeout) + + +def test_connect_timeout_caps_every_blocking_poll(monkeypatch): + clock = _Clock() + _connection, sock, _endpoint, open_calls = _install_handshake( + monkeypatch, + clock, + ) + session = _session() + + with pytest.raises(SessionTimeoutError): + session.connect(timeout=4.75) + + assert clock.now == pytest.approx(104.75) + assert sock.closed + assert open_calls == [ + ( + ("device.example", 5684), + { + "family": socket.AF_UNSPEC, + "local_port": None, + "timeout": 0.5, + }, + ) + ] + assert max(sock.timeouts) <= 0.5 + assert sock.timeouts[-1] == pytest.approx(0.25) + + +def test_short_timeout_is_not_rounded_up_to_poll_interval(monkeypatch): + clock = _Clock() + _connection, sock, _endpoint, open_calls = _install_handshake( + monkeypatch, + clock, + ) + + with pytest.raises(SessionTimeoutError): + _session().connect(timeout=0.125) + + assert clock.now == pytest.approx(100.125) + assert open_calls[0][1]["timeout"] == pytest.approx(0.125) + assert sock.timeouts == pytest.approx([0.125, 0.125]) + + +def test_default_timeout_uses_session_constant(monkeypatch): + clock = _Clock() + _connection, sock, _endpoint, _open_calls = _install_handshake( + monkeypatch, + clock, + ) + session = _session() + session.HANDSHAKE_TIMEOUT_S = 0.2 + + with pytest.raises(SessionTimeoutError): + session.connect() + + assert clock.now == pytest.approx(100.2) + assert sock.closed + + +def test_context_setup_consumes_the_same_deadline(monkeypatch): + clock = _Clock() + socket_opened = False + + def open_socket(*_args, **_kwargs): + nonlocal socket_opened + socket_opened = True + raise AssertionError("expired setup must not open a socket") + + monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object()) + monkeypatch.setattr(dtls_session.SSL, "Connection", lambda *_args: _Connection()) + monkeypatch.setattr(dtls_session, "open_connected_udp_socket", open_socket) + monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic) + session = _session(_Auth(clock, configure_delay=0.2)) + + with pytest.raises(SessionTimeoutError): + session.connect(timeout=0.1) + + assert not socket_opened + + +def test_socket_setup_consumes_the_same_deadline(monkeypatch): + clock = _Clock() + connection = _Connection() + connection.do_handshake = lambda: pytest.fail( + "expired socket setup must not start a handshake" + ) + sock = _Socket(clock) + endpoint = ResolvedUdpEndpoint( + socket.AF_INET, + ("192.0.2.10", 5684), + ) + + def open_socket(*_args, **kwargs): + sock.settimeout(kwargs["timeout"]) + clock.advance(0.2) + return sock, endpoint + + monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object()) + monkeypatch.setattr( + dtls_session.SSL, + "Connection", + lambda *_args: connection, + ) + monkeypatch.setattr(dtls_session, "open_connected_udp_socket", open_socket) + monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic) + + with pytest.raises(SessionTimeoutError): + _session().connect(timeout=0.1) + + assert sock.closed + + +def test_successful_handshake_preserves_connected_session_state(monkeypatch): + clock = _Clock() + connection, sock, endpoint, _open_calls = _install_handshake( + monkeypatch, + clock, + outcomes=("want-read", "success"), + inbound=(b"synthetic server flight",), + ) + session = _session() + + session.connect(timeout=1.0) + + assert connection.bio_writes == [b"synthetic server flight"] + assert session.conn is connection + assert session.sock is sock + assert session.endpoint is endpoint + assert session.dest == endpoint.sockaddr + assert not sock.closed + + +def test_connect_services_openssl_retransmit_timer(monkeypatch): + clock = _Clock() + outbound = b"\x16\xfe\xfd" + b"\x00" * 8 + b"\x00\x01x" + connection, sock, _endpoint, _open_calls = _install_handshake( + monkeypatch, + clock, + outcomes=("want-read", "want-read", "success"), + outputs=(outbound, outbound), + inbound=(TimeoutError(), b"synthetic server flight"), + timer=0.0, + ) + + _session().connect(timeout=2.0) + + assert connection.timeout_calls == 1 + assert sock.sent == [outbound, outbound] + assert connection.bio_writes == [b"synthetic server flight"] + + +@pytest.mark.parametrize("success_delay", (0.1, 0.2)) +def test_handshake_success_at_or_after_deadline_is_rejected( + monkeypatch, + success_delay, +): + clock = _Clock() + connection, sock, _endpoint, _open_calls = _install_handshake( + monkeypatch, + clock, + ) + + def late_success(): + clock.advance(success_delay) + + connection.do_handshake = late_success + + with pytest.raises(SessionTimeoutError): + _session().connect(timeout=0.1) + + assert sock.closed