feat(protocol): add cancellable session interruption
This commit is contained in:
@@ -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:
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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,
|
||||
[
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user