feat(protocol): add cancellable session interruption

This commit is contained in:
Jason Morcos
2026-08-14 14:37:30 -07:00
parent a44930f9df
commit 3f0e437880
5 changed files with 455 additions and 17 deletions
+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()