Merge pull request #35 from Moballo-LLC/codex/py-08b-session-interruption

Add cancellable DTLS connection attempts
This commit is contained in:
Quite Yellow
2026-08-15 10:20:05 +01:00
committed by GitHub
5 changed files with 455 additions and 17 deletions
+16
View File
@@ -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:
+29 -4
View File
@@ -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:
+118 -12
View File
@@ -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')
+12 -1
View File
@@ -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,
[
+280
View File
@@ -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()