From d33f896d84259a8ac6bcaaaad369bbfbd4c52c75 Mon Sep 17 00:00:00 2001 From: Jason Morcos Date: Fri, 28 Aug 2026 12:00:14 -0700 Subject: [PATCH 1/3] feat(protocol): add two-phase session shutdown --- README.md | 22 +- smartthings_local/protocol/dtls_session.py | 199 ++++++++++++---- tests/test_public_api_contract.py | 9 +- tests/test_session_interruption.py | 87 ++++++- tests/test_session_shutdown.py | 256 +++++++++++++++++++++ 5 files changed, 521 insertions(+), 52 deletions(-) create mode 100644 tests/test_session_shutdown.py diff --git a/README.md b/README.md index 7fc547c..f9bcc46 100644 --- a/README.md +++ b/README.md @@ -102,8 +102,26 @@ 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`. +temporary UDP sockets. It does not alter an already established session. +Interrupted attempts raise `SessionClosedError`. + +Hosts that stop network work before their blocking executor drains can use the +session's two-phase shutdown. `quiesce_for_close()` is terminal: it interrupts +an in-progress handshake, wakes pending requests and notification refetches, +and rejects new work while retaining an established DTLS socket. A subsequent +`close()` flushes the authenticated close-notify record before closing that +socket: + +```python +sess.quiesce_for_close() # safe from the host's early shutdown phase +# Later, after session workers have joined: +sess.close() +``` + +Use `abort()` when orderly shutdown is impossible. It performs the same +terminal wakeup but closes the established socket immediately, without waiting +for close-notify. All three methods are idempotent; a quiesced or aborted +session cannot be connected again. Reads retransmit each Block2 request; writes send once. Where a lost write has been shown to be the cause rather than a device that is simply refusing diff --git a/smartthings_local/protocol/dtls_session.py b/smartthings_local/protocol/dtls_session.py index 3f156c9..6da5874 100644 --- a/smartthings_local/protocol/dtls_session.py +++ b/smartthings_local/protocol/dtls_session.py @@ -261,7 +261,8 @@ class ConnectCancellation: 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. + thread. Session shutdown shares that wake path without changing a + caller-supplied signal. """ __slots__ = ("_is_set", "_lock", "_writers") @@ -291,20 +292,33 @@ def is_set(self) -> bool: def _subscribe(self) -> tuple[socket.socket, socket.socket]: reader, writer = socket.socketpair() reader.setblocking(False) + try: + self._subscribe_writer(writer) + except Exception: + reader.close() + writer.close() + raise + return reader, writer + + def _subscribe_writer(self, writer: socket.socket) -> None: + """Attach an existing wake writer without taking ownership of it.""" with self._lock: self._writers.add(writer) if self._is_set: writer.send(b"\0") - return reader, writer + + def _unsubscribe_writer(self, writer: socket.socket) -> bool: + """Detach a shared wake writer and report cancellation state.""" + with self._lock: + self._writers.discard(writer) + return self._is_set def _unsubscribe( self, reader: socket.socket, writer: socket.socket, ) -> bool: - with self._lock: - self._writers.discard(writer) - interrupted = self._is_set + interrupted = self._unsubscribe_writer(writer) reader.close() writer.close() return interrupted @@ -389,6 +403,13 @@ def __init__(self, host, port, cert_path=None, key_path=None, *, self.dest = None self.endpoint = None + # Terminal session shutdown has its own wake signal so + # quiesce_for_close() can interrupt connect() even when the caller did + # not supply a ConnectCancellation. The lifecycle lock makes handshake + # publication atomic with that one-way transition. + self._lifecycle_cancel = ConnectCancellation() + self._lifecycle_lock = threading.Lock() + self._send_lock = threading.Lock() # Guards the MID/token counters and pending-request registries. # The refetch worker makes the session its own second concurrent @@ -461,18 +482,21 @@ def connect( 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(): + if self._lifecycle_cancel.is_set() or \ + (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(): + if self._lifecycle_cancel.is_set() or \ + (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(): + if self._lifecycle_cancel.is_set() or \ + (cancel is not None and cancel.is_set()): raise SessionClosedError() remaining = deadline - time.monotonic() @@ -486,18 +510,32 @@ def connect( timeout=min(_HANDSHAKE_POLL_S, remaining), ) dest = endpoint.sockaddr - if cancel is not None and cancel.is_set(): + if self._lifecycle_cancel.is_set() or \ + (cancel is not None and cancel.is_set()): sock.close() raise SessionClosedError() wake_subscription = None subscription_failed = False - if cancel is not None: + wake_owner = cancel or self._lifecycle_cancel + lifecycle_writer_shared = cancel is not None + # Production endpoints are real sockets and can share select() with + # the wake socket. Retain the timeout-driven path for structural socket + # adapters that intentionally expose no file descriptor. + if callable(getattr(sock, "fileno", None)): try: - wake_subscription = cancel._subscribe() + wake_subscription = wake_owner._subscribe() + if lifecycle_writer_shared: + self._lifecycle_cancel._subscribe_writer( + wake_subscription[1]) except OSError: subscription_failed = True if subscription_failed: + if wake_subscription is not None: + if lifecycle_writer_shared: + self._lifecycle_cancel._unsubscribe_writer( + wake_subscription[1]) + wake_owner._unsubscribe(*wake_subscription) sock.close() raise SessionError() from OSError( "connection cancellation setup failed" @@ -528,7 +566,13 @@ def connect( io_failed = True finally: if wake_subscription is not None: - interrupted = cancel._unsubscribe(*wake_subscription) + if lifecycle_writer_shared: + interrupted = self._lifecycle_cancel._unsubscribe_writer( + wake_subscription[1]) + interrupted = ( + wake_owner._unsubscribe(*wake_subscription) + or interrupted + ) if cancelled or (interrupted and not completed): sock.close() raise SessionClosedError() @@ -542,21 +586,27 @@ def connect( sock.close() raise SessionTimeoutError() - self.sock = sock - self.conn = conn - self.dest = dest - self.endpoint = endpoint - self._stop.clear() + with self._lifecycle_lock: + if self._lifecycle_cancel.is_set(): + sock.close() + raise SessionClosedError() + self.sock = sock + self.conn = conn + self.dest = dest + self.endpoint = endpoint + self._stop.clear() def start_reader(self): """Spawn the reader thread. Must be called after connect().""" - if self.sock is None: - raise RuntimeError("connect() before start_reader()") - self._reader_running.set() - t = threading.Thread(target=self._reader_loop, - daemon=True, name='dtls-reader') - t.start() - self._reader_thread = t + with self._lifecycle_lock: + if self.sock is None: + raise RuntimeError("connect() before start_reader()") + self._check_live() + self._reader_running.set() + t = threading.Thread(target=self._reader_loop, + daemon=True, name='dtls-reader') + t.start() + self._reader_thread = t def _check_live(self): """Raise if the session cannot carry a request. A dead reader is @@ -567,7 +617,7 @@ def _check_live(self): Callers that never start a reader (config-flow style) keep the old behaviour — only the conn check applies while _reader_thread is None.""" - if self.conn is None: + if self._lifecycle_cancel.is_set() or self.conn is None: raise SessionClosedError() if self._reader_thread is not None and \ not self._reader_running.is_set(): @@ -592,6 +642,43 @@ def _send_observe_dereg(self, tok, path_segs): self._send_dgram( build_coap(TYPE_CON, METHOD_GET, mid, tok, opts)) + @staticmethod + def _send_close_notify(connection, sock): + """Best-effort flush the encrypted DTLS close-notify record.""" + try: + connection.shutdown() + except (SSL.WantReadError, SSL.ZeroReturnError): + pass + except Exception: + pass + try: + while True: + try: + outbound = connection.bio_read(65535) + except SSL.WantReadError: + break + if not outbound: + break + for record in _split_dtls(outbound): + if sock.send(record) != len(record): + raise OSError('incomplete UDP send') + except Exception: + pass + + def quiesce_for_close(self): + """Stop new work while retaining an established socket for close().""" + with self._lifecycle_lock: + self._lifecycle_cancel.set() + # Synchronize with an in-progress application send. Once this lock + # is released, _send_dgram() observes lifecycle cancellation before + # touching DTLS. + with self._send_lock: + self._stop.set() + with self._refetch_cond: + self._refetch_pending.clear() + self._refetch_cond.notify_all() + self._close_pending_requests() + def close(self): """Tear down session. Sends best-effort OBSERVE deregisters first so Samsung's RT-OCF cleans up its observer table — @@ -600,7 +687,8 @@ def close(self): # Send dereg for every active observation while the conn is # still healthy. Tiny sleep lets the records reach the wire # before we shut DTLS down. - if self.conn is not None and self._observe_tokens: + if (not self._lifecycle_cancel.is_set() and self.conn is not None + and self._observe_tokens): for tok, href in list(self._observe_tokens.items()): segs = [s for s in href.split('/') if s] try: @@ -609,17 +697,10 @@ def close(self): logger.warning("dereg %s: %s", href, e) time.sleep(0.1) - self._stop.set() - # Wake the refetch worker so it sees _stop instead of sitting on - # its condition for up to a second after the socket is gone. - with self._refetch_cond: - self._refetch_pending.clear() - self._refetch_cond.notify_all() - if self.conn is not None: - try: - self.conn.shutdown() - except Exception: - pass + self.quiesce_for_close() + with self._send_lock: + if self.conn is not None and self.sock is not None: + self._send_close_notify(self.conn, self.sock) if self.sock is not None: try: self.sock.close() @@ -629,12 +710,30 @@ def close(self): # that passed its entry check just before close() will then fail the # post-registration liveness check instead of registering after the # drain and waiting against a session that can no longer respond. - self.sock = None - self.conn = None - self.dest = None - self.endpoint = None - self._close_pending_requests() - self._observe_tokens.clear() + with self._lifecycle_lock: + self.sock = None + self.conn = None + self.dest = None + self.endpoint = None + with self._state_lock: + self._observe_tokens.clear() + + def abort(self): + """Immediately stop work and close the established transport.""" + self.quiesce_for_close() + with self._lifecycle_lock: + sock = self.sock + self.sock = None + self.conn = None + self.dest = None + self.endpoint = None + if sock is not None: + try: + sock.close() + except Exception: + pass + with self._state_lock: + self._observe_tokens.clear() # ---- send / receive plumbing ------------------------------------- @@ -726,7 +825,7 @@ def _send_dgram(self, datagram): """Send a CoAP datagram. Holds the send lock for the BIO-drain so two writers can't interleave records.""" with self._send_lock: - if self.conn is None: + if self._lifecycle_cancel.is_set() or self.conn is None: raise SessionClosedError() send_failed = False try: @@ -1718,11 +1817,17 @@ def subscribe(self, path_segs): # Register the token BEFORE sending — otherwise the device # could respond between send() and the dict insert, and the # reader thread would drop the initial 2.05 as "stale". - self._observe_tokens[tok] = href + with self._state_lock: + self._observe_tokens[tok] = href mid = self._next_mid() opts = [(URI_PATH, s.encode()) for s in path_segs] opts.append((OBSERVE, OBSERVE_REGISTER)) opts.append((ACCEPT, CF_CBOR)) - self._send_dgram( - build_coap(TYPE_CON, METHOD_GET, mid, tok, opts)) + try: + self._send_dgram( + build_coap(TYPE_CON, METHOD_GET, mid, tok, opts)) + except Exception: + with self._state_lock: + self._observe_tokens.pop(tok, None) + raise return tok diff --git a/tests/test_public_api_contract.py b/tests/test_public_api_contract.py index 44d0241..ffce576 100644 --- a/tests/test_public_api_contract.py +++ b/tests/test_public_api_contract.py @@ -176,16 +176,16 @@ def test_dtls_session_keeps_current_consumer_methods(): "delete", "get", "join", + "abort", "pace", "ping", "post", + "quiesce_for_close", "refresh_observes", "start_reader", "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" @@ -198,6 +198,11 @@ def test_dtls_session_keeps_current_consumer_methods(): assert connect_cancel.kind is inspect.Parameter.KEYWORD_ONLY assert connect_cancel.default is None assert callable(ConnectCancellation().set) + _assert_compatible_signature( + DtlsCoapSession.quiesce_for_close, + ["self"], + ) + _assert_compatible_signature(DtlsCoapSession.abort, ["self"]) _assert_compatible_signature( DtlsCoapSession.get, [ diff --git a/tests/test_session_interruption.py b/tests/test_session_interruption.py index 232349b..7cc2396 100644 --- a/tests/test_session_interruption.py +++ b/tests/test_session_interruption.py @@ -100,7 +100,7 @@ def open_socket(*_args, **_kwargs): return endpoint -def _run_connect(session, cancel): +def _run_connect(session, cancel=None): outcome = {} def worker(): @@ -220,6 +220,91 @@ def test_cancel_wakes_blocked_connect_without_poll_latency(monkeypatch): peer.close() +def test_quiesce_wakes_blocked_connect_without_caller_signal(monkeypatch): + 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) + + try: + assert started.wait(1.0) + before = time.monotonic() + session.quiesce_for_close() + 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 session._lifecycle_cancel._writers + finally: + session.abort() + thread.join(1.0) + peer.close() + + +def test_quiesce_wakes_connect_without_setting_caller_signal(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) + session.quiesce_for_close() + thread.join(1.0) + + assert not thread.is_alive() + assert isinstance(outcome.get("error"), SessionClosedError) + assert not cancel.is_set() + assert not cancel._writers + assert not session._lifecycle_cancel._writers + finally: + session.abort() + thread.join(1.0) + peer.close() + + +def test_quiesce_wins_handshake_publication_race(monkeypatch): + handshake_entered = threading.Event() + release_handshake = threading.Event() + + def pause_at_success(): + handshake_entered.set() + assert release_handshake.wait(1.0) + + connection = _Connection(succeed=True, on_success=pause_at_success) + data_socket, peer = socket.socketpair() + _install_connection(monkeypatch, connection, data_socket) + session = _session() + thread, outcome = _run_connect(session) + + try: + assert handshake_entered.wait(1.0) + session.quiesce_for_close() + release_handshake.set() + thread.join(1.0) + + assert not thread.is_alive() + assert isinstance(outcome.get("error"), SessionClosedError) + assert data_socket.fileno() == -1 + assert session.sock is None + assert session.conn is None + finally: + release_handshake.set() + session.abort() + thread.join(1.0) + peer.close() + + def test_reported_handshake_success_wins_cancel_during_unsubscribe( monkeypatch, ): diff --git a/tests/test_session_shutdown.py b/tests/test_session_shutdown.py new file mode 100644 index 0000000..ce47758 --- /dev/null +++ b/tests/test_session_shutdown.py @@ -0,0 +1,256 @@ +"""Two-phase DTLS session shutdown contracts.""" + +from __future__ import annotations + +import threading +import time + +import pytest +from OpenSSL import SSL + +from smartthings_local.errors import SessionClosedError +from smartthings_local.protocol.dtls_session import DtlsCoapSession + + +class _NullAuth: + def configure_context(self, _context): + return None + + +class _CloseNotifyConnection: + def __init__(self): + self.shutdown_calls = 0 + self._outbound = [] + + def shutdown(self): + self.shutdown_calls += 1 + self._outbound.append( + b"\x15\xfe\xfd\x00\x00" + b"\x00" * 6 + + b"\x00\x02\x01\x00" + ) + + def bio_read(self, _size): + if self._outbound: + return self._outbound.pop(0) + raise SSL.WantReadError() + + +class _Socket: + def __init__(self): + self.closed = False + self.sent = [] + + def send(self, datagram): + if self.closed: + raise OSError("closed") + self.sent.append(datagram) + return len(datagram) + + def close(self): + self.closed = True + + +def _session(*, rate_limit_rps=1_000_000): + session = DtlsCoapSession( + "device.example", + 5684, + auth=_NullAuth(), + rate_limit_rps=rate_limit_rps, + ) + session.conn = _CloseNotifyConnection() + session.sock = _Socket() + session.dest = ("192.0.2.10", 5684) + session.endpoint = object() + return session + + +def _run_request(session, operation, started): + outcome = {} + + def send(_datagram): + started.set() + + session._send_dgram = send + + def request(): + try: + if operation == "get": + session.get(["device", "0"], timeout=30.0) + elif operation == "post": + session.post(["mode", "vs", "0"], b"body", timeout=30.0) + else: + session.delete(["oic", "sec", "cred"], timeout=30.0) + except Exception as error: # noqa: BLE001 - captured for assertion + outcome["error"] = error + + thread = threading.Thread(target=request) + thread.start() + return thread, outcome + + +@pytest.mark.parametrize("operation", ("get", "post", "delete")) +def test_quiesce_wakes_pending_requests_and_retains_transport(operation): + session = _session() + connection = session.conn + sock = session.sock + started = threading.Event() + thread, outcome = _run_request(session, operation, started) + + assert started.wait(1.0) + session.quiesce_for_close() + thread.join(1.0) + + assert not thread.is_alive() + assert isinstance(outcome.get("error"), SessionClosedError) + assert session.conn is connection + assert session.sock is sock + assert session._pending == {} + assert session._pending_mids == {} + with pytest.raises(SessionClosedError): + session._check_live() + + session.close() + + +def test_quiesce_wakes_paced_observe_without_registering_or_sending(): + session = _session(rate_limit_rps=1.0) + session._last_send_ts = time.monotonic() + sends = [] + outcome = {} + session._send_dgram = sends.append + + def subscribe(): + try: + session.subscribe(["mode", "vs", "0"]) + except Exception as error: # noqa: BLE001 - captured for assertion + outcome["error"] = error + + thread = threading.Thread(target=subscribe) + thread.start() + time.sleep(0.02) + session.quiesce_for_close() + thread.join(1.0) + + assert not thread.is_alive() + assert isinstance(outcome.get("error"), SessionClosedError) + assert session._observe_tokens == {} + assert sends == [] + session.close() + + +def test_quiesce_during_observe_send_retires_registered_token(): + session = _session() + + def quiesce_before_send(_datagram): + session.quiesce_for_close() + raise SessionClosedError() + + session._send_dgram = quiesce_before_send + + with pytest.raises(SessionClosedError): + session.subscribe(["mode", "vs", "0"]) + + assert session._observe_tokens == {} + session.close() + + +def test_quiesce_blocks_direct_application_send(): + session = _session() + session.quiesce_for_close() + + with pytest.raises(SessionClosedError): + session._send_dgram(b"request") + + assert session.conn.shutdown_calls == 0 + assert session.sock.sent == [] + session.close() + + +def test_quiesce_wakes_idle_refetch_worker(): + session = _session() + thread = threading.Thread(target=session._refetch_loop) + thread.start() + time.sleep(0.02) + + session.quiesce_for_close() + thread.join(0.25) + + assert not thread.is_alive() + assert session._refetch_pending == {} + session.close() + + +def test_close_after_quiesce_flushes_close_notify_without_deregistering(): + session = _session() + connection = session.conn + sock = session.sock + session._observe_tokens = {b"o": "/mode/vs/0"} + session._send_observe_dereg = lambda *_args: pytest.fail( + "quiesced close sent an application request" + ) + + session.quiesce_for_close() + session.close() + + assert connection.shutdown_calls == 1 + assert sock.sent == [ + b"\x15\xfe\xfd\x00\x00" + b"\x00" * 6 + + b"\x00\x02\x01\x00" + ] + assert sock.closed + assert session.sock is None + assert session.conn is None + assert session.dest is None + assert session.endpoint is None + assert session._observe_tokens == {} + + session.close() + assert connection.shutdown_calls == 1 + + +def test_abort_closes_without_close_notify_and_is_idempotent(): + session = _session() + connection = session.conn + sock = session.sock + event = threading.Event() + container = {} + session._register_pending_request(b"token", event, container) + + session.abort() + + assert connection.shutdown_calls == 0 + assert sock.sent == [] + assert sock.closed + assert event.is_set() + assert isinstance(container.get("err"), SessionClosedError) + assert session.sock is None + assert session.conn is None + assert session.dest is None + assert session.endpoint is None + assert session._pending == {} + assert session._pending_mids == {} + + session.abort() + + +def test_quiesced_session_cannot_start_reader(): + session = _session() + session.quiesce_for_close() + + with pytest.raises(SessionClosedError): + session.start_reader() + + assert session._reader_thread is None + session.close() + + +def test_normal_close_still_deregisters_before_quiescing(monkeypatch): + session = _session() + session._observe_tokens = {b"o": "/mode/vs/0"} + calls = [] + session._send_observe_dereg = lambda *args: calls.append(args) + monkeypatch.setattr(time, "sleep", lambda _seconds: None) + + session.close() + + assert calls == [(b"o", ["mode", "vs", "0"])] From 9d38bccc54c8e483828bcdd9feddea5cd08c8618 Mon Sep 17 00:00:00 2001 From: Jason Morcos Date: Fri, 28 Aug 2026 12:10:55 -0700 Subject: [PATCH 2/3] feat(protocol): track confirmed Observe relations --- README.md | 20 +- smartthings_local/protocol/dtls_session.py | 274 ++++++++++++++---- tests/test_dtls_session_post_retry.py | 2 +- tests/test_observe_relations.py | 320 +++++++++++++++++++++ tests/test_public_api_contract.py | 13 + tests/test_session_shutdown.py | 2 +- 6 files changed, 571 insertions(+), 60 deletions(-) create mode 100644 tests/test_observe_relations.py diff --git a/README.md b/README.md index f9bcc46..3e9f619 100644 --- a/README.md +++ b/README.md @@ -37,14 +37,14 @@ auth = CertificateAuth.from_files( sess = DtlsCoapSession( "192.0.2.100", 49154, auth=auth, + on_notification=lambda href, payload: ..., ) sess.connect() sess.start_reader() code, body = sess.get(["device", "0"]) # Block2-aware read code, _ = sess.post(["mode", "vs", "0"], cbor2.dumps({})) # write -sess.subscribe(["operational", "state", "vs", "0"], # OBSERVE - on_notification=lambda href, payload: ...) +sess.subscribe(["operational", "state", "vs", "0"]) # OBSERVE sess.close() ``` @@ -69,6 +69,22 @@ repeatable. `delete()` uses the same path, query, extension-option, timeout, and response contract without sending a request payload. +Observe relations may also include repeated URI-query strings. The same query +is retained for a blockwise notification refetch and for best-effort +deregistration: + +```python +sess.subscribe(["mode", "vs", "0"], query=("if=oic.if.b",)) +``` + +An RFC 7641 relation is confirmed only by a valid Observe response option; +duplicate and stale 24-bit sequence values are not delivered. Some older +Samsung firmware omits that option. For those devices, a plain initial `2.05` +is probationary until a later packet arrives on the same token with a different +Message ID. Optional `on_observe_pending`, `on_legacy_notification`, and +`on_observe_error` constructor callbacks let consumers keep that compatibility +path distinct from confirmed RFC notifications and ordinary polling. + POST bodies through 1024 bytes retain the single-request behavior. Larger bodies use token-stable Block1 requests under one monotonic timeout, include Size1 on the first request, honor a server-requested smaller block size, and diff --git a/smartthings_local/protocol/dtls_session.py b/smartthings_local/protocol/dtls_session.py index 6da5874..d035ede 100644 --- a/smartthings_local/protocol/dtls_session.py +++ b/smartthings_local/protocol/dtls_session.py @@ -133,7 +133,7 @@ # once the ceiling is measured empirically. _DEFAULT_RATE_LIMIT_RPS = 5.0 -# Maximum hrefs held for OBSERVE refetch at once. A notification storm +# Maximum relations held for OBSERVE refetch at once. A notification storm # on more resources than this is already past what the 5/s ceiling can # drain, so the excess is dropped rather than queued indefinitely. _MAX_PENDING_REFETCH = 16 @@ -143,6 +143,13 @@ # is serialized, so a slow one delays only later refetches. _REFETCH_TIMEOUT_S = 15.0 +# RFC 7641 section 3.4 compares the 24-bit Observe value as serial-number +# arithmetic. After 128 seconds without an accepted notification, receipt time +# is allowed to re-establish ordering after a server restart. +_OBSERVE_SEQUENCE_MODULUS = 1 << 24 +_OBSERVE_SEQUENCE_HALF_RANGE = 1 << 23 +_OBSERVE_SEQUENCE_RESET_S = 128.0 + class _EtagChanged(Exception): """Internal: the server's ETag changed partway through a Block2 @@ -351,7 +358,10 @@ def __init__(self, host, port, cert_path=None, key_path=None, *, rate_limit_rps: float = _DEFAULT_RATE_LIMIT_RPS, local_port=None, family=socket.AF_UNSPEC, write_max_attempts: int = 1, - auth: AuthenticationProvider | None = None): + auth: AuthenticationProvider | None = None, + on_legacy_notification=None, + on_observe_pending=None, + on_observe_error=None): file_supplied = cert_path is not None or key_path is not None memory_supplied = cert_pem is not None or key_pem is not None if auth is not None and (file_supplied or memory_supplied): @@ -382,6 +392,9 @@ def __init__(self, host, port, cert_path=None, key_path=None, *, auth = CertificateAuth.from_files(self.cert_path, self.key_path) self.auth = auth self.on_notification = on_notification # fn(href, payload_bytes) + self.on_legacy_notification = on_legacy_notification + self.on_observe_pending = on_observe_pending + self.on_observe_error = on_observe_error self.mtu = mtu self._min_req_interval = 1.0 / rate_limit_rps self._write_max_attempts = max(1, int(write_max_attempts)) @@ -436,9 +449,20 @@ def __init__(self, host, port, cert_path=None, key_path=None, *, self._pending_mids = {} # token (bytes) → href (str) self._observe_tokens = {} - - # OBSERVE refetch queue: href → sequence number of the newest - # notification that asked for it. Drained by a worker thread + # An Observe relation is identified by path plus URI query. Keep the + # exact registration options for refetch and deregistration. + self._observe_queries = {} + # Some Samsung generations return a plain initial 2.05 and later push + # on the same token without RFC 7641's Observe option. One later packet + # with a different MID is required before that relation is trusted. + self._observe_plain_response_mids = {} + self._legacy_observe_tokens = set() + self._legacy_observe_mids = {} + # token → (last accepted 24-bit Observe value, monotonic receipt time) + self._observe_sequences = {} + + # OBSERVE refetch queue: (href, query, legacy) → sequence number of the + # newest notification that asked for it. Drained by a worker thread # because _dispatch_coap cannot block (see _queue_refetch). self._refetch_cond = threading.Condition() self._refetch_pending = {} @@ -630,13 +654,15 @@ def join(self): if self._refetch_thread is not None: self._refetch_thread.join() - def _send_observe_dereg(self, tok, path_segs): + def _send_observe_dereg(self, tok, path_segs, query=()): """Send a single OBSERVE deregister GET (Observe option = 1) on the existing token. Best-effort — caller swallows errors.""" if self.conn is None: return mid = self._next_mid() opts = [(URI_PATH, s.encode()) for s in path_segs] + for value in query: + opts.append((URI_QUERY, value.encode())) opts.append((OBSERVE, OBSERVE_DEREGISTER)) opts.append((ACCEPT, CF_CBOR)) self._send_dgram( @@ -689,10 +715,14 @@ def close(self): # before we shut DTLS down. if (not self._lifecycle_cancel.is_set() and self.conn is not None and self._observe_tokens): - for tok, href in list(self._observe_tokens.items()): + with self._state_lock: + observations = tuple(self._observe_tokens.items()) + observe_queries = dict(self._observe_queries) + for tok, href in observations: segs = [s for s in href.split('/') if s] try: - self._send_observe_dereg(tok, segs) + self._send_observe_dereg( + tok, segs, observe_queries.get(tok, ())) except Exception as e: logger.warning("dereg %s: %s", href, e) time.sleep(0.1) @@ -715,8 +745,7 @@ def close(self): self.conn = None self.dest = None self.endpoint = None - with self._state_lock: - self._observe_tokens.clear() + self._clear_observe_relations() def abort(self): """Immediately stop work and close the established transport.""" @@ -732,8 +761,7 @@ def abort(self): sock.close() except Exception: pass - with self._state_lock: - self._observe_tokens.clear() + self._clear_observe_relations() # ---- send / receive plumbing ------------------------------------- @@ -807,19 +835,45 @@ def _next_tok(self): # avoids collisions across long-running OBSERVE subscriptions. return self._tok_counter.to_bytes(4, 'big') - def _next_observe_tok(self): - with self._state_lock: - # Single-byte tokens for OBSERVE registrations. Samsung - # RT-OCF accepts these but silently drops TKL=4 OBSERVE - # registrations. Counter is randomly seeded per session so - # reconnects don't collide with stale observer state Samsung - # may still be holding from the previous run. + def _next_available_observe_tok_locked(self): + """Return an unused nonzero one-byte Observe token.""" + for _ in range(min(len(self._observe_tokens) + 1, 0xFF)): self._observe_tok_counter = (self._observe_tok_counter + 1) & 0xFF - # Avoid 0x00 — some CoAP stacks treat an all-zero token as - # equivalent to "no token" / empty (TKL=0). if self._observe_tok_counter == 0: self._observe_tok_counter = 1 - return bytes([self._observe_tok_counter]) + token = bytes([self._observe_tok_counter]) + if token not in self._observe_tokens: + return token + raise SessionIdentifierError() + + def _retire_observe_token_locked(self, tok): + """Remove every relation-state index for one Observe token.""" + self._observe_tokens.pop(tok, None) + self._observe_queries.pop(tok, None) + self._observe_plain_response_mids.pop(tok, None) + self._legacy_observe_tokens.discard(tok) + self._legacy_observe_mids.pop(tok, None) + self._observe_sequences.pop(tok, None) + + def _clear_observe_relations(self): + with self._state_lock: + self._observe_tokens.clear() + self._observe_queries.clear() + self._observe_plain_response_mids.clear() + self._legacy_observe_tokens.clear() + self._legacy_observe_mids.clear() + self._observe_sequences.clear() + + @staticmethod + def _observe_sequence_is_fresh(previous, current, received_at): + """Apply RFC 7641's 24-bit serial-number freshness comparison.""" + if previous is None: + return True + previous_value, previous_received_at = previous + if received_at - previous_received_at > _OBSERVE_SEQUENCE_RESET_S: + return True + delta = (current - previous_value) % _OBSERVE_SEQUENCE_MODULUS + return 0 < delta < _OBSERVE_SEQUENCE_HALF_RANGE def _send_dgram(self, datagram): """Send a CoAP datagram. Holds the send lock for the @@ -1007,23 +1061,101 @@ def _dispatch_coap(self, datagram): return # OBSERVE notification? - href = self._observe_tokens.get(tok) + with self._state_lock: + href = self._observe_tokens.get(tok) + observe_query = self._observe_queries.get(tok, ()) if href is not None: if code != 0x45: - logger.warning("observe %s: non-2.05 %s", - href, fmt_code(code)) + with self._state_lock: + self._retire_observe_token_locked(tok) + logger.debug("observe %s: non-2.05 %s", + href, fmt_code(code)) + cb = self.on_observe_error + if cb is not None: + try: + cb(href, code) + except Exception as e: + logger.debug("observe error callback %s: %s", href, e) return + block_values = [ + value for number, value in ropts if number == BLOCK2 + ] + blockwise_refetch = False + if block_values: + if len(block_values) != 1 or len(block_values[0]) > 3: + logger.debug("observe %s: malformed Block2 option", href) + return + try: + block_number, more, _ = block_fields(block_values[0]) + except ValueError: + logger.debug("observe %s: malformed Block2 option", href) + return + blockwise_refetch = bool(more or block_number) + observe_values = [ + value for number, value in ropts if number == OBSERVE + ] + legacy = False + if observe_values: + if len(observe_values) != 1 or len(observe_values[0]) > 3: + logger.debug("observe %s: malformed Observe option", href) + return + sequence = int.from_bytes(observe_values[0], 'big') + received_at = time.monotonic() + with self._state_lock: + previous = self._observe_sequences.get(tok) + if not self._observe_sequence_is_fresh( + previous, sequence, received_at): + return + self._observe_sequences[tok] = (sequence, received_at) + self._observe_plain_response_mids.pop(tok, None) + self._legacy_observe_tokens.discard(tok) + self._legacy_observe_mids.pop(tok, None) + else: + pending = False + with self._state_lock: + if tok in self._observe_sequences: + return + initial_mid = self._observe_plain_response_mids.get(tok) + if tok in self._legacy_observe_tokens: + if self._legacy_observe_mids.get(tok) == mid: + return + self._legacy_observe_mids[tok] = mid + legacy = True + elif initial_mid is None: + self._observe_plain_response_mids[tok] = mid + pending = True + elif initial_mid == mid: + return + else: + self._legacy_observe_tokens.add(tok) + self._legacy_observe_mids[tok] = mid + legacy = True + if pending: + logger.debug( + "observe %s: probationary 2.05 without Observe option", + href, + ) + cb = self.on_observe_pending + if cb is not None: + try: + cb(href) + except Exception as e: + logger.debug( + "observe pending callback %s: %s", href, e) + return # RFC 7959 §2.6: a notification carries only the first block # of the representation. Handing the callback a partial CBOR # buffer is what #39 was about, so anything with M=1 (or a # block past the first) goes to the refetch worker instead. - b2 = [v for n, v in ropts if n == BLOCK2] - if b2: - num, more, _ = block_fields(b2[0]) - if more or num: - self._queue_refetch(href) - return - cb = self.on_notification + if blockwise_refetch: + self._queue_refetch( + href, tuple(observe_query), legacy=legacy) + return + cb = ( + self.on_legacy_notification + if legacy and self.on_legacy_notification is not None + else self.on_notification + ) if cb is not None: try: cb(href, payload) @@ -1045,23 +1177,24 @@ def _log_refetch(msg, *args): without also turning on every per-block retransmit line.""" (logger.info if DEBUG_BRIDGE else logger.debug)(msg, *args) - def _queue_refetch(self, href): + def _queue_refetch(self, href, query=(), *, legacy=False): """Queue a blockwise notification for re-reading. Called from the reader thread, so it must not block: _dispatch_coap runs there and _blockwise_get waits on an Event only that same thread can set, which would deadlock the session outright. Latest - wins per href — a burst of notifications on one resource collapses - into a single re-read of its final state.""" + wins per relation — a burst of notifications for one path/query shape + collapses into a single re-read of its final state.""" + key = (href, tuple(query), bool(legacy)) with self._refetch_cond: - if (href not in self._refetch_pending + if (key not in self._refetch_pending and len(self._refetch_pending) >= _MAX_PENDING_REFETCH): self._log_refetch( "refetch %s dropped: queue full (%d pending)", href, len(self._refetch_pending)) return self._refetch_seq += 1 - self._refetch_pending[href] = self._refetch_seq + self._refetch_pending[key] = self._refetch_seq self._refetch_cond.notify() self._start_refetch_worker() @@ -1091,9 +1224,9 @@ def _refetch_loop(self): self._refetch_cond.wait(1.0) if not self._refetch_pending: return - href, seq = next(iter(self._refetch_pending.items())) - del self._refetch_pending[href] - self._refetch_one(href, seq) + key, seq = next(iter(self._refetch_pending.items())) + del self._refetch_pending[key] + self._refetch_one(key, seq) def _refetch_alive(self): """False once the session is closing or the reader has died. A @@ -1103,14 +1236,15 @@ def _refetch_alive(self): return False return self._reader_thread is None or self._reader_running.is_set() - def _refetch_one(self, href, seq): + def _refetch_one(self, key, seq): """Re-read one href from block 0 and deliver it if it is still the freshest thing we know about that resource.""" + href, query, legacy = key self.pace() segs = [s for s in href.split('/') if s] try: code, payload, blocks, tok = self._blockwise_get( - segs, (), _REFETCH_TIMEOUT_S) + segs, query, _REFETCH_TIMEOUT_S) except Exception as e: # Device silent, session gone, ETag never settled, block cap # hit. Whatever the reason, dropping the notification is the @@ -1124,14 +1258,18 @@ def _refetch_one(self, href, seq): with self._refetch_cond: # A newer notification landed while we were reading. That one # has its own refetch queued, so this result is already stale. - if self._refetch_pending.get(href, 0) > seq: + if self._refetch_pending.get(key, 0) > seq: self._log_refetch( "refetch %s tok=%s blocks=%d bytes=%d superseded", href, tok.hex(), blocks, len(payload)) return self._log_refetch("refetch %s tok=%s blocks=%d bytes=%d ok", href, tok.hex(), blocks, len(payload)) - cb = self.on_notification + cb = ( + self.on_legacy_notification + if legacy and self.on_legacy_notification is not None + else self.on_notification + ) if cb is not None: try: cb(href, payload) @@ -1788,21 +1926,34 @@ def refresh_observes(self, paths): # unpaced OBSERVE burst is what wedges an appliance until something # forces a new session (LocalThings#396). The subscribe sweep below # needs nothing here — subscribe() paces its own send. - for tok, href in list(self._observe_tokens.items()): + normalized_paths = tuple(tuple(path) for path in paths) + with self._state_lock: + observations = tuple(self._observe_tokens.items()) + observe_queries = dict(self._observe_queries) + queries_by_href = {} + for tok, href in observations: + queries_by_href.setdefault(href, []).append( + observe_queries.get(tok, ())) + for tok, href in observations: segs = [s for s in href.split('/') if s] try: self.pace() - self._send_observe_dereg(tok, segs) + self._send_observe_dereg( + tok, segs, observe_queries.get(tok, ())) except Exception as e: logger.warning("refresh dereg %s: %s", href, e) - self._observe_tokens.clear() - for path in paths: - try: - self.subscribe(list(path)) - except Exception as e: - logger.warning("refresh subscribe %s: %s", path, e) + with self._state_lock: + self._retire_observe_token_locked(tok) + for path in normalized_paths: + href = '/' + '/'.join(path) + queries = queries_by_href.get(href, [()]) + for query in queries: + try: + self.subscribe(list(path), query=query) + except Exception as e: + logger.warning("refresh subscribe %s: %s", path, e) - def subscribe(self, path_segs): + def subscribe(self, path_segs, *, query=()): """Register an OBSERVE on the given path. The initial 2.05 notification and all subsequent state-change notifications will fire on_notification(href, payload_bytes). @@ -1810,17 +1961,28 @@ def subscribe(self, path_segs): Returns the token used (in case the caller wants to deregister later).""" self._check_live() + path_segs = _validated_text_options( + path_segs, name='path_segs', allow_empty=False) + query = _validated_text_options( + query, name='query', allow_empty=False) self.pace() self._check_live() - tok = self._next_observe_tok() href = '/' + '/'.join(path_segs) # Register the token BEFORE sending — otherwise the device # could respond between send() and the dict insert, and the # reader thread would drop the initial 2.05 as "stale". with self._state_lock: + tok = self._next_available_observe_tok_locked() self._observe_tokens[tok] = href - mid = self._next_mid() + self._observe_queries[tok] = query + try: + mid = self._next_available_mid_locked() + except Exception: + self._retire_observe_token_locked(tok) + raise opts = [(URI_PATH, s.encode()) for s in path_segs] + for value in query: + opts.append((URI_QUERY, value.encode())) opts.append((OBSERVE, OBSERVE_REGISTER)) opts.append((ACCEPT, CF_CBOR)) try: @@ -1828,6 +1990,6 @@ def subscribe(self, path_segs): build_coap(TYPE_CON, METHOD_GET, mid, tok, opts)) except Exception: with self._state_lock: - self._observe_tokens.pop(tok, None) + self._retire_observe_token_locked(tok) raise return tok diff --git a/tests/test_dtls_session_post_retry.py b/tests/test_dtls_session_post_retry.py index 87743ef..46013d0 100644 --- a/tests/test_dtls_session_post_retry.py +++ b/tests/test_dtls_session_post_retry.py @@ -404,7 +404,7 @@ def test_refresh_observes_paces_the_dereg_sweep(): calls = [] sess.pace = lambda: calls.append("pace") sess._send_observe_dereg = lambda *_a: calls.append("dereg") - sess.subscribe = lambda *_a: calls.append("subscribe") + sess.subscribe = lambda *_a, **_k: calls.append("subscribe") sess.refresh_observes([("power", "vs", "0"), ("oven", "vs", "0")]) diff --git a/tests/test_observe_relations.py b/tests/test_observe_relations.py new file mode 100644 index 0000000..012dea6 --- /dev/null +++ b/tests/test_observe_relations.py @@ -0,0 +1,320 @@ +"""Observe relation identity, confirmation, and sequence handling.""" + +from __future__ import annotations + +from unittest.mock import Mock + +import pytest + +from smartthings_local.errors import SessionIdentifierError +from smartthings_local.protocol import dtls_session +from smartthings_local.protocol.coap import ( + BLOCK2, + OBSERVE, + TYPE_ACK, + TYPE_CON, + TYPE_NON, + URI_QUERY, + block_value, + build_coap, + parse_coap, +) +from smartthings_local.protocol.dtls_session import DtlsCoapSession + + +class _NullAuth: + def configure_context(self, _context): + return None + + +def _session(**callbacks): + session = DtlsCoapSession( + "device.example", + 5684, + auth=_NullAuth(), + rate_limit_rps=1_000_000, + **callbacks, + ) + session.conn = object() + return session + + +def _observe_value(value): + if value == 0: + return b"" + return value.to_bytes((value.bit_length() + 7) // 8, "big") + + +def _notify(session, token, mid, payload, *, sequence=None, options=()): + observe_options = ( + () if sequence is None else ((OBSERVE, _observe_value(sequence)),) + ) + session._dispatch_coap( + build_coap( + TYPE_NON, + 0x45, + mid, + token, + (*observe_options, *options), + payload, + ) + ) + + +def test_subscribe_registers_path_and_query_before_immediate_response(): + delivered = [] + session = _session( + on_notification=lambda href, payload: delivered.append((href, payload)) + ) + requests = [] + + def send(datagram): + request = parse_coap(datagram) + requests.append(request) + _mtype, _code, mid, token, options, _payload = request + assert session._observe_tokens[token] == "/mode/vs/0" + assert session._observe_queries[token] == ( + "if=oic.if.a", + "rt=x.test", + ) + session._dispatch_coap( + build_coap( + TYPE_ACK, + 0x45, + mid, + token, + [(OBSERVE, b"")], + b"initial", + ) + ) + + session._send_dgram = send + + token = session.subscribe( + ["mode", "vs", "0"], + query=("if=oic.if.a", "rt=x.test"), + ) + + assert len(token) == 1 + assert delivered == [("/mode/vs/0", b"initial")] + assert [value for number, value in requests[0][4] if number == URI_QUERY] == [ + b"if=oic.if.a", + b"rt=x.test", + ] + + +def test_plain_initial_response_needs_different_mid_to_confirm_legacy(): + delivered = [] + pending = [] + session = _session( + on_notification=lambda href, payload: delivered.append( + ("standard", href, payload) + ), + on_legacy_notification=lambda href, payload: delivered.append( + ("legacy", href, payload) + ), + on_observe_pending=pending.append, + ) + session._send_dgram = Mock() + token = session.subscribe(["doors", "vs", "0"]) + + _notify(session, token, 10, b"initial") + _notify(session, token, 10, b"initial-retransmit") + + assert pending == ["/doors/vs/0"] + assert delivered == [] + assert token not in session._legacy_observe_tokens + + _notify(session, token, 11, b"changed") + _notify(session, token, 11, b"changed-retransmit") + _notify(session, token, 12, b"changed-again") + + assert token in session._legacy_observe_tokens + assert delivered == [ + ("legacy", "/doors/vs/0", b"changed"), + ("legacy", "/doors/vs/0", b"changed-again"), + ] + + +def test_legacy_notification_falls_back_to_main_callback(): + delivered = [] + session = _session( + on_notification=lambda href, payload: delivered.append((href, payload)) + ) + session._send_dgram = Mock() + token = session.subscribe(["doors", "vs", "0"]) + + _notify(session, token, 20, b"initial") + _notify(session, token, 21, b"changed") + + assert delivered == [("/doors/vs/0", b"changed")] + + +def test_rejected_observe_retires_every_relation_index(): + errors = [] + session = _session( + on_observe_error=lambda href, code: errors.append((href, code)) + ) + session._send_dgram = Mock() + token = session.subscribe( + ["mode", "vs", "0"], query=("if=oic.if.a",) + ) + session._observe_plain_response_mids[token] = 1 + session._legacy_observe_tokens.add(token) + session._legacy_observe_mids[token] = 2 + session._observe_sequences[token] = (3, 4.0) + + session._dispatch_coap( + build_coap(TYPE_ACK, 0x80, 10, token, [], b"rejected") + ) + + assert errors == [("/mode/vs/0", 0x80)] + assert token not in session._observe_tokens + assert token not in session._observe_queries + assert token not in session._observe_plain_response_mids + assert token not in session._legacy_observe_tokens + assert token not in session._legacy_observe_mids + assert token not in session._observe_sequences + + +def test_rfc_observe_sequence_drops_duplicate_and_stale_values(): + delivered = [] + session = _session( + on_notification=lambda _href, payload: delivered.append(payload) + ) + session._send_dgram = Mock() + token = session.subscribe(["mode", "vs", "0"]) + + _notify(session, token, 1, b"zero", sequence=0) + _notify(session, token, 2, b"one", sequence=1) + _notify(session, token, 3, b"duplicate", sequence=1) + _notify(session, token, 4, b"stale", sequence=0) + + assert delivered == [b"zero", b"one"] + + +def test_rfc_observe_sequence_accepts_24_bit_wrap(): + delivered = [] + session = _session( + on_notification=lambda _href, payload: delivered.append(payload) + ) + session._send_dgram = Mock() + token = session.subscribe(["mode", "vs", "0"]) + + _notify(session, token, 1, b"before-wrap", sequence=0xFFFFFE) + _notify(session, token, 2, b"after-wrap", sequence=1) + + assert delivered == [b"before-wrap", b"after-wrap"] + + +def test_receipt_time_can_reestablish_sequence_after_128_seconds(monkeypatch): + now = [100.0] + delivered = [] + monkeypatch.setattr(dtls_session.time, "monotonic", lambda: now[0]) + session = _session( + on_notification=lambda _href, payload: delivered.append(payload) + ) + session._send_dgram = Mock() + token = session.subscribe(["mode", "vs", "0"]) + + _notify(session, token, 1, b"newer", sequence=100) + now[0] += 128.001 + _notify(session, token, 2, b"server-restarted", sequence=1) + + assert delivered == [b"newer", b"server-restarted"] + + +@pytest.mark.parametrize( + "options", + ( + ((OBSERVE, b""), (OBSERVE, b"\x01")), + ((OBSERVE, b"\x00\x00\x00\x01"),), + ((OBSERVE, b""), (BLOCK2, b""), (BLOCK2, b"")), + ((OBSERVE, b""), (BLOCK2, b"\x00\x00\x00\x00")), + ), +) +def test_malformed_observe_transport_options_do_not_advance_relation(options): + delivered = [] + session = _session( + on_notification=lambda _href, payload: delivered.append(payload) + ) + session._send_dgram = Mock() + token = session.subscribe(["mode", "vs", "0"]) + + session._dispatch_coap( + build_coap(TYPE_NON, 0x45, 1, token, options, b"invalid") + ) + + assert delivered == [] + assert token not in session._observe_sequences + + +def test_blockwise_refetch_keeps_query_and_callback_kind(): + standard = [] + legacy = [] + session = _session( + on_notification=lambda href, payload: standard.append((href, payload)), + on_legacy_notification=lambda href, payload: legacy.append((href, payload)), + ) + session._blockwise_get = Mock( + return_value=(0x45, b"complete", 2, b"fresh") + ) + + key = ("/mode/vs/0", ("if=oic.if.a",), True) + session._refetch_one(key, 1) + + session._blockwise_get.assert_called_once_with( + ["mode", "vs", "0"], + ("if=oic.if.a",), + dtls_session._REFETCH_TIMEOUT_S, + ) + assert standard == [] + assert legacy == [("/mode/vs/0", b"complete")] + + +def test_refresh_preserves_each_query_separated_relation(): + session = _session() + session.pace = Mock() + requests = [] + session._send_dgram = requests.append + path = ["mode", "vs", "0"] + session.subscribe(path, query=("if=oic.if.a",)) + session.subscribe(path, query=("if=oic.if.s",)) + requests.clear() + + session.refresh_observes([path]) + + parsed = [parse_coap(request) for request in requests] + deregisters = [ + message for message in parsed + if (OBSERVE, b"\x01") in message[4] + ] + registers = [ + message for message in parsed + if (OBSERVE, b"") in message[4] + ] + assert { + tuple(value for number, value in message[4] if number == URI_QUERY) + for message in deregisters + } == {(b"if=oic.if.a",), (b"if=oic.if.s",)} + assert { + tuple(value for number, value in message[4] if number == URI_QUERY) + for message in registers + } == {(b"if=oic.if.a",), (b"if=oic.if.s",)} + + +def test_observe_token_space_fails_closed_without_overwriting_relation(): + session = _session() + session._send_dgram = Mock() + session._observe_tokens = { + bytes([value]): f"/resource/{value}" for value in range(1, 256) + } + session._observe_queries = { + token: () for token in session._observe_tokens + } + + with pytest.raises(SessionIdentifierError): + session.subscribe(["mode", "vs", "0"]) + + assert len(session._observe_tokens) == 255 + session._send_dgram.assert_not_called() diff --git a/tests/test_public_api_contract.py b/tests/test_public_api_contract.py index ffce576..de83f13 100644 --- a/tests/test_public_api_contract.py +++ b/tests/test_public_api_contract.py @@ -63,6 +63,14 @@ def test_dtls_session_constructor_keeps_file_memory_and_local_port_inputs(): auth_parameter = inspect.signature(DtlsCoapSession).parameters["auth"] assert auth_parameter.kind is inspect.Parameter.KEYWORD_ONLY assert auth_parameter.default is None + for callback in ( + "on_legacy_notification", + "on_observe_pending", + "on_observe_error", + ): + parameter = inspect.signature(DtlsCoapSession).parameters[callback] + assert parameter.kind is inspect.Parameter.KEYWORD_ONLY + assert parameter.default is None def test_known_host_multicast_discovery_has_a_bounded_explicit_interface_api(): @@ -212,6 +220,11 @@ def test_dtls_session_keeps_current_consumer_methods(): "timeout", ], ) + subscribe_query = inspect.signature(DtlsCoapSession.subscribe).parameters[ + "query" + ] + assert subscribe_query.kind is inspect.Parameter.KEYWORD_ONLY + assert subscribe_query.default == () get_extra_options = inspect.signature(DtlsCoapSession.get).parameters[ "extra_options" ] diff --git a/tests/test_session_shutdown.py b/tests/test_session_shutdown.py index ce47758..9dc6bc6 100644 --- a/tests/test_session_shutdown.py +++ b/tests/test_session_shutdown.py @@ -253,4 +253,4 @@ def test_normal_close_still_deregisters_before_quiescing(monkeypatch): session.close() - assert calls == [(b"o", ["mode", "vs", "0"])] + assert calls == [(b"o", ["mode", "vs", "0"], ())] From 987d46db7e21a67e927b87a04c8b15cb182187ab Mon Sep 17 00:00:00 2001 From: Jason Morcos Date: Fri, 28 Aug 2026 12:26:59 -0700 Subject: [PATCH 3/3] feat(protocol): add targeted Observe operations --- README.md | 19 + smartthings_local/protocol/dtls_session.py | 242 +++++++++--- tests/test_observe_operations.py | 405 +++++++++++++++++++++ tests/test_observe_relations.py | 4 + tests/test_public_api_contract.py | 10 + 5 files changed, 631 insertions(+), 49 deletions(-) create mode 100644 tests/test_observe_operations.py diff --git a/README.md b/README.md index 3e9f619..0eae7d1 100644 --- a/README.md +++ b/README.md @@ -85,6 +85,25 @@ Message ID. Optional `on_observe_pending`, `on_legacy_notification`, and `on_observe_error` constructor callbacks let consumers keep that compatibility path distinct from confirmed RFC notifications and ordinary polling. +Periodic renewal can target only the relations that need it; unrelated +observations remain active. Existing query variants are preserved unless the +caller supplies an explicit replacement: + +```python +successful, failures = sess.refresh_observes( + (("mode", "vs", "0"),), + queries_by_href={"/mode/vs/0": ("if=oic.if.b",)}, +) +removed = sess.unsubscribe(("mode", "vs", "0")) +``` + +`successful` reports hrefs whose replacement registration datagram was sent; +confirmation still comes from the Observe callbacks. `unsubscribe()` retires +every query-qualified relation for that exact path without disturbing sibling +paths. Refresh, unsubscribe, and orderly close pace every deregistration just +as `subscribe()` paces each registration, avoiding request bursts during +relation maintenance. + POST bodies through 1024 bytes retain the single-request behavior. Larger bodies use token-stable Block1 requests under one monotonic timeout, include Size1 on the first request, honor a server-requested smaller block size, and diff --git a/smartthings_local/protocol/dtls_session.py b/smartthings_local/protocol/dtls_session.py index d035ede..ade3e0f 100644 --- a/smartthings_local/protocol/dtls_session.py +++ b/smartthings_local/protocol/dtls_session.py @@ -32,6 +32,7 @@ import socket import threading import time +from collections.abc import Mapping from dataclasses import dataclass from OpenSSL import SSL @@ -149,6 +150,7 @@ _OBSERVE_SEQUENCE_MODULUS = 1 << 24 _OBSERVE_SEQUENCE_HALF_RANGE = 1 << 23 _OBSERVE_SEQUENCE_RESET_S = 128.0 +_MAX_OBSERVE_RELATIONS = 0xFF class _EtagChanged(Exception): @@ -228,6 +230,49 @@ def _validated_text_options(values, *, name, allow_empty): return tuple(result) +def _validated_observe_paths(paths): + """Return unique bounded Observe paths and their canonical hrefs.""" + if isinstance(paths, (str, bytes, bytearray, memoryview)): + raise TypeError('paths must be an iterable of path iterables') + try: + iterator = iter(paths) + except TypeError: + raise TypeError( + 'paths must be an iterable of path iterables') from None + result = [] + seen_hrefs = set() + for index, path in enumerate(iterator): + if index >= _MAX_OBSERVE_RELATIONS: + raise ValueError('paths must contain at most 255 paths') + normalized = _validated_text_options( + path, name='path', allow_empty=False) + href = '/' + '/'.join(normalized) + if href in seen_hrefs: + continue + seen_hrefs.add(href) + result.append((normalized, href)) + return tuple(result) + + +def _validated_observe_query_overrides(queries_by_href, target_hrefs): + """Validate explicit path-to-query overrides before relation mutation.""" + if queries_by_href is None: + return {} + if not isinstance(queries_by_href, Mapping): + raise TypeError('queries_by_href must be a mapping or None') + if len(queries_by_href) > len(target_hrefs): + raise ValueError('queries_by_href contains a non-target href') + result = {} + for href, query in queries_by_href.items(): + if not isinstance(href, str): + raise TypeError('queries_by_href keys must be href strings') + if href not in target_hrefs: + raise ValueError('queries_by_href contains a non-target href') + result[href] = _validated_text_options( + query, name='query', allow_empty=False) + return result + + def _validated_extra_options(extra_options): """Return bounded, ordered options not owned by the request methods.""" if isinstance(extra_options, (str, bytes, bytearray, memoryview)): @@ -430,6 +475,10 @@ def __init__(self, host, port, cert_path=None, key_path=None, *, # this they can collide and one transfer silently absorbs the # other's blocks. self._state_lock = threading.Lock() + # Serialize public relation mutations without holding _state_lock over + # network sends. The reader must remain free to dispatch an immediate + # registration response while refresh, unsubscribe, or close runs. + self._observe_operation_lock = threading.RLock() # Randomize MID and token counter starting points so reconnects # don't reuse identifiers from previous sessions — Samsung's # RT-OCF appears to remember observer state across DTLS @@ -657,8 +706,6 @@ def join(self): def _send_observe_dereg(self, tok, path_segs, query=()): """Send a single OBSERVE deregister GET (Observe option = 1) on the existing token. Best-effort — caller swallows errors.""" - if self.conn is None: - return mid = self._next_mid() opts = [(URI_PATH, s.encode()) for s in path_segs] for value in query: @@ -710,6 +757,10 @@ def close(self): first so Samsung's RT-OCF cleans up its observer table — without this, the per-cert observer state survives DTLS close and a quick reconnect with the same tokens silently no-ops.""" + with self._observe_operation_lock: + self._close_orderly() + + def _close_orderly(self): # Send dereg for every active observation while the conn is # still healthy. Tiny sleep lets the records reach the wire # before we shut DTLS down. @@ -721,6 +772,7 @@ def close(self): for tok, href in observations: segs = [s for s in href.split('/') if s] try: + self.pace() self._send_observe_dereg( tok, segs, observe_queries.get(tok, ())) except Exception as e: @@ -750,18 +802,19 @@ def close(self): def abort(self): """Immediately stop work and close the established transport.""" self.quiesce_for_close() - with self._lifecycle_lock: - sock = self.sock - self.sock = None - self.conn = None - self.dest = None - self.endpoint = None - if sock is not None: - try: - sock.close() - except Exception: - pass - self._clear_observe_relations() + with self._observe_operation_lock: + with self._lifecycle_lock: + sock = self.sock + self.sock = None + self.conn = None + self.dest = None + self.endpoint = None + if sock is not None: + try: + sock.close() + except Exception: + pass + self._clear_observe_relations() # ---- send / receive plumbing ------------------------------------- @@ -864,6 +917,31 @@ def _clear_observe_relations(self): self._legacy_observe_mids.clear() self._observe_sequences.clear() + def _observe_relation_active(self, href, query, legacy): + """Return whether one confirmed relation still owns this identity.""" + with self._state_lock: + for tok, observed_href in self._observe_tokens.items(): + if observed_href != href or \ + self._observe_queries.get(tok, ()) != query: + continue + if legacy: + if tok in self._legacy_observe_tokens: + return True + elif tok in self._observe_sequences: + return True + return False + + def _discard_refetches_for_hrefs(self, hrefs): + """Discard queued work for retired paths; running work self-checks.""" + hrefs = frozenset(hrefs) + if not hrefs: + return + with self._refetch_cond: + for key in tuple(self._refetch_pending): + if key[0] in hrefs: + self._refetch_pending.pop(key, None) + self._refetch_cond.notify_all() + @staticmethod def _observe_sequence_is_fresh(previous, current, received_at): """Apply RFC 7641's 24-bit serial-number freshness comparison.""" @@ -980,6 +1058,7 @@ def _reader_loop(self): with self._refetch_cond: self._refetch_pending.clear() self._refetch_cond.notify_all() + self._clear_observe_relations() def _dispatch_coap(self, datagram): try: @@ -1240,6 +1319,8 @@ def _refetch_one(self, key, seq): """Re-read one href from block 0 and deliver it if it is still the freshest thing we know about that resource.""" href, query, legacy = key + if not self._observe_relation_active(href, query, legacy): + return self.pace() segs = [s for s in href.split('/') if s] try: @@ -1263,6 +1344,8 @@ def _refetch_one(self, key, seq): "refetch %s tok=%s blocks=%d bytes=%d superseded", href, tok.hex(), blocks, len(payload)) return + if not self._observe_relation_active(href, query, legacy): + return self._log_refetch("refetch %s tok=%s blocks=%d bytes=%d ok", href, tok.hex(), blocks, len(payload)) cb = ( @@ -1907,51 +1990,108 @@ def ping(self): self._send_dgram(build_coap(TYPE_CON, 0, mid, b'', [])) return mid - def refresh_observes(self, paths): - """Drop all current OBSERVE registrations and re-subscribe to - the given paths. Used as a periodic safety net — CoAP OBSERVE - has no built-in TTL but Samsung's RT-OCF can age out its - observer table during cloud blips even while the DTLS session - stays healthy. Without this, internet recovery on a still- - reachable device leaves push permanently dead. - - Best-effort: dereg failures are logged and we still bind fresh - tokens via subscribe. Brief race window where a notify on the - old token gets dropped as 'stale' — acceptable for a 6h-scale - safety net.""" + def refresh_observes(self, paths, *, queries_by_href=None): + """Replace only the requested Observe paths and report send results. + + Existing query-separated relations are preserved unless an exact href + is present in ``queries_by_href``. Deregistration is best-effort; each + old local relation is retired even when its send fails. A successful + href means its replacement registration datagram was sent, not that a + later response has confirmed the relation. + """ + with self._observe_operation_lock: + return self._refresh_observes_locked( + paths, queries_by_href=queries_by_href) + + def _refresh_observes_locked(self, paths, *, queries_by_href=None): self._check_live() - # Paced, unlike the teardown dereg in close(): that one wants out - # quickly and the session is finished either way, while this one - # runs against a session that has to keep working afterwards, and an - # unpaced OBSERVE burst is what wedges an appliance until something - # forces a new session (LocalThings#396). The subscribe sweep below - # needs nothing here — subscribe() paces its own send. - normalized_paths = tuple(tuple(path) for path in paths) + normalized_paths = _validated_observe_paths(paths) + target_hrefs = frozenset(href for _path, href in normalized_paths) + query_overrides = _validated_observe_query_overrides( + queries_by_href, target_hrefs) + + # Snapshot every target before any network work so query preservation + # and the mutation set describe the same relation generation. with self._state_lock: - observations = tuple(self._observe_tokens.items()) - observe_queries = dict(self._observe_queries) - queries_by_href = {} - for tok, href in observations: - queries_by_href.setdefault(href, []).append( - observe_queries.get(tok, ())) - for tok, href in observations: - segs = [s for s in href.split('/') if s] + observations = tuple( + (tok, href, self._observe_queries.get(tok, ())) + for tok, href in self._observe_tokens.items() + if href in target_hrefs + ) + + preserved_queries = {} + for _tok, href, query in observations: + queries = preserved_queries.setdefault(href, []) + if query not in queries: + queries.append(query) + + for tok, href, query in observations: try: self.pace() self._send_observe_dereg( - tok, segs, observe_queries.get(tok, ())) + tok, [segment for segment in href.split('/') if segment], + query) except Exception as e: logger.warning("refresh dereg %s: %s", href, e) - with self._state_lock: - self._retire_observe_token_locked(tok) - for path in normalized_paths: - href = '/' + '/'.join(path) - queries = queries_by_href.get(href, [()]) + finally: + with self._state_lock: + self._retire_observe_token_locked(tok) + self._discard_refetches_for_hrefs(target_hrefs) + + successful = [] + successful_hrefs = set() + failures = 0 + for path, href in normalized_paths: + queries = ( + (query_overrides[href],) + if href in query_overrides + else tuple(preserved_queries.get(href, ())) or ((),) + ) for query in queries: try: - self.subscribe(list(path), query=query) + # subscribe() owns pacing for its registration send. + self.subscribe(path, query=query) except Exception as e: - logger.warning("refresh subscribe %s: %s", path, e) + failures += 1 + logger.warning("refresh subscribe %s: %s", href, e) + else: + if href not in successful_hrefs: + successful.append(href) + successful_hrefs.add(href) + return tuple(successful), failures + + def unsubscribe(self, path_segs): + """Deregister and retire every active relation for one exact path.""" + with self._observe_operation_lock: + return self._unsubscribe_locked(path_segs) + + def _unsubscribe_locked(self, path_segs): + self._check_live() + path_segs = _validated_text_options( + path_segs, name='path_segs', allow_empty=False) + href = '/' + '/'.join(path_segs) + with self._state_lock: + observations = tuple( + (tok, self._observe_queries.get(tok, ())) + for tok, observed_href in self._observe_tokens.items() + if observed_href == href + ) + + first_error = None + for tok, query in observations: + try: + self.pace() + self._send_observe_dereg(tok, path_segs, query) + except Exception as e: + if first_error is None: + first_error = e + finally: + with self._state_lock: + self._retire_observe_token_locked(tok) + self._discard_refetches_for_hrefs((href,)) + if first_error is not None: + raise first_error + return len(observations) def subscribe(self, path_segs, *, query=()): """Register an OBSERVE on the given path. The initial 2.05 @@ -1960,6 +2100,10 @@ def subscribe(self, path_segs, *, query=()): Returns the token used (in case the caller wants to deregister later).""" + with self._observe_operation_lock: + return self._subscribe_locked(path_segs, query=query) + + def _subscribe_locked(self, path_segs, *, query=()): self._check_live() path_segs = _validated_text_options( path_segs, name='path_segs', allow_empty=False) diff --git a/tests/test_observe_operations.py b/tests/test_observe_operations.py new file mode 100644 index 0000000..f856f9d --- /dev/null +++ b/tests/test_observe_operations.py @@ -0,0 +1,405 @@ +"""Targeted Observe refresh, unsubscribe, and cleanup contracts.""" + +from __future__ import annotations + +import threading +from unittest.mock import Mock + +import pytest + +from smartthings_local.errors import EndpointError +from smartthings_local.protocol.coap import ( + OBSERVE, + URI_PATH, + URI_QUERY, + parse_coap, +) +from smartthings_local.protocol.dtls_session import DtlsCoapSession + + +class _NullAuth: + def configure_context(self, _context): + return None + + +class _Socket: + def __init__(self): + self.closed = False + + def close(self): + self.closed = True + + +def _session(): + session = DtlsCoapSession( + "device.example", + 5684, + auth=_NullAuth(), + rate_limit_rps=1_000_000, + ) + session.conn = object() + return session + + +def _add_relation(session, path, query=()): + session._send_dgram = Mock() + return session.subscribe(path, query=query) + + +def _options(datagram, number): + return tuple(value for option, value in parse_coap(datagram)[4] + if option == number) + + +def test_refresh_targets_only_requested_path_and_preserves_sibling(): + session = _session() + mode = ["mode", "vs", "0"] + door = ["door", "vs", "0"] + mode_token = _add_relation(session, mode, ("if=oic.if.a",)) + door_token = _add_relation(session, door, ("if=oic.if.s",)) + requests = [] + session._send_dgram = requests.append + + successful, failures = session.refresh_observes((mode,)) + + assert successful == ("/mode/vs/0",) + assert failures == 0 + assert mode_token not in session._observe_tokens + assert session._observe_tokens[door_token] == "/door/vs/0" + assert session._observe_queries[door_token] == ("if=oic.if.s",) + assert len(requests) == 2 + assert _options(requests[0], URI_PATH) == (b"mode", b"vs", b"0") + assert _options(requests[0], URI_QUERY) == (b"if=oic.if.a",) + assert _options(requests[0], OBSERVE) == (b"\x01",) + assert _options(requests[1], URI_QUERY) == (b"if=oic.if.a",) + assert _options(requests[1], OBSERVE) == (b"",) + + +def test_refresh_override_replaces_all_old_query_variants_once(): + session = _session() + path = ["mode", "vs", "0"] + _add_relation(session, path, ("if=oic.if.a",)) + _add_relation(session, path, ("if=oic.if.s",)) + requests = [] + session._send_dgram = requests.append + + result = session.refresh_observes( + (path,), + queries_by_href={"/mode/vs/0": ("if=oic.if.b",)}, + ) + + assert result == (("/mode/vs/0",), 0) + assert len(requests) == 3 + assert { + _options(datagram, URI_QUERY) for datagram in requests[:2] + } == {(b"if=oic.if.a",), (b"if=oic.if.s",)} + assert _options(requests[2], URI_QUERY) == (b"if=oic.if.b",) + + +def test_refresh_reports_registration_sends_without_claiming_confirmation(): + session = _session() + session.subscribe = Mock( + side_effect=(b"\x41", EndpointError(), b"\x43") + ) + + successful, failures = session.refresh_observes( + ( + ("door", "cooler", "0"), + ("door", "freezer", "0"), + ("mode", "vs", "0"), + ) + ) + + assert successful == ("/door/cooler/0", "/mode/vs/0") + assert failures == 1 + assert session.subscribe.call_count == 3 + + +@pytest.mark.parametrize( + ("paths", "queries_by_href", "error"), + ( + ((("mode", "vs", "0"), ("door", 7)), None, TypeError), + ((("mode", "vs", "0"),), {"/other": ()}, ValueError), + ((("mode", "vs", "0"),), {"/mode/vs/0": (b"bad",)}, TypeError), + ("/mode/vs/0", None, TypeError), + ), +) +def test_refresh_validation_is_atomic(paths, queries_by_href, error): + session = _session() + token = _add_relation(session, ["door", "vs", "0"]) + original_tokens = dict(session._observe_tokens) + session.pace = Mock() + session._send_observe_dereg = Mock() + + with pytest.raises(error): + session.refresh_observes( + paths, + queries_by_href=queries_by_href, + ) + + assert session._observe_tokens == original_tokens + assert token in session._observe_tokens + session.pace.assert_not_called() + session._send_observe_dereg.assert_not_called() + + +def test_refresh_dereg_failure_retires_every_old_relation_and_continues(): + session = _session() + path = ["mode", "vs", "0"] + first = _add_relation(session, path, ("if=oic.if.a",)) + second = _add_relation(session, path, ("if=oic.if.s",)) + for token in (first, second): + session._observe_plain_response_mids[token] = 10 + session._legacy_observe_tokens.add(token) + session._legacy_observe_mids[token] = 11 + session._observe_sequences[token] = (12, 13.0) + session._send_observe_dereg = Mock( + side_effect=(EndpointError(), None) + ) + session.subscribe = Mock(return_value=b"\x50") + + assert session.refresh_observes((path,)) == (("/mode/vs/0",), 0) + + assert session._send_observe_dereg.call_count == 2 + for token in (first, second): + assert token not in session._observe_tokens + assert token not in session._observe_queries + assert token not in session._observe_plain_response_mids + assert token not in session._legacy_observe_tokens + assert token not in session._legacy_observe_mids + assert token not in session._observe_sequences + + +def test_refresh_paces_each_deregister_and_registration_exactly_once(): + session = _session() + path = ["mode", "vs", "0"] + _add_relation(session, path) + session.pace = Mock() + session._send_dgram = Mock() + + session.refresh_observes((path,)) + + assert session.pace.call_count == 2 + assert session._send_dgram.call_count == 2 + + +def test_refresh_serializes_a_concurrent_subscription(): + session = _session() + path = ["mode", "vs", "0"] + _add_relation(session, path) + deregister_started = threading.Event() + release_deregister = threading.Event() + subscribe_started = threading.Event() + subscribe_finished = threading.Event() + outcomes = {} + + def deregister(*_args): + deregister_started.set() + assert release_deregister.wait(1.0) + + def refresh(): + outcomes["refresh"] = session.refresh_observes((path,)) + + def subscribe(): + subscribe_started.set() + outcomes["subscribe"] = session.subscribe(("door", "vs", "0")) + subscribe_finished.set() + + session._send_observe_dereg = deregister + session._send_dgram = Mock() + refresh_thread = threading.Thread(target=refresh) + subscribe_thread = threading.Thread(target=subscribe) + refresh_thread.start() + assert deregister_started.wait(1.0) + subscribe_thread.start() + assert subscribe_started.wait(1.0) + assert not subscribe_finished.wait(0.05) + + release_deregister.set() + refresh_thread.join(1.0) + subscribe_thread.join(1.0) + + assert not refresh_thread.is_alive() + assert not subscribe_thread.is_alive() + assert outcomes["refresh"] == (("/mode/vs/0",), 0) + assert isinstance(outcomes["subscribe"], bytes) + + +def test_refresh_deduplicates_requested_paths(): + session = _session() + session.subscribe = Mock(return_value=b"\x41") + path = ("mode", "vs", "0") + + result = session.refresh_observes((path, path)) + + assert result == (("/mode/vs/0",), 0) + session.subscribe.assert_called_once_with(path, query=()) + + +def test_unsubscribe_retires_all_exact_path_queries_and_preserves_sibling(): + session = _session() + target = ["mode", "vs", "0"] + first = _add_relation(session, target, ("if=oic.if.a",)) + second = _add_relation(session, target, ("if=oic.if.s",)) + sibling = _add_relation(session, ["door", "vs", "0"]) + for token in (first, second): + session._observe_plain_response_mids[token] = 10 + session._legacy_observe_tokens.add(token) + session._legacy_observe_mids[token] = 11 + session._observe_sequences[token] = (12, 13.0) + session._refetch_pending = { + ("/mode/vs/0", ("if=oic.if.a",), True): 1, + ("/door/vs/0", (), False): 2, + } + session.pace = Mock() + session._send_observe_dereg = Mock() + + assert session.unsubscribe(target) == 2 + + assert session.pace.call_count == 2 + assert session._send_observe_dereg.call_args_list == [ + ((first, tuple(target), ("if=oic.if.a",)), {}), + ((second, tuple(target), ("if=oic.if.s",)), {}), + ] + for token in (first, second): + assert token not in session._observe_tokens + assert token not in session._observe_queries + assert token not in session._observe_plain_response_mids + assert token not in session._legacy_observe_tokens + assert token not in session._legacy_observe_mids + assert token not in session._observe_sequences + assert session._observe_tokens[sibling] == "/door/vs/0" + assert tuple(session._refetch_pending) == ( + ("/door/vs/0", (), False), + ) + + +def test_unsubscribe_retires_every_match_before_reraising_first_error(): + session = _session() + target = ["mode", "vs", "0"] + first = _add_relation(session, target, ("if=oic.if.a",)) + second = _add_relation(session, target, ("if=oic.if.s",)) + sibling = _add_relation(session, ["door", "vs", "0"]) + first_error = EndpointError() + session._send_observe_dereg = Mock( + side_effect=(first_error, EndpointError()) + ) + + with pytest.raises(EndpointError) as raised: + session.unsubscribe(target) + + assert raised.value is first_error + assert session._send_observe_dereg.call_count == 2 + assert first not in session._observe_tokens + assert second not in session._observe_tokens + assert sibling in session._observe_tokens + + +def test_unsubscribe_validation_has_no_side_effects(): + session = _session() + token = _add_relation(session, ["mode", "vs", "0"]) + session.pace = Mock() + session._send_observe_dereg = Mock() + + with pytest.raises(TypeError): + session.unsubscribe("/mode/vs/0") + + assert token in session._observe_tokens + session.pace.assert_not_called() + session._send_observe_dereg.assert_not_called() + + +def test_normal_close_paces_every_exact_deregister_and_clears_state( + monkeypatch, +): + session = _session() + session.sock = _Socket() + first = _add_relation( + session, ["mode", "vs", "0"], ("if=oic.if.a",) + ) + second = _add_relation(session, ["door", "vs", "0"]) + session._legacy_observe_tokens.add(first) + session._observe_sequences[second] = (1, 2.0) + session.pace = Mock() + session._send_observe_dereg = Mock() + monkeypatch.setattr( + DtlsCoapSession, + "_send_close_notify", + Mock(), + ) + monkeypatch.setattr( + "smartthings_local.protocol.dtls_session.time.sleep", + Mock(), + ) + + session.close() + + assert session.pace.call_count == 2 + assert session._send_observe_dereg.call_args_list == [ + ((first, ["mode", "vs", "0"], ("if=oic.if.a",)), {}), + ((second, ["door", "vs", "0"], ()), {}), + ] + assert session._observe_tokens == {} + assert session._observe_queries == {} + assert session._legacy_observe_tokens == set() + assert session._observe_sequences == {} + + +def test_quiesced_close_skips_deregister_pacing(): + session = _session() + session.sock = _Socket() + _add_relation(session, ["mode", "vs", "0"]) + session.pace = Mock() + session._send_observe_dereg = Mock() + + session.quiesce_for_close() + session.close() + + session.pace.assert_not_called() + session._send_observe_dereg.assert_not_called() + + +def test_reader_exit_clears_all_relation_state(): + session = _session() + token = _add_relation(session, ["mode", "vs", "0"]) + session._observe_plain_response_mids[token] = 10 + session._legacy_observe_tokens.add(token) + session._legacy_observe_mids[token] = 11 + session._observe_sequences[token] = (12, 13.0) + session.sock = Mock() + session._stop.set() + + session._reader_loop() + + assert session._observe_tokens == {} + assert session._observe_queries == {} + assert session._observe_plain_response_mids == {} + assert session._legacy_observe_tokens == set() + assert session._legacy_observe_mids == {} + assert session._observe_sequences == {} + + +def test_running_refetch_drops_result_after_relation_is_retired(): + delivered = [] + session = _session() + session.on_notification = lambda href, payload: delivered.append( + (href, payload) + ) + token = _add_relation( + session, ["mode", "vs", "0"], ("if=oic.if.a",) + ) + session._observe_sequences[token] = (1, 2.0) + + def retire_during_read(*_args, **_kwargs): + with session._state_lock: + session._retire_observe_token_locked(token) + return 0x45, b"complete", 2, b"fresh" + + session._blockwise_get = retire_during_read + + session._refetch_one( + ("/mode/vs/0", ("if=oic.if.a",), False), + 1, + ) + + assert delivered == [] diff --git a/tests/test_observe_relations.py b/tests/test_observe_relations.py index 012dea6..7c4ffa7 100644 --- a/tests/test_observe_relations.py +++ b/tests/test_observe_relations.py @@ -261,6 +261,10 @@ def test_blockwise_refetch_keeps_query_and_callback_kind(): ) key = ("/mode/vs/0", ("if=oic.if.a",), True) + token = b"\x41" + session._observe_tokens[token] = key[0] + session._observe_queries[token] = key[1] + session._legacy_observe_tokens.add(token) session._refetch_one(key, 1) session._blockwise_get.assert_called_once_with( diff --git a/tests/test_public_api_contract.py b/tests/test_public_api_contract.py index de83f13..8358466 100644 --- a/tests/test_public_api_contract.py +++ b/tests/test_public_api_contract.py @@ -192,6 +192,7 @@ def test_dtls_session_keeps_current_consumer_methods(): "refresh_observes", "start_reader", "subscribe", + "unsubscribe", } assert expected <= set(dir(DtlsCoapSession)) _assert_compatible_signature(DtlsCoapSession.connect, ["self"]) @@ -225,6 +226,15 @@ def test_dtls_session_keeps_current_consumer_methods(): ] assert subscribe_query.kind is inspect.Parameter.KEYWORD_ONLY assert subscribe_query.default == () + _assert_compatible_signature( + DtlsCoapSession.unsubscribe, + ["self", "path_segs"], + ) + refresh_queries = inspect.signature( + DtlsCoapSession.refresh_observes + ).parameters["queries_by_href"] + assert refresh_queries.kind is inspect.Parameter.KEYWORD_ONLY + assert refresh_queries.default is None get_extra_options = inspect.signature(DtlsCoapSession.get).parameters[ "extra_options" ]