From 3f0e43788087ab9b6019f9a687c1e776a201363a Mon Sep 17 00:00:00 2001 From: Jason Morcos Date: Sun, 9 Aug 2026 12:05:08 -0700 Subject: [PATCH] feat(protocol): add cancellable session interruption --- README.md | 16 ++ smartthings_local/protocol/dtls_handshake.py | 33 ++- smartthings_local/protocol/dtls_session.py | 130 ++++++++- tests/test_public_api_contract.py | 13 +- tests/test_session_interruption.py | 280 +++++++++++++++++++ 5 files changed, 455 insertions(+), 17 deletions(-) create mode 100644 tests/test_session_interruption.py diff --git a/README.md b/README.md index 1355ef1..b0b279d 100644 --- a/README.md +++ b/README.md @@ -57,6 +57,22 @@ retransmissions within that same deadline: sess.connect(timeout=4.0) ``` +Connection attempts can also use a one-way cancellation signal. The signal is +backed by a socketpair, so setting it wakes the network wait immediately while +OpenSSL retains control of DTLS retransmission timing: + +```python +from smartthings_local.protocol.dtls_session import ConnectCancellation + +cancel_connect = ConnectCancellation() +# Another thread may call cancel_connect.set(). +sess.connect(timeout=8.0, cancel=cancel_connect) +``` + +Setting the signal stops subscribed connection attempts and closes their +temporary UDP sockets. It does not alter an already established session or add +new session lifecycle methods. Interrupted attempts raise `SessionClosedError`. + 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 index 8280e3b..ffc8e63 100644 --- a/smartthings_local/protocol/dtls_handshake.py +++ b/smartthings_local/protocol/dtls_handshake.py @@ -2,6 +2,7 @@ from __future__ import annotations +import select import time from collections.abc import Callable @@ -13,6 +14,10 @@ _HANDSHAKE_POLL_S = 0.5 _MAX_DATAGRAM_SIZE = 65535 +class _HandshakeCancelled(Exception): + """Internal signal that a handshake wake socket became readable.""" + + def _drive_dtls_handshake( connection, sock, @@ -20,6 +25,7 @@ def _drive_dtls_handshake( deadline: float, retries: int | None = None, on_datagram: Callable[[bytes], None] | None = None, + wake_socket=None, ) -> bool: """Drive one memory-BIO DTLS handshake up to a monotonic deadline. @@ -51,10 +57,29 @@ def _drive_dtls_handshake( remaining = deadline - time.monotonic() if remaining <= 0: break - sock.settimeout(min(_HANDSHAKE_POLL_S, remaining)) - try: - datagram = sock.recv(_MAX_DATAGRAM_SIZE) - except TimeoutError: + wait = min(_HANDSHAKE_POLL_S, remaining) + timed_out = False + if wake_socket is None: + sock.settimeout(wait) + try: + datagram = sock.recv(_MAX_DATAGRAM_SIZE) + except TimeoutError: + timed_out = True + else: + readable, _, _ = select.select( + (sock, wake_socket), + (), + (), + wait, + ) + if wake_socket in readable: + raise _HandshakeCancelled() + if sock not in readable: + timed_out = True + else: + datagram = sock.recv(_MAX_DATAGRAM_SIZE) + + if timed_out: timer = connection.DTLSv1_get_timeout() if timer is not None and timer <= 0: if retries is not None and retransmits >= retries: diff --git a/smartthings_local/protocol/dtls_session.py b/smartthings_local/protocol/dtls_session.py index 910bbe1..7165ed0 100644 --- a/smartthings_local/protocol/dtls_session.py +++ b/smartthings_local/protocol/dtls_session.py @@ -49,7 +49,11 @@ from .auth import ( _OCF_ROOT_CA, _load_pem_chain, ) -from .dtls_handshake import _HANDSHAKE_POLL_S, _drive_dtls_handshake +from .dtls_handshake import ( + _HANDSHAKE_POLL_S, + _HandshakeCancelled, + _drive_dtls_handshake, +) from .endpoint import open_connected_udp_socket import logging @@ -104,6 +108,61 @@ def _validate_handshake_timeout(timeout, default): raise ValueError('timeout must be a positive finite number or None') return value + +class ConnectCancellation: + """One-way, socket-backed cancellation signal for ``connect()``. + + Each active connection attempt receives its own wake socket. ``set()`` + makes every subscribed socket readable immediately, without a polling + thread or a session-level abort API. + """ + + __slots__ = ("_is_set", "_lock", "_writers") + + def __init__(self) -> None: + self._is_set = False + self._lock = threading.Lock() + self._writers: set[socket.socket] = set() + + def set(self) -> None: + """Cancel current and future connection attempts using this signal.""" + with self._lock: + if self._is_set: + return + self._is_set = True + for writer in self._writers: + try: + writer.send(b"\0") + except OSError: + pass + + def is_set(self) -> bool: + """Return whether cancellation has been requested.""" + with self._lock: + return self._is_set + + def _subscribe(self) -> tuple[socket.socket, socket.socket]: + reader, writer = socket.socketpair() + reader.setblocking(False) + with self._lock: + self._writers.add(writer) + if self._is_set: + writer.send(b"\0") + return reader, writer + + def _unsubscribe( + self, + reader: socket.socket, + writer: socket.socket, + ) -> bool: + with self._lock: + self._writers.discard(writer) + interrupted = self._is_set + reader.close() + writer.close() + return interrupted + + class DtlsCoapSession: """Single sustained DTLS-CoAP session. @@ -215,22 +274,37 @@ class DtlsCoapSession: # ---- lifecycle --------------------------------------------------- - def connect(self, *, timeout: float | None = None): - """Perform a DTLS handshake within a monotonic deadline. + def connect( + self, + *, + timeout: float | None = None, + cancel: ConnectCancellation | None = None, + ): + """Perform a cancellable 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. + A ``ConnectCancellation`` wakes the network wait immediately and does + not alter an already established session. """ handshake_timeout = _validate_handshake_timeout( timeout, self.HANDSHAKE_TIMEOUT_S) + if cancel is not None and not isinstance(cancel, ConnectCancellation): + raise TypeError("cancel must be a ConnectCancellation or None") + if cancel is not None and cancel.is_set(): + raise SessionClosedError() deadline = time.monotonic() + handshake_timeout ctx = SSL.Context(SSL.DTLS_METHOD) self.auth.configure_context(ctx) + if cancel is not None and cancel.is_set(): + raise SessionClosedError() conn = SSL.Connection(ctx, None) conn.set_connect_state() conn.set_ciphertext_mtu(self.mtu) + if cancel is not None and cancel.is_set(): + raise SessionClosedError() remaining = deadline - time.monotonic() if remaining <= 0: @@ -243,19 +317,51 @@ class DtlsCoapSession: timeout=min(_HANDSHAKE_POLL_S, remaining), ) dest = endpoint.sockaddr + if cancel is not None and cancel.is_set(): + sock.close() + raise SessionClosedError() + + wake_subscription = None + subscription_failed = False + if cancel is not None: + try: + wake_subscription = cancel._subscribe() + except OSError: + subscription_failed = True + if subscription_failed: + sock.close() + raise SessionError() from OSError( + "connection cancellation setup failed" + ) backend_failed = False io_failed = False + cancelled = False + interrupted = False try: - completed = _drive_dtls_handshake( - conn, - sock, - deadline=deadline, - ) - except SSL.Error: - backend_failed = True - except OSError: - io_failed = True + try: + completed = _drive_dtls_handshake( + conn, + sock, + deadline=deadline, + wake_socket=( + wake_subscription[0] + if wake_subscription is not None + else None + ), + ) + except _HandshakeCancelled: + cancelled = True + except SSL.Error: + backend_failed = True + except OSError: + io_failed = True + finally: + if wake_subscription is not None: + interrupted = cancel._unsubscribe(*wake_subscription) + if cancelled or interrupted: + sock.close() + raise SessionClosedError() if backend_failed: sock.close() raise SessionError() from ConnectionError('DTLS backend failed') diff --git a/tests/test_public_api_contract.py b/tests/test_public_api_contract.py index 5abf19a..ee96d3c 100644 --- a/tests/test_public_api_contract.py +++ b/tests/test_public_api_contract.py @@ -11,7 +11,10 @@ from smartthings_local.protocol.auth import ( CertificateAuth, PskAuth, ) -from smartthings_local.protocol.dtls_session import DtlsCoapSession +from smartthings_local.protocol.dtls_session import ( + ConnectCancellation, + DtlsCoapSession, +) def _assert_compatible_signature(callable_object, expected: list[str]) -> None: @@ -81,12 +84,20 @@ def test_dtls_session_keeps_current_consumer_methods(): "subscribe", } assert expected <= set(dir(DtlsCoapSession)) + assert "abort" not in DtlsCoapSession.__dict__ + assert "quiesce_for_close" not in DtlsCoapSession.__dict__ _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 + connect_cancel = inspect.signature(DtlsCoapSession.connect).parameters[ + "cancel" + ] + assert connect_cancel.kind is inspect.Parameter.KEYWORD_ONLY + assert connect_cancel.default is None + assert callable(ConnectCancellation().set) _assert_compatible_signature( DtlsCoapSession.get, [ diff --git a/tests/test_session_interruption.py b/tests/test_session_interruption.py new file mode 100644 index 0000000..519d673 --- /dev/null +++ b/tests/test_session_interruption.py @@ -0,0 +1,280 @@ +"""Deterministic tests for connection-attempt cancellation.""" + +from __future__ import annotations + +import select +import socket +import threading +import time +import traceback + +import pytest +from OpenSSL import SSL + +from smartthings_local.errors import SessionClosedError, SessionError +from smartthings_local.protocol import dtls_session +from smartthings_local.protocol.dtls_session import ( + ConnectCancellation, + DtlsCoapSession, +) +from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint + + +class _Auth: + def __init__(self, on_configure=None): + self.on_configure = on_configure + + def configure_context(self, _context): + if self.on_configure is not None: + self.on_configure() + + +class _Connection: + def __init__(self, *, started=None, on_success=None, succeed=False): + self.started = started + self.on_success = on_success + self.succeed = succeed + self.bio_writes = [] + self.handshake_calls = 0 + + def set_connect_state(self): + return None + + def set_ciphertext_mtu(self, _mtu): + return None + + def do_handshake(self): + self.handshake_calls += 1 + if self.started is not None: + self.started.set() + if self.succeed: + if self.on_success is not None: + self.on_success() + return + raise SSL.WantReadError() + + def bio_read(self, _size): + raise SSL.WantReadError() + + def bio_write(self, datagram): + self.bio_writes.append(datagram) + self.succeed = True + + def DTLSv1_get_timeout(self): + return None + + def shutdown(self): + return None + + +def _session(auth=None): + return DtlsCoapSession( + "device.example", + 5684, + auth=auth or _Auth(), + ) + + +def _install_connection(monkeypatch, connection, data_socket, *, on_open=None): + endpoint = ResolvedUdpEndpoint( + socket.AF_INET, + ("192.0.2.10", 5684), + ) + + def open_socket(*_args, **_kwargs): + if on_open is not None: + on_open() + return data_socket, 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, + ) + return endpoint + + +def _run_connect(session, cancel): + outcome = {} + + def worker(): + try: + session.connect(timeout=2.0, cancel=cancel) + except Exception as error: # noqa: BLE001 - captured for assertion + outcome["error"] = error + else: + outcome["connected"] = True + + thread = threading.Thread(target=worker) + thread.start() + return thread, outcome + + +@pytest.mark.parametrize("cancel", (True, threading.Event(), object(), "signal")) +def test_connect_cancel_type_is_explicit(cancel): + with pytest.raises(TypeError, match="ConnectCancellation or None"): + _session().connect(cancel=cancel) + + +def test_pre_cancelled_connect_stops_before_context_setup(monkeypatch): + cancel = ConnectCancellation() + cancel.set() + monkeypatch.setattr( + dtls_session.SSL, + "Context", + lambda *_args: pytest.fail("cancelled connect configured TLS"), + ) + + with pytest.raises(SessionClosedError): + _session().connect(cancel=cancel) + + +def test_cancel_during_context_setup_stops_before_socket_setup(monkeypatch): + cancel = ConnectCancellation() + monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object()) + monkeypatch.setattr( + dtls_session, + "open_connected_udp_socket", + lambda *_args, **_kwargs: pytest.fail( + "cancelled connect opened a socket" + ), + ) + + with pytest.raises(SessionClosedError): + _session(_Auth(cancel.set)).connect(cancel=cancel) + + +def test_cancel_during_socket_setup_closes_before_handshake(monkeypatch): + cancel = ConnectCancellation() + connection = _Connection(succeed=True) + data_socket, peer = socket.socketpair() + _install_connection( + monkeypatch, + connection, + data_socket, + on_open=cancel.set, + ) + + try: + with pytest.raises(SessionClosedError): + _session().connect(cancel=cancel) + assert data_socket.fileno() == -1 + assert connection.handshake_calls == 0 + finally: + peer.close() + + +def test_socket_signal_wakes_every_subscribed_waiter(): + cancel = ConnectCancellation() + first = cancel._subscribe() + second = cancel._subscribe() + try: + cancel.set() + readable, _, _ = select.select( + (first[0], second[0]), + (), + (), + 0, + ) + assert set(readable) == {first[0], second[0]} + assert cancel._unsubscribe(*first) + assert cancel._unsubscribe(*second) + finally: + for reader, writer in (first, second): + reader.close() + writer.close() + + +def test_cancel_wakes_blocked_connect_without_poll_latency(monkeypatch): + cancel = ConnectCancellation() + started = threading.Event() + connection = _Connection(started=started) + data_socket, peer = socket.socketpair() + _install_connection(monkeypatch, connection, data_socket) + session = _session() + thread, outcome = _run_connect(session, cancel) + + try: + assert started.wait(1.0) + before = time.monotonic() + cancel.set() + thread.join(1.0) + elapsed = time.monotonic() - before + + assert not thread.is_alive() + assert elapsed < 0.25 + assert isinstance(outcome.get("error"), SessionClosedError) + assert data_socket.fileno() == -1 + assert session.sock is None + assert session.conn is None + assert not cancel._writers + finally: + cancel.set() + thread.join(1.0) + peer.close() + + +def test_cancel_wins_race_with_reported_handshake_success(monkeypatch): + cancel = ConnectCancellation() + connection = _Connection(on_success=cancel.set, succeed=True) + data_socket, peer = socket.socketpair() + _install_connection(monkeypatch, connection, data_socket) + session = _session() + + try: + with pytest.raises(SessionClosedError): + session.connect(cancel=cancel) + assert data_socket.fileno() == -1 + assert session.sock is None + assert session.conn is None + finally: + peer.close() + + +def test_successful_connect_does_not_set_cancel_or_close_session(monkeypatch): + cancel = ConnectCancellation() + connection = _Connection() + data_socket, peer = socket.socketpair() + endpoint = _install_connection(monkeypatch, connection, data_socket) + peer.send(b"synthetic server flight") + session = _session() + + try: + session.connect(cancel=cancel) + assert not cancel.is_set() + assert session.sock is data_socket + assert session.conn is connection + assert session.endpoint is endpoint + assert connection.bio_writes == [b"synthetic server flight"] + assert not cancel._writers + finally: + session.close() + peer.close() + + +def test_cancellation_socket_failure_is_redacted_and_closes_udp(monkeypatch): + class FailingCancellation(ConnectCancellation): + def _subscribe(self): + raise OSError("credential-value at device.example") + + connection = _Connection() + data_socket, peer = socket.socketpair() + _install_connection(monkeypatch, connection, data_socket) + + try: + with pytest.raises(SessionError) as exc: + _session().connect(cancel=FailingCancellation()) + + formatted = "".join(traceback.format_exception(exc.value)) + assert data_socket.fileno() == -1 + assert exc.value.__context__ is None + assert "credential-value" not in formatted + assert "device.example" not in formatted + finally: + peer.close()