30 Commits
Author SHA1 Message Date
Quite Yellow 4999e8beac Merge pull request #47 from Moballo-LLC/codex/ocf-multicast-responder
feat(protocol): discover known-host OCF responder ports
2026-08-21 12:17:13 +01:00
Quite Yellow b3045f5ddb Merge pull request #51 from Moballo-LLC/codex/coap-initial-pacing
fix(protocol): pace CoAP request sends
2026-08-21 08:38:01 +01:00
Jason Morcos 0d9d13b8fb fix(protocol): pace CoAP request sends 2026-08-18 14:25:10 -07:00
Quite Yellow 0722c552c9 Merge pull request #52 from QuiteYellow/fix/oven-idle-setpoint
fix(mqtt): withhold the oven setpoint when no cycle is set
2026-08-18 20:27:44 +01:00
Jack Nagy a45ee4c004 fix(mqtt): withhold the oven setpoint when no cycle is set
With no cycle set the oven reports x.com.samsung.da.desired = 0, and
flatten() published that straight through as target_temp_c. Home
Assistant rejects it against the Number entity's declared 30-270 range
on every publish, which produced 66,899 log errors over three weeks:

  Invalid value for number.samsung_oven_setpoint: 0 (range 30.0 - 270.0)

0 is not a 0 degree target, it is the absence of a setpoint, so treat
anything outside the settable band as absent. null lands as unknown on
both the Number and the Setpoint sensor, the way completion_minutes
already reads when the oven is idle. _setpoint applied these bounds on
the write side already; only the read path was missing them.

Adds the first tests for the sample descriptors. One of them pins a
non-obvious asymmetry: the write path snaps to the 5 degree step grid
before bounds-checking, so 29 commits as 30 and 271 as 270, and only 0
is refused outright. The invariant that has to hold is the weaker one,
that every value the write path commits is one flatten() will publish
back, or a write appears to succeed and then reads as unknown.
2026-08-18 20:25:19 +01:00
Jason Morcos 512df7ff36 feat(protocol): discover known-host OCF responder ports 2026-08-18 12:22:11 -07:00
Quite Yellow 6b9a508fd9 Merge pull request #44 from Moballo-LLC/codex/issue-9-session-workers
fix(mqtt): retire session workers on reconnect
2026-08-18 20:19:04 +01:00
Quite Yellow ab0ec4961c Merge pull request #46 from Moballo-LLC/codex/owner-psk-clarifications
docs(protocol): clarify OwnerPSK vector scope
2026-08-18 20:18:58 +01:00
Jason Morcos 6db846563a docs(protocol): clarify OwnerPSK vector scope 2026-08-18 11:10:44 -07:00
Jason Morcos 8c374e17b4 fix(mqtt): retire session workers on reconnect 2026-08-18 11:10:38 -07:00
Quite Yellow cd86424ca0 Merge pull request #45 from Moballo-LLC/codex/owner-psk-derivation
feat(protocol): add pure OwnerPSK derivation
2026-08-17 19:49:06 +01:00
Quite Yellow 6da7e95731 Merge pull request #43 from Moballo-LLC/codex/py-08b1-completion-cancel
fix(protocol): keep completed sessions on late cancel
2026-08-17 19:48:53 +01:00
Quite Yellow a3bec9470d Merge pull request #42 from Moballo-LLC/codex/py-08a1-completed-handshake
fix(protocol): retain completed DTLS handshakes
2026-08-17 19:48:44 +01:00
Jason Morcos 79493fbd48 feat(protocol): add pure OwnerPSK derivation 2026-08-15 12:24:11 -07:00
Jason Morcos 627fcb19da fix(protocol): keep completed sessions on late cancel 2026-08-15 11:44:24 -07:00
Jason Morcos 31be87061a fix(protocol): retain completed DTLS handshakes 2026-08-15 11:42:28 -07:00
Quite Yellow bc4465b274 Merge pull request #40 from QuiteYellow/fix/observe-block2-truncation
fix(protocol): reassemble blockwise OBSERVE notifications
2026-08-15 14:08:04 +01:00
Jack Nagy e63acb759f feat(protocol): surface OBSERVE refetch outcomes under DEBUG_BRIDGE
The refetch path logged only at debug, and the bridge configures logging
at INFO, so a successful re-read and a total failure produced identical
output: nothing. That makes the hardware validation for #39 impossible
to read.

One line per refetch, promoted to INFO when DEBUG_BRIDGE=1 and left at
debug otherwise, naming the href, the one-shot token, the block count,
and the reassembled size. The token is the part that matters: it is what
shows the re-read used a fresh 4-byte token rather than the observation's
1-byte one, which is the assumption the whole design rests on.

Gating on DEBUG_BRIDGE rather than raising the logger keeps the per-block
retransmit lines out of the way, and matches how the module already gates
its frame dump.
2026-08-15 13:28:00 +01:00
Jack Nagy 231e88a8c6 fix(protocol): reassemble blockwise OBSERVE notifications
A notification carries only the first block of a large representation
(RFC 7959 §2.6). _dispatch_coap handed that block straight to
on_notification, so consumers decoded a truncated CBOR buffer. Reported
twice on /mode/vs/0: #37 and mbillow/localthings#361.

The recovery is a re-read from block 0 on a fresh 4-byte one-shot token,
not a §2.6 continuation. §3.4 rules out reusing the observation's token,
and this server drops a transfer that opens at NUM>0 under a token it
has not seen, so a continuation is the one shape that cannot work here.
A truncated notification is now withheld and queued to a worker thread
that re-reads the resource and delivers the reassembled representation.
When the re-read fails the notification is dropped at debug level and
the poll tiers carry freshness, which is what they already did.

The re-read has to run off the reader thread: _dispatch_coap runs there
and the transfer waits on an event only that same thread can set. The
worker is serialized and paces between transfers, so a notification
storm stays under the firmware request ceiling.

Also in the Block2 loop, now extracted and shared by both paths:

- compare the response's Block2 NUM against the one requested, so a
  retransmitted block is no longer concatenated as if it were the next
- compare ETags across blocks (§2.4) and restart once when the
  representation changes mid-transfer
- recompute the next block number from the accumulated byte offset when
  the server negotiates the block size down
- re-check reader liveness while waiting on a block, so a mid-transfer
  reader death fails fast instead of burning the whole timeout
- guard the token counters and _pending with a lock, now that the
  session issues concurrent reads of its own

Closes #39
2026-08-15 12:39:48 +01:00
Jack Nagy dec84c9ad8 docs(readme): fill gaps left by the auth/handshake stack merges
- add dtls_handshake.py (from #34) to the repo-layout tree
- point the issue #16 / #20 notes at SamsungServerProfile / ServerCertificateAuth instead of calling that path unsupported
- list the new certificate-profile, connect-deadline, and session-interruption test modules
2026-08-15 11:17:42 +01:00
Quite Yellow e9aee1c235 Merge pull request #33 from Moballo-LLC/codex/py-07-certificate-profiles
Add bound Samsung server certificate profile
2026-08-15 10:20:14 +01:00
Quite Yellow 9ef5598813 Merge pull request #35 from Moballo-LLC/codex/py-08b-session-interruption
Add cancellable DTLS connection attempts
2026-08-15 10:20:05 +01:00
Quite Yellow 917b0e47c5 Merge pull request #34 from Moballo-LLC/codex/py-08a-bounded-connect
Bound DTLS handshakes with a monotonic deadline
2026-08-15 10:19:52 +01:00
Jason Morcos 83a5973434 feat(protocol): add bound server certificate profile 2026-08-14 14:55:10 -07:00
Jason Morcos 3f0e437880 feat(protocol): add cancellable session interruption 2026-08-14 14:37:30 -07:00
Jason Morcos a44930f9df feat(protocol): bound DTLS handshake deadline 2026-08-14 14:26:33 -07:00
Quite Yellow b0d51abcc8 Merge pull request #38 from QuiteYellow/fix/reader-death-visible
fix(dtls): make reader-thread death visible and fail fast
2026-08-14 17:38:41 +01:00
Jack Nagy 7a74a955f3 fix(dtls): make reader-thread death visible and fail fast
The reader loop exited silently on any socket error, leaving conn/sock
set so the session still looked open. Every later get()/post()/ping()
then waited out its full request timeout on a session nobody was
reading, raising SessionTimeoutError on repeat, forever.

This started biting in v0.1.3 (d677c72), which moved to connected UDP
sockets: a connected socket surfaces ICMP errors on recv, so one
ECONNREFUSED from a rebooting appliance now killed the reader.

- Advisory ICMP errnos (ECONNREFUSED/EHOSTUNREACH/...) no longer kill
  the reader; the next datagram usually works.
- Real reader exits log at WARNING; close()-driven exits stay quiet.
- A _reader_running Event lets get/post/ping/subscribe/refresh_observes
  fail fast via _check_live() with SessionClosedError instead of waiting
  out a timeout. Callers that never start a reader are unaffected.

Refs QuiteYellow/SmartThings-Local#37
2026-08-14 17:33:13 +01:00
Quite Yellow d4aebebca4 Merge pull request #32 from Moballo-LLC/codex/py-06-psk-auth
feat(protocol): add PSK authentication provider
2026-08-09 14:21:54 +01:00
Jason Morcos 2a6fc627f2 feat(protocol): add PSK authentication provider 2026-08-08 16:06:53 -07:00
26 changed files with 5719 additions and 215 deletions
+173 -3
View File
@@ -48,6 +48,35 @@ sess.subscribe(["operational", "state", "vs", "0"], # OBSERVE
sess.close()
```
`connect()` uses a 12-second monotonic DTLS handshake deadline by default. A
caller that needs a shorter bounded attempt can pass a positive finite value
without changing later reader timeouts. OpenSSL's DTLS timer schedules flight
retransmissions within that same deadline:
```python
sess.connect(timeout=4.0)
```
The deadline stops further setup, retries, and network waits. If OpenSSL
reports that the handshake completed at the deadline boundary, the completed
session is retained rather than torn down as a timeout.
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:
@@ -56,11 +85,122 @@ auth = CertificateAuth.from_memory(cert_pem, key_pem)
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
```
Some newer OCF-PKI devices require an exact Samsung DTLS offer and present a
hardware certificate whose subject contains a certificate UUID. That UUID can
be distinct from the runtime OCF device UUID reported by `/oic/d`, so callers
must obtain and verify the certificate identity independently. When the caller
already has an authorized client certificate and a previously verified
hardware-certificate UUID, opt in to both requirements explicitly:
```python
from smartthings_local.protocol.auth import (
CertificateAuth,
SamsungServerProfile,
)
server_profile = SamsungServerProfile.bound_device(
expected_certificate_uuid,
additional_ca_pem=additional_samsung_ca_pem,
)
auth = CertificateAuth.from_memory(
cert_pem,
key_pem,
server_profile=server_profile,
)
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
```
The default profile is restricted to Samsung home-appliance leaves with
`OU=OCF HA Device`. The profile limits the ClientHello to P-256,
`ECDHE-ECDSA-AES128-GCM-SHA256`, and the observed SHA-256/SHA-1 RSA/ECDSA
signature set, disables session tickets, preserves certificate-chain
verification, and requires the exact subject role
`C=KR, O=Samsung Electronics, OU=OCF HA Device` with a common name ending in
the expected certificate UUID. `additional_ca_pem` is optional and accepts
only a bounded PEM CA-certificate chain; it is applied only to this profiled
context. Without a profile, `CertificateAuth` retains its existing verification
behavior.
Samsung VD-family devices can present the same wire profile with the distinct
`OU=OCF VD Device` role. Select that role explicitly; profiles never fall back
between device classes:
```python
from smartthings_local.protocol.auth import (
SamsungServerProfile,
SamsungServerRole,
ServerCertificateAuth,
)
server_profile = SamsungServerProfile.bound_device(
expected_certificate_uuid,
role=SamsungServerRole.VD_DEVICE,
)
auth = ServerCertificateAuth(server_profile=server_profile)
sess = DtlsCoapSession("192.0.2.100", 5684, auth=auth)
```
`ServerCertificateAuth` is for a server-authenticated channel that does not
send a client certificate, such as the initial DTLS carrier used by
manufacturer-certificate OTM. It still verifies the CA chain, exact selected
subject role, and pinned certificate UUID. It does not learn an identity from
the first endpoint it reaches and cannot be combined with client credentials.
This API deliberately does not discover, mint, authorize, provision, rotate,
or persist credentials, and it performs no ownership transfer or OCF security
resource writes. In particular, the server-only provider can authenticate the
initial manufacturer-certificate channel, but it does not implement the OTM
that follows. The already-owned new-PKI case in
[issue #16](https://github.com/QuiteYellow/SmartThings-Local/issues/16) still
requires an authorized client identity before ordinary protected resources
can be used.
For compatibility, the existing `cert_path` / `key_path` and `cert_pem` /
`key_pem` session arguments remain supported without a deprecation warning.
They are routed through `CertificateAuth` internally. Do not combine `auth`
with those legacy arguments.
An existing OCF PSK credential can be supplied through `PskAuth`:
```python
from smartthings_local.protocol.auth import PskAuth
auth = PskAuth(identity=psk_identity, key=psk_key)
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
```
The identity must be the raw 16-byte OCF UUID and cannot contain a NUL byte;
the key must be exactly 16 or 32 bytes. `PskAuth` selects only
`ECDHE-PSK-AES128-CBC-SHA256` and does not acquire, derive, provision, rotate,
or persist credentials. Ownership transfer and credential discovery are
outside this package.
Code that has already completed an authenticated manufacturer-certificate
session can derive IoTivity's 128-bit OwnerPSK from the resulting TLS state:
```python
from smartthings_local.protocol.owner_psk import derive_mfg_certificate_owner_psk
owner_psk = derive_mfg_certificate_owner_psk(
master_secret=master_secret,
client_random=client_random,
server_random=server_random,
owner_uuid=owner_uuid,
device_uuid=device_uuid,
cipher_name=cipher_name,
oxm_label=selected_oxm_label,
)
```
The caller must supply the exact authenticated TLS values, non-nil raw OCF
UUIDs, negotiated cipher name, and label for the selected OXM. Use
`STANDARD_MFG_CERTIFICATE_OXM_LABEL` for `oic.sec.doxm.mfgcert` and
`CONFIRMED_MFG_CERTIFICATE_OXM_LABEL` for
`x.org.iotivity.conmfgcert`; do not infer the label from the appliance model.
The helper performs deterministic key derivation only: it does not access a
session, discover credentials, choose an ownership method, write security
resources, run OTM, or persist the result.
### Classified errors
Runtime transport failures use the public types in
@@ -122,6 +262,34 @@ The resolver retains candidate order and the socket setup tries the next
candidate after a family, bind, or connect failure. Resolution and socket
setup failures raise the redacted `EndpointError` documented above.
### Dynamic plaintext OCF response ports
Some OCF devices listen for multicast discovery on UDP 5683 but send their
response from a different port that changes after a power cycle. A caller that
already knows the device's IPv4 address can discover those plaintext response
port candidates on one explicit LAN interface:
```python
from smartthings_local.protocol.ocf_multicast import (
discover_ocf_responder_ports,
)
result = discover_ocf_responder_ports(
"192.0.2.20",
interface_address="192.0.2.10",
)
for discovery_port in result.ports:
pass # use for a bounded, source-bound /oic/res lookup
```
The call sends unfiltered current OCF and legacy IoTivity directory requests,
plus a legacy DOXM-filtered fallback, under one deadline. It accepts only
token-correlated replies from the expected host, closes its multicast socket
before returning, and omits addresses and ports from its result representation.
Returned ports are unauthenticated candidates, not DTLS endpoints; directory
parsing, DTLS liveness, and authenticated device identity remain separate
checks.
For a full worked integration, the higher-level `smartthings_local.ocf` layer (`StateCache`, `PollScheduler`, `KeepaliveTask`, `ObserveRefreshTask`) coordinates tiered polling and OBSERVE on top of a session. The MQTT bridge demo below wires all of it together.
### What the demo bridge gives you
@@ -140,7 +308,7 @@ For a full worked integration, the higher-level `smartthings_local.ocf` layer (`
Each appliance runs an independent bridge built around three coordinated pieces over one persistent DTLS session: a `StateCache` (single source of truth for all reps), a `PollScheduler` (tiered adaptive polling: hot/warm/cold plus a periodic `/device/0` sweep), and a `KeepaliveTask` (CoAP empty-CON ping for DTLS-layer liveness, with consecutive-failure detection for MQTT availability). Tier cadences are descriptor-declared and were calibrated against the empirically-measured per-firmware ceilings: dryer ~14 req/s, oven ~8 req/s. OBSERVE registrations (RFC 7641) are kept as an opportunistic freshness accelerator: when the appliance has internet and emits notifications, the cache absorbs them and the next-poll timer is reset for that resource; when it's air-gapped, polling alone carries the UX with no other code change. Token-stable Block2 (RFC 7959) handles multi-block reads. Writes are optimistically merged into the cache the moment the device 2.04-confirms, with the scheduler deferring that resource's next poll past the fetchback-revert window. Reconnect with exponential backoff on session errors, gated by a stateless DTLS ClientHello pre-flight (`smartthings_local/protocol/dtls_probe.py`) so a silent/rebooting device or wrong port drops into backoff in ~1 RTT instead of eating the full handshake timeout; when `OCF_PORT` is unset the same probe auto-discovers the live port across the OCF band.
On the currently supported firmware families, authentication uses a client cert keyed to the UUID published in Samsung's own wildcard cloud TLS cert. Their factory ACL grants that UUID `perm=31` (full CRUDN) on `href=*`. That certificate path is not universal: the WD53 profile in issue #16 and the washer in issue #20 reject it and need separate authentication work.
On the currently supported firmware families, authentication uses a client cert keyed to the UUID published in Samsung's own wildcard cloud TLS cert. Their factory ACL grants that UUID `perm=31` (full CRUDN) on `href=*`. That certificate path is not universal: the WD53 profile in issue #16 and the washer in issue #20 reject it. For those newer OCF-PKI devices, the `SamsungServerProfile` and `ServerCertificateAuth` providers (see Quick start) pin and verify the device's hardware certificate, but getting an authorized client credential to reach protected resources is still an open problem.
---
@@ -155,7 +323,7 @@ nmap -Pn -sU -p 5683,5684,49152-49160 "$APPLIANCE_IP"
Read the result:
- **`5684/udp` or a 4915x port with a DTLS first-flight response** → an OCF DTLS listener. Standard-port OCF-PKI firmware may still require an unsupported authentication profile.
- **`5684/udp` or a 4915x port with a DTLS first-flight response** → an OCF DTLS listener. Standard-port OCF-PKI firmware needs the Samsung server-certificate profile (`SamsungServerProfile` / `ServerCertificateAuth`, see Quick start), and no working client credential for it exists yet.
- **`5683/udp` responds to public OCF security/resource GETs** → use `/oic/res` to learn the device's advertised secure endpoint; do not assume that endpoint is fixed.
- **Only `8888/tcp` open (token-based HTTPS)** → older firmware (~2018–2022). **Not supported here.**
@@ -494,6 +662,8 @@ smartthings_local/ The installable library — `pip install sm
coap.py CoAP wire protocol: message encode/decode, token handling
dtls_session.py DTLS session: handshake, client-cert auth (file or in-memory PEM), Block2, liveness
dtls_probe.py Stateless DTLS liveness + opt-in stateful diagnostic
dtls_handshake.py Shared memory-BIO handshake driver, bounded by a monotonic deadline (used by session + probe)
owner_psk.py Pure manufacturer-certificate OwnerPSK derivation
ocf_root_ca.pem Samsung OCF root CA, bundled for handshake verification
ocf/ OCF resource + state layer (reusable)
__init__.py
@@ -520,7 +690,7 @@ mqtt_demo/ MQTT bridge demo (consumes smartthings_loca
.env.example Template — copy to .env, fill in
setup_cert.py One-shot cert minting script (live-fetches AC14K_M + UUID)
pyproject.toml Packaging — PyPI dist `smartthings-local`, hatch-vcs versioning
tests/ pytest suite (CoAP wire, state cache, import isolation, cert loading, DTLS probe, bridge port resolution, cert signing)
tests/ pytest suite (CoAP wire, state cache, import isolation, cert loading, DTLS probe, bridge port resolution, cert signing, certificate profiles, OwnerPSK derivation, connect deadline, session interruption)
.github/workflows/publish.yml Build + PyPI Trusted Publishing on `v*` tags
```
+1 -1
View File
@@ -166,7 +166,7 @@ def main():
stopping.set()
logger.info("shutting down…")
for b in bridges:
b.stop.set()
b.request_stop()
try: b.set_availability(False)
except Exception: pass
+45 -6
View File
@@ -87,6 +87,7 @@ OCF_STANDARD_SECURE_PORT = 5684
# fixed-source-port reconnect invariant is untouched (see session_once).
_GATE_RETRIES = 1
_GATE_TIMEOUT_S = 4.0
_WORKER_JOIN_TIMEOUT_S = 2.0
class PushBridge:
@@ -123,6 +124,8 @@ class PushBridge:
self.last_cycle_pub = None
self.last_avail_pub: str | None = None
self.stop = threading.Event()
self._session_stop_lock = threading.Lock()
self._session_stop: threading.Event | None = None
self.started_ts = time.time()
self.session_started_ts = None
self.last_change_ts = None
@@ -175,6 +178,13 @@ class PushBridge:
app.topic_prefix, shared.HA_DISCOVERY_PREFIX, app.device_name,
model=descriptor.name.title()))
def request_stop(self) -> None:
"""Stop the bridge and wake workers belonging to its current session."""
self.stop.set()
with self._session_stop_lock:
if self._session_stop is not None:
self._session_stop.set()
# ---- cache plumbing ---------------------------------------------
def _on_cache_change(self, changed: bool, source: str) -> None:
@@ -423,25 +433,54 @@ class PushBridge:
self.keepalive = keepalive
self.observe_refresh = observe_refresh
# These workers belong to this DTLS session, not to the bridge
# process. A reconnect must retire them before the replacement
# session starts or they continue operating on the closed session.
session_stop = threading.Event()
with self._session_stop_lock:
self._session_stop = session_stop
# ``request_stop()`` sets the bridge event before taking this
# lock. Checking it while publishing the handle prevents a lost
# wakeup if shutdown races this session handoff.
if self.stop.is_set():
session_stop.set()
sched_t = threading.Thread(
target=scheduler.run_forever, args=(self.stop,),
target=scheduler.run_forever, args=(session_stop,),
daemon=True, name=f'{self.app.klass}-poll')
ka_t = threading.Thread(
target=keepalive.run_forever, args=(self.stop,),
target=keepalive.run_forever, args=(session_stop,),
daemon=True, name=f'{self.app.klass}-ping')
ref_t = threading.Thread(
target=observe_refresh.run_forever, args=(self.stop,),
target=observe_refresh.run_forever, args=(session_stop,),
daemon=True, name=f'{self.app.klass}-obsref')
sched_t.start()
ka_t.start()
ref_t.start()
workers = (sched_t, ka_t, ref_t)
started_workers = []
try:
for worker in workers:
worker.start()
started_workers.append(worker)
sess.join()
finally:
# A worker already inside a tick can finish after the reader
# exits. Disable old-session reachability callbacks first so it
# cannot change availability after a replacement takes over.
keepalive.on_reachable = None
keepalive.on_unreachable = None
session_stop.set()
join_deadline = time.monotonic() + _WORKER_JOIN_TIMEOUT_S
for worker in started_workers:
worker.join(max(0.0, join_deadline - time.monotonic()))
if worker.is_alive():
self.log.warning(
"session worker did not stop: %s", worker.name)
self.scheduler = None
self.keepalive = None
self.observe_refresh = None
with self._session_stop_lock:
if self._session_stop is session_stop:
self._session_stop = None
def _seed_from_device0(self, sess):
code, pl = sess.get(self.descriptor.seed_path, timeout=15.0)
+8
View File
@@ -166,6 +166,14 @@ def flatten(links):
if temps_items:
cur_c = _int(temps_items[0].get('x.com.samsung.da.current'))
des_c = _int(temps_items[0].get('x.com.samsung.da.desired'))
# With no cycle set the oven reports desired=0. That means "no
# setpoint", not a 0 °C target, and HA rejects it against the Number
# entity's 30-270 range on every publish. Anything outside the
# settable band is absent, not a value: null lands as unknown on both
# the Number and the Setpoint sensor, the way completion_minutes
# already reads when idle. _setpoint applies the same bounds on write.
if des_c is not None and not (SETPOINT_MIN_C <= des_c <= SETPOINT_MAX_C):
des_c = None
# Door
doors_items = g('/doors/vs/0', 'x.com.samsung.da.items') or []
+414 -10
View File
@@ -2,15 +2,37 @@
from __future__ import annotations
import logging
import re
import warnings
from enum import Enum
from os import PathLike
from pathlib import Path
from typing import Protocol, runtime_checkable
from uuid import UUID
from OpenSSL import SSL, crypto
from cryptography.x509.oid import ExtensionOID
from OpenSSL import SSL, _util, crypto
logger = logging.getLogger(__name__)
_OCF_ROOT_CA = str(Path(__file__).with_name("ocf_root_ca.pem"))
_DTLS_CIPHERS = b"ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0"
_DTLS_PSK_CIPHERS = b"ECDHE-PSK-AES128-CBC-SHA256:@SECLEVEL=0"
_SAMSUNG_SERVER_CURVES = b"prime256v1"
_SAMSUNG_SERVER_SIGNATURE_ALGORITHMS = (
b"RSA+SHA256:ECDSA+SHA256:RSA+SHA1:ECDSA+SHA1"
)
_SAMSUNG_SERVER_CN_RE = re.compile(
r"\AOCF Device: [^()\r\n]{1,96} "
r"\((?P<device_identity>[0-9a-f]{8}-[0-9a-f]{4}-"
r"[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12})\)\Z",
re.IGNORECASE,
)
_PSK_CLIENT_CALLBACK_CDEF = (
"unsigned int (*)(SSL *, char *, char *, unsigned int, "
"unsigned char *, unsigned int)"
)
_PEM_CERT_RE = re.compile(
rb"-----BEGIN CERTIFICATE-----.*?-----END CERTIFICATE-----",
re.DOTALL,
@@ -45,7 +67,300 @@ class AuthenticationProvider(Protocol):
"""Configure authentication for a newly created DTLS context."""
def configure_context(self, context: SSL.Context) -> None:
"""Configure trust, verification, ciphers, and client credentials."""
"""Configure a context while this provider remains session-owned."""
class SamsungServerRole(Enum):
"""Known Samsung OCF hardware-certificate subject roles."""
HOME_APPLIANCE = "OCF HA Device"
VD_DEVICE = "OCF VD Device"
class SamsungServerProfile:
"""Opt-in Samsung hardware-certificate verification profile."""
__slots__ = (
"_additional_ca_certificates",
"_expected_certificate_identity",
"_role",
)
def __init__(
self,
*,
expected_certificate_identity: UUID | str,
role: SamsungServerRole = SamsungServerRole.HOME_APPLIANCE,
additional_ca_pem: str | None = None,
) -> None:
if type(expected_certificate_identity) is UUID:
parsed_identity = expected_certificate_identity
elif type(expected_certificate_identity) is str:
try:
parsed_identity = UUID(expected_certificate_identity)
except ValueError:
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
) from None
if expected_certificate_identity != str(parsed_identity):
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
)
else:
raise TypeError(
"expected_certificate_identity must be a UUID or string"
)
if parsed_identity.int == 0:
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
)
if type(role) is not SamsungServerRole:
raise TypeError("role must be a SamsungServerRole")
certificates: tuple[bytes, ...] = ()
if additional_ca_pem is not None:
if type(additional_ca_pem) is not str:
raise TypeError("additional_ca_pem must be a string")
try:
raw_ca_pem = additional_ca_pem.encode("ascii")
except UnicodeEncodeError:
raise ValueError(
"additional_ca_pem must contain ASCII PEM certificates"
) from None
parsed_certificates = tuple(_PEM_CERT_RE.findall(raw_ca_pem))
if (
not 1 <= len(parsed_certificates) <= 4
or len(raw_ca_pem) > 32 * 1024
or _PEM_CERT_RE.sub(b"", raw_ca_pem).strip()
):
raise ValueError(
"additional_ca_pem must contain one to four PEM certificates"
)
try:
loaded_certificates = [
crypto.load_certificate(crypto.FILETYPE_PEM, certificate)
for certificate in parsed_certificates
]
basic_constraints = [
[
extension
for extension in certificate.to_cryptography().extensions
if extension.oid == ExtensionOID.BASIC_CONSTRAINTS
]
for certificate in loaded_certificates
]
except (crypto.Error, ValueError):
raise ValueError(
"additional_ca_pem contains an invalid certificate"
) from None
if any(
len(constraints) != 1 or not constraints[0].value.ca
for constraints in basic_constraints
):
raise ValueError(
"additional_ca_pem must contain only CA certificates"
)
fingerprints = {
crypto.dump_certificate(crypto.FILETYPE_ASN1, certificate)
for certificate in loaded_certificates
}
if len(fingerprints) != len(loaded_certificates):
raise ValueError(
"additional_ca_pem must not contain duplicate certificates"
)
certificates = parsed_certificates
object.__setattr__(
self,
"_expected_certificate_identity",
parsed_identity,
)
object.__setattr__(self, "_role", role)
object.__setattr__(self, "_additional_ca_certificates", certificates)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("SamsungServerProfile is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("SamsungServerProfile is immutable")
@classmethod
def bound_device(
cls,
expected_certificate_identity: UUID | str,
*,
role: SamsungServerRole = SamsungServerRole.HOME_APPLIANCE,
additional_ca_pem: str | None = None,
) -> SamsungServerProfile:
"""Bind a verified Samsung hardware leaf to its certificate UUID."""
return cls(
expected_certificate_identity=expected_certificate_identity,
role=role,
additional_ca_pem=additional_ca_pem,
)
def __repr__(self) -> str:
"""Return a representation without device or trust-chain details."""
return "SamsungServerProfile()"
def _configure_context(self, context: SSL.Context) -> None:
curve_setter = getattr(_util.lib, "SSL_CTX_set1_curves_list", None)
if curve_setter is not None:
if curve_setter(context._context, _SAMSUNG_SERVER_CURVES) != 1:
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
)
else:
# pyOpenSSL 23.1 does not expose SSL_CTX_set1_curves_list. Its
# public set_tmp_ecdh fallback produces the same single P-256
# supported-groups ClientHello extension; wire-level tests protect
# that compatibility path. Newer pyOpenSSL uses the exact setter.
try:
with warnings.catch_warnings():
warnings.simplefilter("ignore", DeprecationWarning)
curve = crypto.get_elliptic_curve("prime256v1")
context.set_tmp_ecdh(curve)
except (AttributeError, TypeError, ValueError, SSL.Error):
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
) from None
signature_setter = getattr(
_util.lib,
"SSL_CTX_set1_sigalgs_list",
None,
)
if (
signature_setter is None
or signature_setter(
context._context,
_SAMSUNG_SERVER_SIGNATURE_ALGORITHMS,
)
!= 1
):
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
)
context.set_options(SSL.OP_NO_TICKET)
if self._additional_ca_certificates:
store = context.get_cert_store()
try:
for certificate in self._additional_ca_certificates:
store.add_cert(
crypto.load_certificate(
crypto.FILETYPE_PEM,
certificate,
)
)
except crypto.Error:
raise RuntimeError(
"OpenSSL rejected the Samsung server trust profile"
) from None
def _verify_peer(
self,
_connection,
certificate,
_error,
depth,
ok,
) -> bool:
if not ok or certificate is None or depth < 0:
return False
if depth > 0:
return True
try:
with warnings.catch_warnings():
# pyOpenSSL deprecates this API in favor of cryptography, but
# reparsing Samsung's non-DER factory leaves with cryptography
# rejects certificates that OpenSSL has already verified.
warnings.simplefilter("ignore", DeprecationWarning)
components = [
(name.decode("ascii"), value.decode("ascii"))
for name, value in certificate.get_subject().get_components()
]
except (
AttributeError,
TypeError,
UnicodeDecodeError,
ValueError,
crypto.Error,
):
logger.warning("Unable to parse Samsung server certificate subject")
return False
common_names = [value for name, value in components if name == "CN"]
organizational_units = [
value for name, value in components if name == "OU"
]
organizations = [value for name, value in components if name == "O"]
countries = [value for name, value in components if name == "C"]
# This deliberately pins the complete Samsung subject role. The OCF
# reference implementation reads only the UUID-bearing CN, but
# relaxing C/O/OU here could accept a different certificate cohort.
if (
len(common_names) != 1
or organizational_units != [self._role.value]
or organizations != ["Samsung Electronics"]
or countries != ["KR"]
):
return False
match = _SAMSUNG_SERVER_CN_RE.fullmatch(common_names[0])
return (
match is not None
and UUID(match.group("device_identity"))
== self._expected_certificate_identity
)
def _configure_certificate_server(
context: SSL.Context,
server_profile: SamsungServerProfile | None,
) -> None:
"""Configure certificate-server verification for one DTLS context."""
context.load_verify_locations(_OCF_ROOT_CA)
if server_profile is None:
context.set_verify(SSL.VERIFY_PEER, _verify_peer)
else:
server_profile._configure_context(context)
context.set_verify(
SSL.VERIFY_PEER,
server_profile._verify_peer,
)
# @SECLEVEL=0 permits SHA-1 in Samsung's server cert chain (AC14K_M
# intermediate is SHA-1 signed). This is the only channel that reaches
# the OpenSSL instance cryptography bundles; ctypes and cffi bindings
# do not expose SSL_CTX_set_security_level on this build.
context.set_cipher_list(_DTLS_CIPHERS)
class ServerCertificateAuth:
"""Verify a pinned Samsung server without a client certificate."""
__slots__ = ("_server_profile",)
def __init__(self, *, server_profile: SamsungServerProfile) -> None:
if type(server_profile) is not SamsungServerProfile:
raise TypeError("server_profile must be a SamsungServerProfile")
object.__setattr__(self, "_server_profile", server_profile)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("ServerCertificateAuth is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("ServerCertificateAuth is immutable")
def __repr__(self) -> str:
"""Return a representation without server identity or trust details."""
return "ServerCertificateAuth()"
def configure_context(self, context: SSL.Context) -> None:
"""Verify the selected server profile without loading client material."""
_configure_certificate_server(context, self._server_profile)
class CertificateAuth:
@@ -60,6 +375,7 @@ class CertificateAuth:
"_certificate_pem",
"_private_key_path",
"_private_key_pem",
"_server_profile",
)
def __init__(
@@ -69,6 +385,7 @@ class CertificateAuth:
private_key_path: str | PathLike[str] | None = None,
certificate_pem: str | None = None,
private_key_pem: str | None = None,
server_profile: SamsungServerProfile | None = None,
) -> None:
file_supplied = (
certificate_path is not None or private_key_path is not None
@@ -96,6 +413,11 @@ class CertificateAuth:
"must pass either certificate_path/private_key_path or "
"certificate_pem/private_key_pem"
)
if (
server_profile is not None
and type(server_profile) is not SamsungServerProfile
):
raise TypeError("server_profile must be a SamsungServerProfile")
object.__setattr__(
self,
"_certificate_path",
@@ -108,6 +430,7 @@ class CertificateAuth:
)
object.__setattr__(self, "_certificate_pem", certificate_pem)
object.__setattr__(self, "_private_key_pem", private_key_pem)
object.__setattr__(self, "_server_profile", server_profile)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("CertificateAuth is immutable")
@@ -120,11 +443,14 @@ class CertificateAuth:
cls,
certificate_path: str | PathLike[str],
private_key_path: str | PathLike[str],
*,
server_profile: SamsungServerProfile | None = None,
) -> CertificateAuth:
"""Create a provider backed by certificate-chain and key files."""
return cls(
certificate_path=certificate_path,
private_key_path=private_key_path,
server_profile=server_profile,
)
@classmethod
@@ -132,11 +458,14 @@ class CertificateAuth:
cls,
certificate_pem: str,
private_key_pem: str,
*,
server_profile: SamsungServerProfile | None = None,
) -> CertificateAuth:
"""Create a provider backed by an in-memory PEM chain and key."""
return cls(
certificate_pem=certificate_pem,
private_key_pem=private_key_pem,
server_profile=server_profile,
)
def __repr__(self) -> str:
@@ -145,13 +474,7 @@ class CertificateAuth:
def configure_context(self, context: SSL.Context) -> None:
"""Apply the existing certificate authentication profile to a context."""
context.load_verify_locations(_OCF_ROOT_CA)
context.set_verify(SSL.VERIFY_PEER, _verify_peer)
# @SECLEVEL=0 permits SHA-1 in Samsung's server cert chain (AC14K_M
# intermediate is SHA-1 signed). This is the only channel that reaches
# the OpenSSL instance cryptography bundles; ctypes and cffi bindings
# do not expose SSL_CTX_set_security_level on this build.
context.set_cipher_list(_DTLS_CIPHERS)
_configure_certificate_server(context, self._server_profile)
if self._certificate_pem is not None:
_load_pem_chain(
context,
@@ -164,4 +487,85 @@ class CertificateAuth:
context.check_privatekey()
__all__ = ["AuthenticationProvider", "CertificateAuth"]
class PskAuth:
"""DTLS authentication using an existing OCF PSK credential.
The identity must be a raw 16-byte OCF UUID. The key must contain 16 or
32 bytes. Credential material is intentionally not exposed as public
attributes and is never included in this provider's representation. A
configured context must not outlive this provider; ``DtlsCoapSession``
enforces that lifetime by retaining its provider.
"""
__slots__ = ("_callback",)
def __init__(self, *, identity: bytes, key: bytes) -> None:
if type(identity) is not bytes or type(key) is not bytes:
raise TypeError("identity and key must be bytes")
if len(identity) != 16:
raise ValueError("identity must be a raw 16-byte OCF UUID")
if b"\x00" in identity:
raise ValueError("identity cannot contain a NUL byte")
if len(key) not in (16, 32):
raise ValueError("key must be 16 or 32 bytes")
ffi = _util.ffi
@ffi.callback(_PSK_CLIENT_CALLBACK_CDEF)
def client_callback(
_ssl,
_identity_hint,
identity_buffer,
max_identity_length,
key_buffer,
max_key_length,
):
# OpenSSL callbacks cannot propagate Python exceptions. Fail
# before touching either destination when a buffer is unavailable
# or too small for the complete credential.
if (
identity_buffer == ffi.NULL
or key_buffer == ffi.NULL
or len(identity) + 1 > max_identity_length
or len(key) > max_key_length
):
return 0
ffi.memmove(
identity_buffer,
identity + b"\x00",
len(identity) + 1,
)
ffi.memmove(key_buffer, key, len(key))
return len(key)
object.__setattr__(self, "_callback", client_callback)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("PskAuth is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("PskAuth is immutable")
def __repr__(self) -> str:
"""Return a representation that never includes credential material."""
return "PskAuth()"
def configure_context(self, context: SSL.Context) -> None:
"""Configure one context for the narrow Samsung OCF PSK profile."""
setter = getattr(_util.lib, "SSL_CTX_set_psk_client_callback", None)
if setter is None:
raise RuntimeError(
"the installed OpenSSL binding does not support DTLS PSK"
)
context.set_cipher_list(_DTLS_PSK_CIPHERS)
setter(context._context, self._callback)
__all__ = [
"AuthenticationProvider",
"CertificateAuth",
"PskAuth",
"SamsungServerProfile",
"SamsungServerRole",
"ServerCertificateAuth",
]
+9
View File
@@ -13,6 +13,7 @@ from ..errors import MalformedMessageError
URI_PATH = 11
URI_QUERY = 15
OBSERVE = 6
ETAG = 4
CONTENT_FORMAT = 12
ACCEPT = 17
BLOCK2 = 23
@@ -121,6 +122,14 @@ def block_value(num, more, szx):
return struct.pack('>I', v)[1:]
def block_fields(value):
"""Decode a CoAP Block-N option value. Inverse of block_value().
Returns (num, more, szx). An empty value means block 0, no more,
SZX=0 — RFC 7959 §2.2 allows a zero-length option to elide it."""
v = int.from_bytes(value, 'big')
return v >> 4, (v >> 3) & 1, v & 0x07
def fmt_code(c):
"""0x45 → '2.05', 0x84 → '4.04'. Used in log lines."""
return f"{c >> 5}.{c & 0x1F:02d}"
@@ -0,0 +1,99 @@
"""Shared memory-BIO driver for bounded DTLS handshakes."""
from __future__ import annotations
import select
import time
from collections.abc import Callable
from OpenSSL import SSL
from .coap import split_dtls
_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,
*,
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.
OpenSSL owns the retransmission schedule. ``retries`` optionally limits
how many expired retransmission timers are serviced; a normal session is
bounded only by its deadline, while the diagnostic probe retains its
explicit retry budget.
Return ``True`` once OpenSSL reports the handshake complete. The deadline
prevents another setup, retry, or network-wait iteration; it does not tear
down a session that completed while ``do_handshake()`` was running. TLS and
socket failures are left to the caller to classify.
"""
retransmits = 0
while time.monotonic() < deadline:
try:
connection.do_handshake()
return True
except SSL.WantReadError:
pass
try:
output = connection.bio_read(_MAX_DATAGRAM_SIZE)
except SSL.WantReadError:
output = None
if output:
for record in split_dtls(output):
if sock.send(record) != len(record):
raise OSError("incomplete UDP send")
remaining = deadline - time.monotonic()
if remaining <= 0:
break
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:
break
connection.DTLSv1_handle_timeout()
retransmits += 1
continue
if not datagram:
continue
if on_datagram is not None:
on_datagram(datagram)
connection.bio_write(datagram)
return False
+33 -65
View File
@@ -39,8 +39,9 @@ from dataclasses import dataclass
from OpenSSL import SSL
from ..errors import ProbeError
from .auth import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
from .coap import split_dtls
from .dtls_session import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
from .dtls_handshake import _drive_dtls_handshake
from .endpoint import open_connected_udp_socket
# DTLS record content types (RFC 6347 §4.1)
@@ -603,73 +604,40 @@ def diagnose_dtls_handshake(
started = time.monotonic()
deadline = started + timeout
seen = set()
retransmits = 0
def record_datagram(datagram):
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
result.datagrams.append(datagram)
for content_type, detail in classify_datagram(datagram):
if content_type == _CT_HANDSHAKE:
if detail not in seen:
seen.add(detail)
result.handshake_msgs.append(detail)
if result.outcome == DEAD:
result.outcome = LIVE
elif content_type == _CT_ALERT and detail is not None:
level, name = detail
result.alert = (level, name)
if level == 2: # fatal
result.outcome = REJECTED
try:
while time.monotonic() < deadline:
try:
conn.do_handshake()
result.outcome = COMPLETED
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
break
except SSL.WantReadError:
pass
except SSL.Error:
# A fatal Alert lands here; the alert record was already
# captured below, so classification still works.
result.error = ProbeError()
break
try:
o = conn.bio_read(65535)
if o:
for r in split_dtls(o):
if sock.send(r) != len(r):
raise OSError('short UDP send')
except SSL.WantReadError:
pass
remaining = deadline - time.monotonic()
if remaining <= 0:
break
sock.settimeout(min(0.5, remaining))
try:
d = sock.recv(65535)
except TimeoutError:
# No answer to the last flight. Service OpenSSL's DTLS
# retransmit timer: once it has counted down to 0,
# handle_timeout() re-queues the previous flight into the
# write BIO for the next iteration to flush. A live server
# answers within a flight or two; a silent/non-DTLS port
# never does, so we give up only after `retries`
# retransmits — one dropped ClientHello no longer reads as
# a false DEAD.
to = conn.DTLSv1_get_timeout()
if to is not None and to <= 0:
if retransmits >= retries:
break
conn.DTLSv1_handle_timeout()
retransmits += 1
continue
if not d:
continue
completed = _drive_dtls_handshake(
conn,
sock,
deadline=deadline,
retries=retries,
on_datagram=record_datagram,
)
if completed:
result.outcome = COMPLETED
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
result.datagrams.append(d)
for ct, detail in classify_datagram(d):
if ct == _CT_HANDSHAKE:
if detail not in seen:
seen.add(detail)
result.handshake_msgs.append(detail)
if result.outcome == DEAD:
result.outcome = LIVE
elif ct == _CT_ALERT and detail is not None:
level, name = detail
result.alert = (level, name)
if level == 2: # fatal
result.outcome = REJECTED
conn.bio_write(d)
except SSL.Error:
# A fatal Alert lands here; record_datagram() has already classified
# the alert record before it is fed back into OpenSSL.
result.error = ProbeError()
except OSError:
result.error = ProbeError()
finally:
+575 -126
View File
@@ -12,12 +12,21 @@ Wire-level details that matter (from local-tools/oven-findings.md §17):
or interleaved one-shot / OBSERVE traffic mis-attributes.
* Multi-block GET requires the SAME CoAP token across every block
of the response ("token-stable Block2"). Fresh-token-per-block
is silently dropped by the server.
is silently dropped by the server, and a transfer that opens at
NUM>0 under a token the server has not seen gets no reply at all.
Reader thread owns the UDP socket. Callers issue get()/post() and block
on a per-token Event the reader signals. OBSERVE notifications are
delivered via the on_notification callback.
A notification carries only the first block of a large representation
(RFC 7959 §2.6). Because a continuation cannot borrow the observation's
token (§3.4) and this server will not continue a transfer it did not
start, such a notification is withheld and the resource is re-read from
block 0 on a fresh one-shot token by a worker thread. See #39.
"""
import errno
import math
import os
import socket
import threading
@@ -33,11 +42,12 @@ from ..errors import (
SessionTimeoutError,
)
from .coap import (
URI_PATH, URI_QUERY, OBSERVE, CONTENT_FORMAT, ACCEPT, BLOCK2, SIZE2,
URI_PATH, URI_QUERY, OBSERVE, ETAG, CONTENT_FORMAT, ACCEPT, BLOCK2, SIZE2,
TYPE_CON, TYPE_NON, TYPE_ACK, TYPE_RST,
METHOD_GET, METHOD_POST, CF_CBOR,
OBSERVE_REGISTER, OBSERVE_DEREGISTER, BLOCK_SZX,
encode_options, parse_coap, build_coap, block_value, fmt_code,
encode_options, parse_coap, build_coap, block_value, block_fields,
fmt_code,
split_dtls as _split_dtls,
)
from .auth import (
@@ -47,6 +57,11 @@ from .auth import (
_OCF_ROOT_CA,
_load_pem_chain,
)
from .dtls_handshake import (
_HANDSHAKE_POLL_S,
_HandshakeCancelled,
_drive_dtls_handshake,
)
from .endpoint import open_connected_udp_socket
import logging
@@ -65,6 +80,11 @@ DEBUG_BRIDGE = os.environ.get('DEBUG_BRIDGE') == '1'
_BLOCK_MAX_ATTEMPTS = 3
_BLOCK_ACK_TIMEOUT = 4.0
# How often a block wait re-checks that the reader is still alive. Short
# enough that a mid-transfer reader death fails fast instead of burning
# the whole per-block timeout, long enough to stay off the CPU.
_BLOCK_LIVENESS_POLL_S = 0.25
# Inter-request pacing: minimum seconds between CoAP CON sends on one session.
# Samsung's RT-OCF stacks drop requests when hit faster than their firmware
# ceiling (dryer ~14 req/s, oven ~8 req/s, dishwasher unknown). 5 req/s
@@ -72,6 +92,106 @@ _BLOCK_ACK_TIMEOUT = 4.0
# once the ceiling is measured empirically.
_DEFAULT_RATE_LIMIT_RPS = 5.0
# Maximum hrefs 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
# Timeout for one notification refetch. Generous relative to a poll:
# the resource is known large (that is why it blocked) and the worker
# is serialized, so a slow one delays only later refetches.
_REFETCH_TIMEOUT_S = 15.0
class _EtagChanged(Exception):
"""Internal: the server's ETag changed partway through a Block2
transfer, so the blocks in hand are from two different versions."""
# ICMP errors a connected UDP socket surfaces on the next recv. On these
# appliances they show up while the device is rebooting, while it holds an
# orphaned association, or across a router blip, and the next datagram
# usually works. UDP delivery was never guaranteed, so treat them as
# advisory and keep reading. Unconnected sockets never see any of this,
# which is why the reader survived them before the connected-socket change
# in d677c72 (v0.1.3).
_ADVISORY_ERRNOS = frozenset(
value for value in (
getattr(errno, name, None)
for name in ('ECONNREFUSED', 'EHOSTUNREACH', 'ENETUNREACH',
'EHOSTDOWN', 'ENETDOWN')
) if value is not None
)
def _validate_handshake_timeout(timeout, default):
"""Return one finite, positive DTLS handshake timeout."""
value = default if timeout is None else timeout
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise TypeError('timeout must be a number or None')
try:
value = float(value)
except OverflowError:
raise ValueError(
'timeout must be a positive finite number or None') from None
if not math.isfinite(value) or value <= 0:
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.
@@ -150,6 +270,11 @@ class DtlsCoapSession:
self.endpoint = None
self._send_lock = threading.Lock()
# Guards the MID/token counters and _pending. The refetch worker
# makes the session its own second concurrent get() caller, so
# two threads can mint tokens at once; without this they can
# collide and one transfer silently absorbs the other's blocks.
self._state_lock = threading.Lock()
# 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
@@ -166,8 +291,20 @@ class DtlsCoapSession:
# 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
# because _dispatch_coap cannot block (see _queue_refetch).
self._refetch_cond = threading.Condition()
self._refetch_pending = {}
self._refetch_seq = 0
self._refetch_thread = None
self._stop = threading.Event()
self._reader_thread = None
# Set while the reader owns the socket. Cleared when it exits for
# any reason, so callers fail fast through _check_live() instead of
# waiting out a request timeout against a session nobody is reading.
self._reader_running = threading.Event()
self._last_send_ts = 0.0
def pace(self) -> None:
@@ -179,69 +316,106 @@ class DtlsCoapSession:
# ---- lifecycle ---------------------------------------------------
def connect(self):
"""DTLS handshake. Blocks up to HANDSHAKE_TIMEOUT_S. Raises
ConnectionError / TimeoutError on failure."""
def connect(
self,
*,
timeout: float | None = None,
cancel: ConnectCancellation | None = None,
):
"""Perform a cancellable DTLS handshake using 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.
Once OpenSSL reports completion, that completed session is retained
even if the call returns just after the deadline. 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:
raise SessionTimeoutError()
sock, endpoint = open_connected_udp_socket(
self.host,
self.port,
family=self.family,
local_port=self.local_port,
timeout=2.0,
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"
)
t0 = time.time()
backend_failed = False
while time.time() - t0 < self.HANDSHAKE_TIMEOUT_S:
io_failed = False
cancelled = False
interrupted = False
completed = False
try:
try:
conn.do_handshake()
break
except SSL.WantReadError:
pass
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:
sock.close()
backend_failed = True
break
send_failed = False
try:
o = conn.bio_read(65535)
if o:
for r in _split_dtls(o):
if sock.send(r) != len(r):
raise OSError('incomplete UDP send')
except SSL.WantReadError:
pass
except OSError:
sock.close()
send_failed = True
if send_failed:
raise EndpointError() from OSError('UDP send failed')
receive_failed = False
try:
d = sock.recv(65535)
if d:
conn.bio_write(d)
except socket.timeout:
pass
except OSError:
sock.close()
receive_failed = True
if receive_failed:
raise EndpointError() from OSError('UDP receive failed')
time.sleep(0.05)
else:
io_failed = True
finally:
if wake_subscription is not None:
interrupted = cancel._unsubscribe(*wake_subscription)
if cancelled or (interrupted and not completed):
sock.close()
raise SessionClosedError()
if backend_failed:
sock.close()
raise SessionError() from ConnectionError('DTLS backend failed')
if io_failed:
sock.close()
raise EndpointError() from OSError('UDP handshake I/O failed')
if not completed:
sock.close()
raise SessionTimeoutError()
if backend_failed:
raise SessionError() from ConnectionError('DTLS backend failed')
self.sock = sock
self.conn = conn
@@ -253,15 +427,33 @@ class DtlsCoapSession:
"""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
def _check_live(self):
"""Raise if the session cannot carry a request. A dead reader is
as fatal as a closed connection: the socket may still accept
sends, but no response will ever be dispatched, so waiting out the
request timeout only delays the inevitable SessionClosedError.
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:
raise SessionClosedError()
if self._reader_thread is not None and \
not self._reader_running.is_set():
raise SessionClosedError()
def join(self):
"""Block until the reader thread exits (i.e. socket dies)."""
if self._reader_thread is not None:
self._reader_thread.join()
if self._refetch_thread is not None:
self._refetch_thread.join()
def _send_observe_dereg(self, tok, path_segs):
"""Send a single OBSERVE deregister GET (Observe option = 1)
@@ -293,6 +485,11 @@ class DtlsCoapSession:
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()
@@ -303,10 +500,12 @@ class DtlsCoapSession:
self.sock.close()
except Exception:
pass
for tok, (ev, container) in list(self._pending.items()):
with self._state_lock:
pending = list(self._pending.items())
self._pending.clear()
for tok, (ev, container) in pending:
container.setdefault('err', SessionClosedError())
ev.set()
self._pending.clear()
self._observe_tokens.clear()
self.sock = None
self.conn = None
@@ -316,27 +515,30 @@ class DtlsCoapSession:
# ---- send / receive plumbing -------------------------------------
def _next_mid(self):
self._mid = (self._mid + 1) & 0xFFFF
return self._mid
with self._state_lock:
self._mid = (self._mid + 1) & 0xFFFF
return self._mid
def _next_tok(self):
self._tok_counter = (self._tok_counter + 1) & 0xFFFFFFFF
# 4-byte tokens — fits within tkl=8 cap with headroom and
# avoids collisions across long-running OBSERVE subscriptions.
return self._tok_counter.to_bytes(4, 'big')
with self._state_lock:
self._tok_counter = (self._tok_counter + 1) & 0xFFFFFFFF
# 4-byte tokens — fits within tkl=8 cap with headroom and
# avoids collisions across long-running OBSERVE subscriptions.
return self._tok_counter.to_bytes(4, 'big')
def _next_observe_tok(self):
# 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.
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])
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.
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])
def _send_dgram(self, datagram):
"""Send a CoAP datagram. Holds the send lock for the
@@ -374,7 +576,23 @@ class DtlsCoapSession:
d = sock.recv(65535)
except socket.timeout:
continue
except (OSError, ValueError):
except OSError as e:
if self._stop.is_set():
return # close() got here first
if e.errno in _ADVISORY_ERRNOS:
logger.debug("reader: advisory %s from %s, continuing",
errno.errorcode.get(e.errno, e.errno),
self.host)
continue
logger.warning("reader exiting: socket error %s from %s",
errno.errorcode.get(e.errno, e.errno),
self.host)
return
except ValueError:
# recv on a socket closed underneath the reader.
if not self._stop.is_set():
logger.warning("reader exiting: socket closed "
"underneath it")
return
if not d:
continue
@@ -419,10 +637,18 @@ class DtlsCoapSession:
if exit_reader:
return
finally:
# Reader no longer owns the socket — callers must fail fast.
self._reader_running.clear()
# Make sure pending waiters don't hang if the reader dies.
for tok, (ev, container) in list(self._pending.items()):
with self._state_lock:
pending = list(self._pending.items())
for tok, (ev, container) in pending:
container.setdefault('err', SessionClosedError())
ev.set()
# Nothing will answer a refetch now either.
with self._refetch_cond:
self._refetch_pending.clear()
self._refetch_cond.notify_all()
def _dispatch_coap(self, datagram):
try:
@@ -452,7 +678,8 @@ class DtlsCoapSession:
return
# Pending one-shot? Resolve and return.
rec = self._pending.get(tok)
with self._state_lock:
rec = self._pending.get(tok)
if rec is not None:
ev, container = rec
container['code'] = code
@@ -470,6 +697,16 @@ class DtlsCoapSession:
logger.warning("observe %s: non-2.05 %s",
href, fmt_code(code))
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 cb is not None:
try:
@@ -481,6 +718,110 @@ class DtlsCoapSession:
# Stale token (post-reconnect or unknown) — drop quietly.
# ---- OBSERVE refetch ---------------------------------------------
@staticmethod
def _log_refetch(msg, *args):
"""Refetch outcomes are debug-level in normal operation, which is
below the bridge's INFO default, so a healthy session stays quiet.
DEBUG_BRIDGE=1 promotes them to INFO for hardware validation:
that shows which token the re-read used and whether it completed,
without also turning on every per-block retransmit line."""
(logger.info if DEBUG_BRIDGE else logger.debug)(msg, *args)
def _queue_refetch(self, href):
"""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."""
with self._refetch_cond:
if (href 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_cond.notify()
self._start_refetch_worker()
def _start_refetch_worker(self):
"""Start the refetch worker on first use. Sessions that never see
a blockwise notification never grow the thread."""
if self._refetch_thread is not None or self._stop.is_set():
return
with self._state_lock:
if self._refetch_thread is not None or self._stop.is_set():
return
self._refetch_thread = threading.Thread(
target=self._refetch_loop, daemon=True,
name=f'stl-refetch-{self.host}')
self._refetch_thread.start()
def _refetch_loop(self):
"""Re-read blockwise-notified resources, one at a time.
Serialized on purpose. Each re-read is a multi-block transfer and
_blockwise_get paces between blocks, so running one at a time is
what keeps a notification storm under the firmware's request
ceiling."""
while self._refetch_alive():
with self._refetch_cond:
while not self._refetch_pending and self._refetch_alive():
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)
def _refetch_alive(self):
"""False once the session is closing or the reader has died. A
refetch needs the reader to resolve its token, so outliving it
would leave join() waiting on a thread with nothing to do."""
if self._stop.is_set():
return False
return self._reader_thread is None or self._reader_running.is_set()
def _refetch_one(self, href, seq):
"""Re-read one href from block 0 and deliver it if it is still
the freshest thing we know about that resource."""
self.pace()
segs = [s for s in href.split('/') if s]
try:
code, payload, blocks, tok = self._blockwise_get(
segs, (), _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
# contract: the poll tiers still carry freshness, and handing
# over the first block is the bug this replaced.
self._log_refetch("refetch %s failed: %s", href, e)
return
if code != 0x45:
self._log_refetch("refetch %s returned %s", href, fmt_code(code))
return
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:
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
if cb is not None:
try:
cb(href, payload)
except Exception as e:
logger.warning("notification callback %s: %s", href, e)
# ---- request primitives ------------------------------------------
def get(self, path_segs, query=(), timeout=10.0):
@@ -490,84 +831,189 @@ class DtlsCoapSession:
response — Samsung's server keys per-transfer state on the
token, and dropping a fresh token on block 1+ silently drops
the request."""
if self.conn is None:
raise SessionClosedError()
self._check_live()
code, blob, _blocks, _tok = self._blockwise_get(
path_segs, query, timeout)
return code, blob
def _blockwise_get(self, path_segs, query=(), timeout=10.0):
"""Shared token-stable Block2 reassembly (RFC 7959 §2.4).
Returns (code, payload, block_count, token). The last two are
diagnostics for the refetch log; get() drops them.
Mints one fresh 4-byte token and holds it across every block of
the transfer. Also the notification-refetch primitive: RFC 7959
§3.4 forbids continuing a blockwise notification on the
observation's token, and Samsung's RT-OCF drops a transfer that
opens at NUM>0 under a token it has not seen, so a truncated
notification is recovered by re-reading from block 0 through
this same path rather than by a §2.6 continuation.
Restarts once if the server's ETag changes mid-transfer, then
gives up: RFC 7959 §2.4 requires the client to compare ETags
when the server supplies them. None of the tested appliances
emit option 4, so on those this is inert."""
try:
return self._blockwise_get_once(path_segs, query, timeout)
except _EtagChanged:
logger.debug("GET %s /%s: ETag changed mid-transfer, restarting",
self.host, '/'.join(path_segs))
try:
return self._blockwise_get_once(path_segs, query, timeout)
except _EtagChanged:
logger.debug(
"GET %s /%s: representation kept changing mid-transfer",
self.host, '/'.join(path_segs))
raise BlockwiseError() from None
def _blockwise_get_once(self, path_segs, query, timeout):
"""One attempt at a full Block2 transfer. Raises _EtagChanged if
the server's representation changed while we were reassembling."""
tok = self._next_tok()
blob = b''
num = 0
blocks = 0
last_code = None
last_opts = []
etag = None
deadline = time.time() + timeout
szx = BLOCK_SZX # server may negotiate down; track per-transfer
while True:
if num > 0:
self.pace()
container = {}
for attempt in range(_BLOCK_MAX_ATTEMPTS):
ev = threading.Event()
container = {}
self._pending[tok] = (ev, container)
try:
mid = self._next_mid()
opts = [(URI_PATH, s.encode()) for s in path_segs]
for q in query:
opts.append((URI_QUERY, q.encode()))
opts.append((ACCEPT, CF_CBOR))
if num > 0:
opts.append((BLOCK2, block_value(num, 0, szx)))
self._send_dgram(
build_coap(TYPE_CON, METHOD_GET, mid, tok, opts))
per_wait = min(_BLOCK_ACK_TIMEOUT,
max(0.1, deadline - time.time()))
if ev.wait(per_wait):
break # got a response
remaining = deadline - time.time()
if remaining <= 0 or attempt == _BLOCK_MAX_ATTEMPTS - 1:
logger.debug(
"GET %s /%s block %d: timed out after %d attempt(s)",
self.host, '/'.join(path_segs), num, attempt + 1,
)
raise SessionTimeoutError()
logger.debug(
"GET %s /%s block %d: attempt %d/%d timeout, retrying",
self.host, '/'.join(path_segs), num,
attempt + 1, _BLOCK_MAX_ATTEMPTS,
)
finally:
self._pending.pop(tok, None)
self.pace()
self._check_live()
container = self._exchange_block(
tok, path_segs, query, num, szx, deadline)
if 'err' in container:
raise container['err']
blocks += 1
code = container['code']
payload = container['payload']
ropts = container['options']
last_code = code
last_opts = ropts
blob += payload
# 4.xx / 5.xx responses don't carry Block2 continuation —
# bail with whatever we got. Caller decides if 4.xx is fatal.
if code >> 5 != 2:
return code, blob
return code, blob, blocks, tok
# RFC 7959 §2.4: compare ETags across blocks, or we splice
# two versions of the resource into one buffer.
block_etag = next((v for n, v in ropts if n == ETAG), None)
if num == 0:
etag = block_etag
elif etag is not None and block_etag != etag:
raise _EtagChanged()
blob += payload
b2 = [v for n, v in ropts if n == BLOCK2]
more = 0
if b2:
bv = int.from_bytes(b2[0], 'big')
more = (bv >> 3) & 1
server_szx = bv & 0x07
if server_szx != szx:
szx = server_szx
if not b2:
break
_, more, server_szx = block_fields(b2[0])
if not more:
break
num += 1
if server_szx != szx:
# Server negotiated the block size down. Block numbers
# are indices into the new size, so the next one has to
# come off the byte offset we have actually accumulated,
# not off num + 1.
szx = server_szx
num = len(blob) >> (szx + 4)
else:
num += 1
if num > self.MAX_BLOCKS:
raise BlockwiseError()
return last_code, blob
return last_code, blob, blocks, tok
def _exchange_block(self, tok, path_segs, query, num, szx, deadline):
"""Send one block request under `tok` and return its response
container, retransmitting up to _BLOCK_MAX_ATTEMPTS times.
A response whose Block2 NUM is not the one we asked for is a
retransmit of an earlier block, not the next one. Concatenating
it would corrupt the buffer, so keep waiting on the same
attempt budget instead."""
for attempt in range(_BLOCK_MAX_ATTEMPTS):
ev = threading.Event()
container = {}
with self._state_lock:
self._pending[tok] = (ev, container)
try:
mid = self._next_mid()
opts = [(URI_PATH, s.encode()) for s in path_segs]
for q in query:
opts.append((URI_QUERY, q.encode()))
opts.append((ACCEPT, CF_CBOR))
if num > 0:
opts.append((BLOCK2, block_value(num, 0, szx)))
self._send_dgram(
build_coap(TYPE_CON, METHOD_GET, mid, tok, opts))
while True:
per_wait = min(_BLOCK_ACK_TIMEOUT,
max(0.1, deadline - time.time()))
if not self._wait_for_block(ev, per_wait):
break # attempt timed out
if 'err' in container or self._block_num_matches(
container, num):
return container
logger.debug(
"GET %s /%s block %d: stale block, still waiting",
self.host, '/'.join(path_segs), num)
ev.clear()
container.clear()
if deadline - time.time() <= 0:
break
remaining = deadline - time.time()
if remaining <= 0 or attempt == _BLOCK_MAX_ATTEMPTS - 1:
logger.debug(
"GET %s /%s block %d: timed out after %d attempt(s)",
self.host, '/'.join(path_segs), num, attempt + 1,
)
raise SessionTimeoutError()
logger.debug(
"GET %s /%s block %d: attempt %d/%d timeout, retrying",
self.host, '/'.join(path_segs), num,
attempt + 1, _BLOCK_MAX_ATTEMPTS,
)
finally:
with self._state_lock:
self._pending.pop(tok, None)
raise SessionTimeoutError()
def _wait_for_block(self, ev, per_wait):
"""Wait for one block response, giving up early if the reader
dies underneath us.
Only the reader thread can resolve a token, so once it is gone
the wait can never succeed. Polling in slices turns what would
be a full per-block timeout into an immediate SessionClosedError,
which is the same fail-fast contract get() gets from _check_live()
at entry — it just has to hold for every block, not only the
first."""
deadline = time.time() + per_wait
while True:
slice_s = min(_BLOCK_LIVENESS_POLL_S, deadline - time.time())
if slice_s <= 0:
return False
if ev.wait(slice_s):
return True
self._check_live()
@staticmethod
def _block_num_matches(container, num):
"""True if this response carries the block we asked for. A
response with no Block2 option is the whole representation, so
it only answers block 0."""
if container.get('code', 0) >> 5 != 2:
return True # error responses end the transfer either way
b2 = [v for n, v in container.get('options', ()) if n == BLOCK2]
if not b2:
return num == 0
return block_fields(b2[0])[0] == num
def post(self, path_segs, body_cbor, timeout=8.0):
"""Single-frame POST with a CBOR-encoded body. Returns
(code, payload_bytes). body_cbor must already be encoded."""
if self.conn is None:
raise SessionClosedError()
self._check_live()
tok = self._next_tok()
mid = self._next_mid()
opts = [(URI_PATH, s.encode()) for s in path_segs]
@@ -577,8 +1023,11 @@ class DtlsCoapSession:
body_cbor)
ev = threading.Event()
container = {}
self._pending[tok] = (ev, container)
with self._state_lock:
self._pending[tok] = (ev, container)
try:
self.pace()
self._check_live()
self._send_dgram(datagram)
if not ev.wait(timeout):
raise SessionTimeoutError()
@@ -586,7 +1035,8 @@ class DtlsCoapSession:
raise container['err']
return container['code'], container['payload']
finally:
self._pending.pop(tok, None)
with self._state_lock:
self._pending.pop(tok, None)
def ping(self):
"""RFC 7252 §4.4 CoAP Ping — empty CON, no token, no payload.
@@ -600,8 +1050,7 @@ class DtlsCoapSession:
Real half-open-session detection lives in PollScheduler's
`last_success_ts`, surfaced through KeepaliveTask's
`liveness_fn`."""
if self.conn is None:
raise SessionClosedError()
self._check_live()
mid = self._next_mid()
self._send_dgram(build_coap(TYPE_CON, 0, mid, b'', []))
return mid
@@ -618,8 +1067,7 @@ class DtlsCoapSession:
tokens via subscribe. Brief race window where a notify on the
old token gets dropped as 'stale' — acceptable for a 6h-scale
safety net."""
if self.conn is None:
raise SessionClosedError()
self._check_live()
for tok, href in list(self._observe_tokens.items()):
segs = [s for s in href.split('/') if s]
try:
@@ -642,8 +1090,9 @@ class DtlsCoapSession:
Returns the token used (in case the caller wants to deregister
later)."""
if self.conn is None:
raise SessionClosedError()
self._check_live()
self.pace()
self._check_live()
tok = self._next_observe_tok()
href = '/' + '/'.join(path_segs)
# Register the token BEFORE sending — otherwise the device
+322
View File
@@ -0,0 +1,322 @@
"""Bounded discovery of a known host's plaintext OCF response port.
Some OCF devices receive multicast discovery on UDP 5683 but reply from an
ephemeral port. This module records only token-correlated response ports from
the caller's expected IPv4 address. The results are candidates: callers still
need directory parsing, a DTLS probe, and authenticated identity validation.
"""
from __future__ import annotations
import ipaddress
import math
import secrets
import selectors
import socket
import time
from dataclasses import dataclass
from ..errors import MalformedMessageError
from .coap import (
ACCEPT,
CF_CBOR,
METHOD_GET,
TYPE_ACK,
TYPE_CON,
TYPE_NON,
URI_PATH,
URI_QUERY,
build_coap,
parse_coap,
)
__all__ = [
"OcfResponderPortDiscoveryResult",
"discover_ocf_responder_ports",
]
_OCF_MULTICAST_GROUP = socket.inet_ntoa(bytes((224, 0, 1, 187)))
_OCF_DISCOVERY_PORT = 5683
_OCF_CBOR = (10_000).to_bytes(2, "big")
_OCF_CONTENT_FORMAT_VERSION = 2049
_OCF_VERSION_1_0 = (2048).to_bytes(2, "big")
_CONTENT = 0x45
_MAX_DATAGRAM_BYTES = 8192
_MAX_DATAGRAMS_PER_ROUND = 64
_MAX_PORTS = 8
@dataclass(frozen=True, slots=True, repr=False)
class OcfResponderPortDiscoveryResult:
"""Redacted result of one known-host multicast discovery operation."""
ports: tuple[int, ...]
attempts: int
responses: int
error_code: str | None = None
@property
def found(self) -> bool:
"""Return whether at least one response port was discovered."""
return bool(self.ports)
def __repr__(self) -> str:
return (
"OcfResponderPortDiscoveryResult("
f"found={self.found!r}, port_count={len(self.ports)}, "
f"attempts={self.attempts}, responses={self.responses}, "
f"error_code={self.error_code!r})"
)
def _validate_address(value: object, name: str) -> tuple[str, bytes]:
if not isinstance(value, str):
raise TypeError(f"{name} must be an IPv4 address string")
try:
address = ipaddress.IPv4Address(value)
except ipaddress.AddressValueError as exc:
raise ValueError(f"{name} must be a valid IPv4 address") from exc
if address.is_multicast or address.is_unspecified or address.is_reserved:
raise ValueError(f"{name} must be a unicast IPv4 address")
return str(address), address.packed
def _validate_options(
*,
discovery_port: object,
timeout: object,
rounds: object,
) -> tuple[int, float, int]:
if isinstance(discovery_port, bool) or not isinstance(discovery_port, int):
raise TypeError("discovery_port must be an integer")
if not 1 <= discovery_port <= 65535:
raise ValueError("discovery_port must be between 1 and 65535")
if isinstance(timeout, bool) or not isinstance(timeout, (int, float)):
raise TypeError("timeout must be a number")
timeout_value = float(timeout)
if not math.isfinite(timeout_value) or not 0 < timeout_value <= 30:
raise ValueError("timeout must be greater than zero and at most 30")
if isinstance(rounds, bool) or not isinstance(rounds, int):
raise TypeError("rounds must be an integer")
if not 1 <= rounds <= 4:
raise ValueError("rounds must be between one and four")
return discovery_port, timeout_value, rounds
def _request(
token: bytes,
message_id: int,
*,
versioned: bool,
filtered: bool,
) -> bytes:
options = [
(URI_PATH, b"oic"),
(URI_PATH, b"res"),
(ACCEPT, _OCF_CBOR if versioned else CF_CBOR),
]
if filtered:
options.append((URI_QUERY, b"rt=oic.r.doxm"))
if versioned:
options.append((_OCF_CONTENT_FORMAT_VERSION, _OCF_VERSION_1_0))
return build_coap(TYPE_NON, METHOD_GET, message_id, token, options)
def _result(
ports: tuple[int, ...],
attempts: int,
responses: int,
error_code: str | None = None,
) -> OcfResponderPortDiscoveryResult:
return OcfResponderPortDiscoveryResult(
ports=ports,
attempts=attempts,
responses=responses,
error_code=error_code,
)
def discover_ocf_responder_ports(
target_address: str,
*,
interface_address: str,
discovery_port: int = _OCF_DISCOVERY_PORT,
timeout: float = 3.0,
rounds: int = 2,
) -> OcfResponderPortDiscoveryResult:
"""Find plaintext OCF response ports for one known IPv4 host.
Each round sends modern OCF and legacy IoTivity NON requests to the
link-local multicast group. Only a 2.05 response with a request token and
the exact target source address contributes a candidate. One monotonic
deadline bounds all rounds, and every socket is closed before return.
"""
_target_address, target_key = _validate_address(target_address, "target_address")
interface_address, interface_key = _validate_address(
interface_address, "interface_address"
)
discovery_port, timeout, rounds = _validate_options(
discovery_port=discovery_port,
timeout=timeout,
rounds=rounds,
)
try:
selector = selectors.DefaultSelector()
except (OSError, ValueError):
return _result((), 0, 0, "interface_unavailable")
active = None
try:
active = socket.socket(socket.AF_INET, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_IF, interface_key)
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_TTL, 1)
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_LOOP, 0)
active.bind((interface_address, 0))
active.setblocking(False)
selector.register(active, selectors.EVENT_READ)
except (OSError, ValueError):
if active is not None:
try:
active.close()
except OSError:
pass
selector.close()
return _result((), 0, 0, "interface_unavailable")
started = time.monotonic()
deadline = started + timeout
accepted_tokens: set[bytes] = set()
observations: set[tuple[bytes, int]] = set()
ports: list[int] = []
seen_ports: set[int] = set()
attempts = 0
responses = 0
too_many_ports = False
try:
for round_number in range(rounds):
if time.monotonic() >= deadline:
break
# Preserve the unfiltered modern and legacy requests used by the
# installed appliance generations. Older media firmware can omit
# usable endpoint policy from its large unfiltered directory but
# answer the smaller legacy DOXM-filtered lookup, so send that as
# a third bounded fallback rather than narrowing every request.
for versioned, filtered in (
(True, False),
(False, False),
(False, True),
):
token = secrets.token_bytes(8)
while token in accepted_tokens:
token = secrets.token_bytes(8)
accepted_tokens.add(token)
request = _request(
token,
secrets.randbits(16),
versioned=versioned,
filtered=filtered,
)
try:
sent = active.sendto(
request,
(_OCF_MULTICAST_GROUP, discovery_port),
)
except OSError:
continue
attempts += 1
if sent != len(request):
continue
round_deadline = started + timeout * (round_number + 1) / rounds
datagrams = 0
while datagrams < _MAX_DATAGRAMS_PER_ROUND:
remaining = min(deadline, round_deadline) - time.monotonic()
if remaining <= 0:
break
try:
events = selector.select(remaining)
except (OSError, ValueError):
break
if not events:
break
try:
datagram, source = active.recvfrom(_MAX_DATAGRAM_BYTES + 1)
except (BlockingIOError, OSError):
continue
datagrams += 1
if len(datagram) > _MAX_DATAGRAM_BYTES:
continue
if (
len(datagram) < 4
or datagram[0] >> 6 != 1
or datagram[0] & 0x0F > 8
or 4 + (datagram[0] & 0x0F) > len(datagram)
):
continue
if not isinstance(source, tuple) or len(source) != 2:
continue
source_host, source_port = source
if not isinstance(source_host, str):
continue
try:
source_key = socket.inet_pton(socket.AF_INET, source_host)
except OSError:
continue
if source_key != target_key:
continue
if (
isinstance(source_port, bool)
or not isinstance(source_port, int)
or not 1 <= source_port <= 65535
):
continue
try:
message_type, code, mid, token, _options, payload = parse_coap(
datagram
)
except (IndexError, ValueError, MalformedMessageError):
continue
if (
token not in accepted_tokens
or code != _CONTENT
or message_type not in (TYPE_NON, TYPE_CON)
or not payload
):
continue
if message_type == TYPE_CON:
try:
active.sendto(build_coap(TYPE_ACK, 0, mid, b"", []), source)
except OSError:
pass
observation = (token, source_port)
if observation in observations:
continue
observations.add(observation)
responses += 1
if source_port in seen_ports:
continue
if len(ports) >= _MAX_PORTS:
too_many_ports = True
continue
seen_ports.add(source_port)
ports.append(source_port)
if too_many_ports:
return _result((), attempts, responses, "ambiguous_response")
if ports:
return _result(tuple(ports), attempts, responses)
if attempts == 0:
return _result((), attempts, responses, "interface_unavailable")
return _result((), attempts, responses, "no_response")
finally:
try:
selector.unregister(active)
except (KeyError, OSError, ValueError):
pass
try:
active.close()
except OSError:
pass
selector.close()
+128
View File
@@ -0,0 +1,128 @@
"""Pure IoTivity manufacturer-certificate OwnerPSK derivation."""
from __future__ import annotations
from collections.abc import Mapping
import hashlib
import hmac
from types import MappingProxyType
from typing import Final
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL: Final = b"x.org.iotivity.conmfgcert"
STANDARD_MFG_CERTIFICATE_OXM_LABEL: Final = b"oic.sec.doxm.mfgcert"
# OpenSSL cipher names mapped to the key-block lengths used by IoTivity's
# CAGenerateOwnerPSK implementation.
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS: Final[Mapping[str, int]] = MappingProxyType(
{
"ECDHE-ECDSA-AES128-SHA256": 96,
"ECDHE-ECDSA-AES128-CCM": 40,
"ECDHE-ECDSA-AES128-CCM8": 40,
"ECDHE-ECDSA-AES128-GCM-SHA256": 120,
"AES256-SHA256": 128,
"ECDHE-ECDSA-AES256-SHA384": 160,
"ECDHE-ECDSA-AES256-GCM-SHA384": 184,
"AES128-GCM-SHA256": 120,
}
)
_TLS_MASTER_SECRET_BYTES: Final = 48
_TLS_RANDOM_BYTES: Final = 32
_OCF_UUID_BYTES: Final = 16
_OWNER_PSK_BYTES: Final = 16
def _require_bytes(name: str, value: bytes, length: int) -> bytes:
if not isinstance(value, bytes):
raise TypeError(f"{name} must be bytes")
if len(value) != length:
raise ValueError(f"{name} must be exactly {length} bytes")
return value
def _require_uuid(name: str, value: bytes) -> bytes:
value = _require_bytes(name, value, _OCF_UUID_BYTES)
if not any(value):
raise ValueError(f"{name} must not be the nil UUID")
return value
def _tls12_p_hash_sha256(
key: bytes,
label: bytes,
random1: bytes,
random2: bytes,
length: int,
) -> bytes:
seed = label + random1 + random2
a_value = hmac.new(key, seed, hashlib.sha256).digest()
output = bytearray()
while len(output) < length:
output.extend(hmac.new(key, a_value + seed, hashlib.sha256).digest())
a_value = hmac.new(key, a_value, hashlib.sha256).digest()
return bytes(output[:length])
def derive_mfg_certificate_owner_psk(
*,
master_secret: bytes,
client_random: bytes,
server_random: bytes,
owner_uuid: bytes,
device_uuid: bytes,
cipher_name: str,
oxm_label: bytes,
) -> bytes:
"""Derive a 128-bit OwnerPSK from caller-supplied DTLS state.
This implements IoTivity's two-stage TLS 1.2 SHA-256 P_hash operation.
It performs no session access, network I/O, ownership writes, or storage.
The caller must supply state from an authenticated manufacturer-certificate
session and explicitly select the OXM label used by that transaction.
IoTivity's other 96-byte ECDH_ANON, ECDHE_PSK, and ECDHE_RSA mappings are
intentionally outside this helper's manufacturer-certificate allowlist.
"""
if not isinstance(cipher_name, str):
raise TypeError("cipher_name must be a string")
key_block_bytes = MFG_CERTIFICATE_KEY_BLOCK_LENGTHS.get(cipher_name)
if key_block_bytes is None:
raise ValueError("unexpected manufacturer-certificate DTLS cipher")
master_secret = _require_bytes(
"master_secret", master_secret, _TLS_MASTER_SECRET_BYTES
)
client_random = _require_bytes(
"client_random", client_random, _TLS_RANDOM_BYTES
)
server_random = _require_bytes(
"server_random", server_random, _TLS_RANDOM_BYTES
)
owner_uuid = _require_uuid("owner_uuid", owner_uuid)
device_uuid = _require_uuid("device_uuid", device_uuid)
if not isinstance(oxm_label, bytes):
raise TypeError("oxm_label must be bytes")
if oxm_label not in {
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
STANDARD_MFG_CERTIFICATE_OXM_LABEL,
}:
raise ValueError("unexpected manufacturer-certificate OXM label")
key_block = _tls12_p_hash_sha256(
master_secret,
b"key expansion",
server_random,
client_random,
key_block_bytes,
)
# IoTivity's OTM callers pass the owner UUID first and the target device
# UUID second. The lower adapter's historical rsrc/prov parameter names
# describe those arguments inconsistently, so preserve the caller order.
return _tls12_p_hash_sha256(
key_block,
oxm_label,
owner_uuid,
device_uuid,
_OWNER_PSK_BYTES,
)
File diff suppressed because it is too large Load Diff
+19 -1
View File
@@ -2,7 +2,8 @@ import pytest
from smartthings_local.errors import MalformedMessageError
from smartthings_local.protocol.coap import (
build_coap, parse_coap, encode_options, block_value, fmt_code,
build_coap, parse_coap, encode_options, block_value, block_fields,
fmt_code,
TYPE_CON, METHOD_GET, URI_PATH, ACCEPT, CF_CBOR, BLOCK2,
)
@@ -43,6 +44,23 @@ def test_block_value_promotes_to_two_bytes_when_num_is_large():
assert len(v) == 2
@pytest.mark.parametrize('num, more, szx', [
(0, 0, 0),
(0, 1, 6),
(2, 1, 6),
(1, 0, 4),
(0xFFF, 0, 0),
(0xFFFFF, 1, 7),
])
def test_block_fields_inverts_block_value(num, more, szx):
assert block_fields(block_value(num, more, szx)) == (num, more, szx)
def test_block_fields_treats_empty_value_as_block_zero():
# RFC 7959 §2.2: a zero-length Block option means num=0, m=0, szx=0.
assert block_fields(b'') == (0, 0, 0)
def test_fmt_code_formats_class_dot_detail():
assert fmt_code(0x45) == '2.05'
assert fmt_code(0x84) == '4.04'
+173
View File
@@ -0,0 +1,173 @@
"""Reader-thread death visibility (QuiteYellow/SmartThings-Local#37).
A connected UDP socket surfaces ICMP errors on recv; before this the
reader exited silently on the first one and every later request waited
out its full timeout against a session nobody was reading. These tests
pin the three behaviours that fixed it: advisory ICMP errnos keep the
reader alive, a real socket error exits with a WARNING and clears
_reader_running, and callers then fail fast with SessionClosedError.
"""
import errno
import logging
import socket
import threading
import time
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionClosedError
from smartthings_local.protocol.dtls_session import DtlsCoapSession
_LOGGER_NAME = "smartthings_local.protocol.dtls_session"
class _NullAuth:
"""Structural AuthenticationProvider — never configured, we skip connect()."""
def configure_context(self, _context):
return None
class _FakeConn:
"""Minimal stand-in for SSL.Connection: each datagram written to the
BIO surfaces as one decrypted packet on the next recv(), then
WantReadError like a drained DTLS record buffer."""
def __init__(self):
self._decrypted = []
def bio_write(self, datagram):
self._decrypted.append(datagram)
def recv(self, _n):
if self._decrypted:
return self._decrypted.pop(0)
raise SSL.WantReadError()
def bio_read(self, _n):
return b""
def send(self, _datagram):
return None
def shutdown(self):
return None
class _FakeSock:
"""Scripted UDP socket. Each step is bytes to return, an exception to
raise, or a callable to run (then a timeout, so the loop re-checks
_stop). An exhausted script blocks like a real recv timeout; once
close()d it raises EBADF the way a closed fd does."""
def __init__(self, steps=()):
self._steps = list(steps)
self.closed = False
self.timeout = None
def settimeout(self, value):
self.timeout = value
def recv(self, _n):
if self.closed:
raise OSError(errno.EBADF, "bad file descriptor")
if self._steps:
step = self._steps.pop(0)
if callable(step):
step()
raise socket.timeout()
if isinstance(step, BaseException):
raise step
return step
time.sleep(0.01)
raise socket.timeout()
def send(self, data):
return len(data)
def close(self):
self.closed = True
def _make_session():
sess = DtlsCoapSession("host", 1234, auth=_NullAuth())
sess.conn = _FakeConn()
return sess
def _run_reader(sess, steps, timeout=2.0):
sess.sock = _FakeSock(steps)
sess.start_reader()
sess._reader_thread.join(timeout)
assert not sess._reader_thread.is_alive(), "reader thread did not exit"
def test_advisory_icmp_error_does_not_kill_reader(caplog):
sess = _make_session()
dispatched = []
sess._dispatch_coap = dispatched.append
with caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME):
_run_reader(sess, [
OSError(errno.ECONNREFUSED, "connection refused"),
b"\x60\x00\x00\x00", # survives, gets dispatched
lambda: sess._stop.set(), # end the loop cleanly
])
assert dispatched == [b"\x60\x00\x00\x00"]
assert not sess._reader_running.is_set()
assert any(r.levelno == logging.DEBUG and "advisory" in r.getMessage()
for r in caplog.records)
# An advisory errno is not a real exit — no WARNING.
assert not any(r.levelno >= logging.WARNING for r in caplog.records)
def test_fatal_socket_error_exits_with_warning(caplog):
sess = _make_session()
with caplog.at_level(logging.WARNING, logger=_LOGGER_NAME):
_run_reader(sess, [OSError(errno.EBADF, "bad file descriptor")])
assert not sess._reader_running.is_set()
warnings = [r for r in caplog.records if r.levelno == logging.WARNING]
assert len(warnings) == 1
assert "reader exiting" in warnings[0].getMessage()
def test_request_fails_fast_after_reader_death():
sess = _make_session()
_run_reader(sess, [OSError(errno.EBADF, "bad file descriptor")])
assert not sess._reader_running.is_set()
start = time.monotonic()
with pytest.raises(SessionClosedError):
sess.get(["oic", "d"], timeout=10.0)
elapsed = time.monotonic() - start
# The whole point: no waiting out the request timeout.
assert elapsed < 1.0, f"get() waited {elapsed:.2f}s instead of failing fast"
def test_close_does_not_log_warning_on_teardown(caplog):
sess = _make_session()
sess.sock = _FakeSock() # empty script: blocks on recv
sess.start_reader()
time.sleep(0.05) # let the reader reach recv
with caplog.at_level(logging.WARNING, logger=_LOGGER_NAME):
sess.close()
sess._reader_thread.join(2.0)
assert not sess._reader_thread.is_alive()
assert not sess._reader_running.is_set()
assert not any(r.levelno >= logging.WARNING for r in caplog.records)
def test_check_live_without_reader_matches_old_conn_guard():
sess = _make_session() # conn set, reader never started
assert sess._reader_thread is None
sess._check_live() # must not raise — config-flow behaviour
sess.conn = None
with pytest.raises(SessionClosedError):
sess._check_live()
+1 -1
View File
@@ -376,7 +376,7 @@ def test_session_uses_connected_socket_send_and_recv(monkeypatch):
assert open_calls == [(('device.example', 5684), {
'family': socket.AF_INET6,
'local_port': None,
'timeout': 2.0,
'timeout': 0.5,
})]
session.close()
+1
View File
@@ -17,6 +17,7 @@ def test_smartthings_local_imports_without_mqtt_demo_present(tmp_path):
import_lines = [
"import smartthings_local.protocol.coap",
"import smartthings_local.protocol.ocf_multicast",
"import smartthings_local.protocol.dtls_session",
"import smartthings_local.ocf.state_cache",
"import smartthings_local.ocf.poll_scheduler",
+508
View File
@@ -0,0 +1,508 @@
"""Blockwise OBSERVE notifications (QuiteYellow/SmartThings-Local#39).
A notification carries only the first block of the representation
(RFC 7959 §2.6). Before this, _dispatch_coap handed that first block
straight to on_notification and the consumer decoded a truncated CBOR
buffer. These tests pin the replacement: a truncated notification is
withheld, the resource is re-read from block 0 under a fresh one-shot
token, and only the reassembled representation reaches the callback.
The re-read starts at block 0 rather than continuing at NUM=1 for two
reasons, both recorded on #39: RFC 7959 §3.4 forbids continuing on the
observation's token, and Samsung's RT-OCF drops a transfer that opens
at NUM>0 under a token it has not seen.
"""
import logging
import socket
import threading
import time
import pytest
from OpenSSL import SSL
from smartthings_local.errors import BlockwiseError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.coap import (
BLOCK2, ETAG, METHOD_GET, OBSERVE, TYPE_ACK, TYPE_NON,
block_fields, block_value, build_coap, parse_coap,
)
from smartthings_local.protocol.dtls_session import DtlsCoapSession
SZX = 6 # 1024-byte blocks, the only size these appliances honour
_LOGGER_NAME = "smartthings_local.protocol.dtls_session"
class _NullAuth:
"""Structural AuthenticationProvider — never configured, we skip connect()."""
def configure_context(self, _context):
return None
class _LoopbackConn:
"""SSL.Connection stand-in that answers requests from a script.
`responder` is called with each parsed request and returns a list of
CoAP datagrams to hand back (possibly empty, to model a silent
device). Responses surface on recv() the way decrypted records do.
"""
def __init__(self, responder):
self._responder = responder
self._inbox = []
self._lock = threading.Lock()
self.sent = []
# -- client -> device
def send(self, datagram):
self.sent.append(parse_coap(datagram))
for reply in self._responder(parse_coap(datagram)):
with self._lock:
self._inbox.append(reply)
return len(datagram)
def bio_read(self, _n):
return b""
# -- device -> client
def inject(self, datagram):
"""Push a device-initiated frame (an OBSERVE notification)."""
with self._lock:
self._inbox.append(datagram)
def bio_write(self, _datagram):
return None
def recv(self, _n):
with self._lock:
if self._inbox:
return self._inbox.pop(0)
raise SSL.WantReadError()
def shutdown(self):
return None
def pending(self):
with self._lock:
return len(self._inbox)
class _PumpSock:
"""UDP socket stand-in. recv() returns a dummy datagram whenever the
connection has something decrypted waiting, so the reader loop keeps
pumping; otherwise it times out like a real socket."""
def __init__(self, conn):
self._conn = conn
self.closed = False
def settimeout(self, _value):
return None
def recv(self, _n):
for _ in range(20):
if self.closed:
raise OSError("closed")
if self._conn.pending():
return b"\x00"
time.sleep(0.005)
raise socket.timeout()
def send(self, data):
return len(data)
def close(self):
self.closed = True
def _make_session(responder, **kwargs):
calls = []
sess = DtlsCoapSession(
"host", 1234, auth=_NullAuth(),
on_notification=lambda href, payload: calls.append((href, payload)),
**kwargs)
sess.conn = _LoopbackConn(responder)
sess.sock = _PumpSock(sess.conn)
sess.start_reader()
return sess, calls
def _notification(tok, payload, *, block2=None, mtype=TYPE_NON, obs=1):
opts = [(OBSERVE, bytes([obs]))]
if block2 is not None:
opts.append((BLOCK2, block2))
return build_coap(mtype, 0x45, 0x1234, tok, opts, payload)
def _content(tok, mid, payload, *, block2=None, etag=None):
opts = []
if etag is not None:
opts.append((ETAG, etag))
if block2 is not None:
opts.append((BLOCK2, block2))
return build_coap(TYPE_ACK, 0x45, mid, tok, opts, payload)
def _requested_block(request):
"""(num, szx) the request asked for, or (0, None) with no Block2."""
_, _, _, _, opts, _ = request
b2 = [v for n, v in opts if n == BLOCK2]
if not b2:
return 0, None
num, _, szx = block_fields(b2[0])
return num, szx
def _wait_for(predicate, timeout=3.0):
deadline = time.time() + timeout
while time.time() < deadline:
if predicate():
return True
time.sleep(0.01)
return False
def _close(sess):
sess.close()
sess.join()
# --------------------------------------------------------------------
# The no-regression case
def test_single_block_notification_is_delivered_inline():
sess, calls = _make_session(lambda request: [])
try:
tok = sess.subscribe(["oven", "vs", "0"])
sess.conn.inject(_notification(tok, b"\xa1\x01\x02"))
assert _wait_for(lambda: calls)
assert calls == [("/oven/vs/0", b"\xa1\x01\x02")]
# The subscribe GET is the only thing we sent — no refetch.
assert len(sess.conn.sent) == 1
finally:
_close(sess)
def test_complete_block_zero_notification_is_delivered_inline():
"""Block2 present but M=0 and NUM=0 means the whole representation
fit in one block. Nothing to fetch back."""
sess, calls = _make_session(lambda request: [])
try:
tok = sess.subscribe(["oven", "vs", "0"])
sess.conn.inject(
_notification(tok, b"\xa1\x01\x02",
block2=block_value(0, 0, SZX)))
assert _wait_for(lambda: calls)
assert calls == [("/oven/vs/0", b"\xa1\x01\x02")]
assert len(sess.conn.sent) == 1
finally:
_close(sess)
# --------------------------------------------------------------------
# The fix
def test_truncated_notification_is_refetched_and_reassembled():
blocks = [b"A" * 1024, b"B" * 40]
def responder(request):
_mtype, code, mid, tok, opts, _ = request
if code != METHOD_GET or any(n == OBSERVE for n, _ in opts):
return [] # the subscribe registration itself
num, _ = _requested_block(request)
more = 1 if num + 1 < len(blocks) else 0
return [_content(tok, mid, blocks[num],
block2=block_value(num, more, SZX))]
sess, calls = _make_session(responder)
try:
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, blocks[0],
block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: calls), "callback never fired"
assert calls == [("/mode/vs/0", b"".join(blocks))]
finally:
_close(sess)
def test_refetch_uses_a_fresh_one_shot_token_not_the_observe_token():
"""RFC 7959 §3.4: the requests for additional blocks cannot use the
token of the observation relationship."""
blocks = [b"A" * 1024, b"B" * 40]
def responder(request):
_mtype, _code, mid, tok, opts, _ = request
if any(n == OBSERVE for n, _ in opts):
return []
num, _ = _requested_block(request)
more = 1 if num + 1 < len(blocks) else 0
return [_content(tok, mid, blocks[num],
block2=block_value(num, more, SZX))]
sess, calls = _make_session(responder)
try:
observe_tok = sess.subscribe(["mode", "vs", "0"])
assert len(observe_tok) == 1, "OBSERVE registrations use 1-byte tokens"
sess.conn.inject(
_notification(observe_tok, blocks[0],
block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: calls)
refetch = [r for r in sess.conn.sent
if not any(n == OBSERVE for n, _ in r[4])]
assert refetch, "no refetch request was sent"
tokens = {r[3] for r in refetch}
assert observe_tok not in tokens
assert all(len(t) == 4 for t in tokens), "one-shot tokens are 4-byte"
assert len(tokens) == 1, "the transfer must hold one token throughout"
# And the transfer restarts at block 0 rather than continuing at 1.
assert _requested_block(refetch[0])[0] == 0
assert [_requested_block(r)[0] for r in refetch] == [0, 1]
# No Observe option on a continuation request.
assert not any(n == OBSERVE for r in refetch for n, _ in r[4])
finally:
_close(sess)
def test_silent_device_drops_the_notification_without_delivering_a_partial():
sess, calls = _make_session(lambda request: [])
try:
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, b"A" * 1024,
block2=block_value(0, 1, SZX)))
# Give the worker a chance to try and fail. _BLOCK_ACK_TIMEOUT is
# 4s per attempt, so we only need to see that nothing partial got
# through in the meantime.
assert not _wait_for(lambda: calls, timeout=0.6)
assert sess._reader_thread.is_alive(), "reader must survive"
finally:
_close(sess)
def test_non_2xx_refetch_is_dropped_rather_than_delivered():
def responder(request):
_mtype, _code, mid, tok, opts, _ = request
if any(n == OBSERVE for n, _ in opts):
return []
return [build_coap(TYPE_ACK, 0x84, mid, tok, [], b"")]
sess, calls = _make_session(responder)
try:
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, b"A" * 1024,
block2=block_value(0, 1, SZX)))
assert not _wait_for(lambda: calls, timeout=0.6)
finally:
_close(sess)
def test_notification_burst_collapses_to_one_refetch_per_resource():
blocks = [b"A" * 1024, b"B" * 40]
gate = threading.Event()
def responder(request):
_mtype, _code, mid, tok, opts, _ = request
if any(n == OBSERVE for n, _ in opts):
return []
gate.wait(2.0) # hold the first transfer open
num, _ = _requested_block(request)
more = 1 if num + 1 < len(blocks) else 0
return [_content(tok, mid, blocks[num],
block2=block_value(num, more, SZX))]
sess, calls = _make_session(responder)
try:
tok = sess.subscribe(["mode", "vs", "0"])
for seq in range(5):
sess.conn.inject(
_notification(tok, blocks[0], obs=seq + 1,
block2=block_value(0, 1, SZX)))
# Five notifications, one queue entry: latest wins per href.
assert _wait_for(lambda: sess._refetch_pending or sess.conn.sent[1:])
assert len(sess._refetch_pending) <= 1
gate.set()
assert _wait_for(lambda: calls)
assert _wait_for(
lambda: not sess._refetch_pending and len(calls) >= 1)
time.sleep(0.2)
# Two transfers at most: the one in flight when the burst landed,
# plus one for the final state.
starts = [r for r in sess.conn.sent
if not any(n == OBSERVE for n, _ in r[4])
and _requested_block(r)[0] == 0]
assert len(starts) <= 2, f"{len(starts)} refetches for one burst"
assert calls[-1] == ("/mode/vs/0", b"".join(blocks))
finally:
gate.set()
_close(sess)
def test_close_during_a_queued_refetch_stops_the_worker():
sess, _calls = _make_session(lambda request: [])
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, b"A" * 1024, block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: sess._refetch_thread is not None)
sess.close()
sess.join() # hangs if the worker outlives the session
assert not sess._refetch_thread.is_alive()
def test_refetch_worker_exits_when_the_reader_dies():
sess, _calls = _make_session(lambda request: [])
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, b"A" * 1024, block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: sess._refetch_thread is not None)
# Kill the reader the way a socket error does, without close().
sess.sock.closed = True
assert _wait_for(lambda: not sess._reader_running.is_set(), timeout=5.0)
sess._refetch_thread.join(6.0)
assert not sess._refetch_thread.is_alive()
sess.close()
def test_debug_bridge_promotes_the_refetch_outcome_to_info(monkeypatch, caplog):
"""The hardware validation for #39 reads this line to confirm which
token the re-read used, so it has to survive the bridge's INFO
default. Without DEBUG_BRIDGE it stays at debug."""
monkeypatch.setattr(dtls_session, "DEBUG_BRIDGE", True)
blocks = [b"A" * 1024, b"B" * 40]
def responder(request):
_mtype, _code, mid, tok, opts, _ = request
if any(n == OBSERVE for n, _ in opts):
return []
num, _ = _requested_block(request)
more = 1 if num + 1 < len(blocks) else 0
return [_content(tok, mid, blocks[num],
block2=block_value(num, more, SZX))]
sess, calls = _make_session(responder)
try:
with caplog.at_level(logging.INFO, logger=_LOGGER_NAME):
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, blocks[0], block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: calls)
line = next((r.getMessage() for r in caplog.records
if r.getMessage().startswith("refetch /mode/vs/0")), None)
assert line is not None, "no refetch line at INFO"
assert "blocks=2" in line
assert f"bytes={sum(len(b) for b in blocks)}" in line
assert line.endswith("ok")
# The token in the line is the one-shot token, not the observe one.
assert f"tok={tok.hex()} " not in line
finally:
_close(sess)
# --------------------------------------------------------------------
# Shared Block2 loop hardening
def test_stale_block_number_is_not_concatenated():
"""A retransmit of block 0 arriving while we wait for block 1 must
not be appended as if it were block 1."""
served = []
def responder(request):
_mtype, _code, mid, tok, _opts, _ = request
num, _ = _requested_block(request)
served.append(num)
if num == 0:
return [_content(tok, mid, b"A" * 1024,
block2=block_value(0, 1, SZX))]
# Answer the block-1 request with a duplicate of block 0 first.
return [
_content(tok, mid, b"A" * 1024, block2=block_value(0, 1, SZX)),
_content(tok, mid, b"B" * 40, block2=block_value(1, 0, SZX)),
]
sess, _calls = _make_session(responder)
try:
code, payload = sess.get(["mode", "vs", "0"], timeout=5.0)
assert code == 0x45
assert payload == b"A" * 1024 + b"B" * 40
finally:
_close(sess)
def test_etag_change_mid_transfer_restarts_then_fails():
etags = [b"\x01", b"\x02", b"\x03", b"\x04"]
def responder(request):
_mtype, _code, mid, tok, _opts, _ = request
num, _ = _requested_block(request)
# A different ETag on every single response: the representation
# never settles, so reassembly can never be consistent.
etag = etags.pop(0) if etags else b"\xff"
payload = b"A" * 1024 if num == 0 else b"B" * 40
more = 1 if num == 0 else 0
return [_content(tok, mid, payload, etag=etag,
block2=block_value(num, more, SZX))]
sess, _calls = _make_session(responder)
try:
with pytest.raises(BlockwiseError):
sess.get(["mode", "vs", "0"], timeout=5.0)
finally:
_close(sess)
def test_stable_etag_across_blocks_reassembles():
def responder(request):
_mtype, _code, mid, tok, _opts, _ = request
num, _ = _requested_block(request)
payload = b"A" * 1024 if num == 0 else b"B" * 40
more = 1 if num == 0 else 0
return [_content(tok, mid, payload, etag=b"\x77",
block2=block_value(num, more, SZX))]
sess, _calls = _make_session(responder)
try:
code, payload = sess.get(["mode", "vs", "0"], timeout=5.0)
assert code == 0x45
assert payload == b"A" * 1024 + b"B" * 40
finally:
_close(sess)
def test_szx_downshift_asks_for_the_block_after_what_we_have():
"""Server answers block 0 at SZX=6 (1024B) then drops to SZX=4
(256B). Block numbers index the new size, so the next request is
block 4, not block 1."""
requested = []
def responder(request):
_mtype, _code, mid, tok, _opts, _ = request
num, szx = _requested_block(request)
requested.append((num, szx))
if num == 0:
return [_content(tok, mid, b"A" * 1024,
block2=block_value(0, 1, 4))]
return [_content(tok, mid, b"B" * 100,
block2=block_value(num, 0, 4))]
sess, _calls = _make_session(responder)
try:
code, payload = sess.get(["mode", "vs", "0"], timeout=5.0)
assert code == 0x45
assert payload == b"A" * 1024 + b"B" * 100
# 1024 bytes in hand at 256B blocks = blocks 0..3 done, ask for 4.
assert requested[1] == (4, 4)
finally:
_close(sess)
+392
View File
@@ -0,0 +1,392 @@
"""Known-host OCF multicast responder-port discovery tests."""
from __future__ import annotations
from types import SimpleNamespace
import pytest
from smartthings_local.protocol.coap import (
ACCEPT,
TYPE_CON,
TYPE_NON,
URI_PATH,
URI_QUERY,
build_coap,
parse_coap,
)
from smartthings_local.protocol.ocf_multicast import (
_OCF_MULTICAST_GROUP,
OcfResponderPortDiscoveryResult,
discover_ocf_responder_ports,
)
_TARGET = "192.0.2.20"
_OTHER = "192.0.2.21"
_INTERFACE = "192.0.2.10"
class _FakeSocket:
def __init__(self, responder=None):
self.responder = responder
self.incoming = []
self.sent = []
self.options = []
self.bound = None
self.blocking = None
self.closed = False
def setsockopt(self, level, option, value):
self.options.append((level, option, value))
def bind(self, address):
self.bound = address
def setblocking(self, enabled):
self.blocking = enabled
def sendto(self, datagram, destination):
self.sent.append((datagram, destination))
if self.responder is not None:
self.incoming.extend(self.responder(datagram, destination))
return len(datagram)
def recvfrom(self, _size):
return self.incoming.pop(0)
def close(self):
self.closed = True
class _FakeSelector:
def __init__(self, active):
self.active = active
self.closed = False
def register(self, *_args):
return None
def unregister(self, *_args):
return None
def select(self, _timeout):
if not self.active.incoming:
return []
return [(SimpleNamespace(fileobj=self.active), selectors_event_read())]
def close(self):
self.closed = True
def selectors_event_read():
return 1
def _response(
datagram, _destination=None, *, host=_TARGET, port=43123, message_type=TYPE_NON
):
_mtype, _code, mid, token, _options, _payload = parse_coap(datagram)
response_mid = mid if message_type != TYPE_CON else (mid + 1) & 0xFFFF
return [
(
build_coap(message_type, 0x45, response_mid, token, [], b"directory"),
(host, port),
)
]
@pytest.fixture
def patch_socket(monkeypatch):
created = []
def install(responder=None):
active = _FakeSocket(responder)
selector = _FakeSelector(active)
created.append((active, selector))
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.socket.socket",
lambda *_args: active,
)
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.selectors.DefaultSelector",
lambda: selector,
)
return active, selector
return install
def test_sends_proven_directory_requests_and_filtered_fallback_on_interface(
patch_socket,
):
active, selector = patch_socket(_response)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.ports == (43123,)
assert result.attempts == 3
assert result.responses == 3
assert active.bound == (_INTERFACE, 0)
assert active.blocking is False
assert active.closed
assert selector.closed
assert all(
destination == (_OCF_MULTICAST_GROUP, 5683) for _, destination in active.sent
)
requests = [parse_coap(datagram) for datagram, _ in active.sent]
assert all(
[value for number, value in item[4] if number == URI_PATH] == [b"oic", b"res"]
for item in requests
)
queries = [
[value for number, value in item[4] if number == URI_QUERY] for item in requests
]
assert queries == [[], [], [b"rt=oic.r.doxm"]]
accepts = [
[value for number, value in item[4] if number == ACCEPT] for item in requests
]
assert accepts == [[b"\x27\x10"], [b"\x3c"], [b"\x3c"]]
assert [value for number, value in requests[0][4] if number == 2049] == [
b"\x08\x00"
]
assert [value for number, value in requests[1][4] if number == 2049] == []
assert [value for number, value in requests[2][4] if number == 2049] == []
@pytest.mark.parametrize(
("accepted_accept", "accepted_query"),
(
(b"\x27\x10", ()),
(b"\x3c", ()),
(b"\x3c", (b"rt=oic.r.doxm",)),
),
)
def test_each_directory_request_profile_can_find_the_responder(
patch_socket, accepted_accept, accepted_query
):
def responder(datagram, destination):
_mtype, _code, _mid, _token, options, _payload = parse_coap(datagram)
accept = next(value for number, value in options if number == ACCEPT)
query = tuple(value for number, value in options if number == URI_QUERY)
if (accept, query) != (accepted_accept, accepted_query):
return []
return _response(datagram, destination)
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.ports == (43123,)
assert result.responses == 1
def test_accepts_only_token_correlated_content_from_the_target(patch_socket):
def responder(datagram, _destination):
responses = _response(datagram)
_mtype, _code, mid, token, _options, _payload = parse_coap(datagram)
responses.extend(
[
(build_coap(TYPE_NON, 0x45, mid, b"wrong", [], b"x"), (_TARGET, 49999)),
(build_coap(TYPE_NON, 0x44, mid, token, [], b"x"), (_TARGET, 49998)),
(build_coap(TYPE_NON, 0x45, mid, token, [], b"x"), (_OTHER, 49997)),
(build_coap(TYPE_NON, 0x45, mid, token, [], b""), (_TARGET, 49996)),
]
)
return responses
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.ports == (43123,)
assert result.responses == 3
def test_ignores_malformed_and_oversized_datagrams(patch_socket):
def responder(datagram, _destination):
valid = _response(datagram)
invalid_version = bytes([valid[0][0][0] & 0x3F]) + valid[0][0][1:]
invalid_token_length = bytes([0x49]) + valid[0][0][1:]
return [
(b"\x40", (_TARGET, 49999)),
(invalid_version, (_TARGET, 49998)),
(invalid_token_length, (_TARGET, 49997)),
(b"x" * 8193, (_TARGET, 49996)),
*valid,
]
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.ports == (43123,)
assert result.responses == 3
def test_acknowledges_confirmable_responses(patch_socket):
active, _selector = patch_socket(
lambda datagram, _destination: _response(datagram, message_type=TYPE_CON)
)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.found
acknowledgements = [
(parse_coap(datagram), destination)
for datagram, destination in active.sent
if parse_coap(datagram)[0] == 2
]
assert len(acknowledgements) == 3
assert all(item[0][1] == 0 and item[0][3] == b"" for item in acknowledgements)
assert all(
destination == (_TARGET, 43123) for _item, destination in acknowledgements
)
def test_collects_a_small_bounded_candidate_set(patch_socket):
response_number = 0
def responder(datagram, _destination):
nonlocal response_number
port = 40000 + response_number % 8
response_number += 1
return _response(datagram, port=port)
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=4,
)
assert result.ports == tuple(range(40000, 40008))
assert result.error_code is None
def test_fails_closed_when_too_many_distinct_ports_answer(patch_socket):
next_port = 40000
def responder(datagram, _destination):
nonlocal next_port
responses = [
*_response(datagram, port=next_port),
*_response(datagram, port=next_port + 1),
]
next_port += 2
return responses
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=4,
)
assert result.ports == ()
assert result.responses == 24
assert result.error_code == "ambiguous_response"
def test_no_response_and_interface_failures_are_fixed_results(
patch_socket, monkeypatch
):
active, _selector = patch_socket()
no_response = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert no_response == OcfResponderPortDiscoveryResult(
ports=(),
attempts=3,
responses=0,
error_code="no_response",
)
assert active.closed
class BrokenSocket:
def setsockopt(self, *_args):
raise OSError("synthetic")
def close(self):
return None
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.socket.socket",
lambda *_args: BrokenSocket(),
)
unavailable = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
)
assert unavailable.error_code == "interface_unavailable"
assert unavailable.attempts == 0
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.selectors.DefaultSelector",
lambda: (_ for _ in ()).throw(OSError("synthetic")),
)
unavailable = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
)
assert unavailable.error_code == "interface_unavailable"
assert unavailable.attempts == 0
@pytest.mark.parametrize(
("kwargs", "exception"),
[
({"target_address": 123}, TypeError),
({"target_address": "not-an-address"}, ValueError),
({"target_address": _OCF_MULTICAST_GROUP}, ValueError),
({"interface_address": "0.0.0.0"}, ValueError),
({"discovery_port": True}, TypeError),
({"discovery_port": 0}, ValueError),
({"timeout": float("nan")}, ValueError),
({"rounds": 0}, ValueError),
({"rounds": 5}, ValueError),
],
)
def test_rejects_invalid_options_without_opening_a_socket(
monkeypatch, kwargs, exception
):
values = {
"target_address": _TARGET,
"interface_address": _INTERFACE,
**kwargs,
}
socket_factory = pytest.fail
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.socket.socket",
socket_factory,
)
with pytest.raises(exception):
discover_ocf_responder_ports(**values)
def test_result_repr_omits_ports_and_addresses():
result = OcfResponderPortDiscoveryResult(ports=(43123,), attempts=2, responses=1)
rendered = repr(result)
assert "43123" not in rendered
assert _TARGET not in rendered
assert "port_count=1" in rendered
+86
View File
@@ -0,0 +1,86 @@
"""Oven descriptor flatten() contracts for the HA Number entity's range.
The oven reports ``x.com.samsung.da.desired = 0`` whenever no cycle is
set. That is "no setpoint", not a 0 °C target, and publishing it as one
makes Home Assistant reject every state message against the Number
entity's declared 30-270 range.
"""
from __future__ import annotations
import pytest
from mqtt_demo.samples import oven
def _links(desired, current=180):
"""A /temperatures/vs/0 link tree carrying one desired/current pair."""
return {
'/temperatures/vs/0': {
'x.com.samsung.da.items': [{
'x.com.samsung.da.current': str(current),
'x.com.samsung.da.desired': str(desired),
}],
},
}
@pytest.mark.parametrize('desired', [
oven.SETPOINT_MIN_C,
oven.SETPOINT_MIN_C + oven.SETPOINT_STEP_C,
180,
oven.SETPOINT_MAX_C,
])
def test_settable_setpoints_are_published_unchanged(desired):
assert oven.flatten(_links(desired))['target_temp_c'] == desired
@pytest.mark.parametrize('desired', [
0, # the idle oven; see module docstring
oven.SETPOINT_MIN_C - 1,
oven.SETPOINT_MAX_C + 1,
])
def test_unsettable_setpoints_are_published_as_absent(desired):
assert oven.flatten(_links(desired))['target_temp_c'] is None
def test_out_of_range_setpoint_does_not_suppress_current_temperature():
"""The guard applies to the setpoint alone. A cooling oven still
reports its cavity temperature after the cycle ends."""
sensors = oven.flatten(_links(0, current=210))
assert sensors['target_temp_c'] is None
assert sensors['current_temp_c'] == 210
def test_missing_temperature_resource_leaves_both_absent():
sensors = oven.flatten({})
assert sensors['target_temp_c'] is None
assert sensors['current_temp_c'] is None
def test_every_committed_write_is_a_value_flatten_will_publish():
"""The write path snaps to the step grid *before* bounds-checking, so
it accepts more than flatten() publishes: 29 commits as 30, and 271 as
270. That is fine for a slider, but it means the two range checks are
not symmetric. What has to hold is the weaker invariant: any setpoint
the oven is actually told to adopt is one flatten() will show back,
otherwise a write appears to succeed and then reads as unknown."""
handler = oven.command_handlers()[oven.CMD_SETPOINT]
for requested in range(-20, oven.SETPOINT_MAX_C + 40):
write = handler(str(requested), _links(180))
if write is None:
continue
_path, body = write
committed = int(body['x.com.samsung.da.items'][0][
'x.com.samsung.da.desired'])
assert oven.flatten(_links(committed))['target_temp_c'] == committed
def test_zero_is_rejected_on_the_write_path_too():
"""0 is the one value that neither snaps into range nor publishes."""
handler = oven.command_handlers()[oven.CMD_SETPOINT]
assert handler('0', _links(180)) is None
+144
View File
@@ -0,0 +1,144 @@
"""IoTivity manufacturer-certificate OwnerPSK derivation contracts."""
from __future__ import annotations
import pytest
from smartthings_local.protocol.owner_psk import (
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS,
STANDARD_MFG_CERTIFICATE_OXM_LABEL,
derive_mfg_certificate_owner_psk,
)
_VALID_INPUTS = {
"master_secret": bytes(range(48)),
"client_random": bytes(range(32)),
"server_random": bytes(range(32, 64)),
"owner_uuid": bytes.fromhex("00112233445566778899aabbccddeeff"),
"device_uuid": bytes.fromhex("ffeeddccbbaa99887766554433221100"),
"cipher_name": "ECDHE-ECDSA-AES128-GCM-SHA256",
"oxm_label": CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
}
# These synthetic expected values were generated from the fixed inputs above;
# they were not captured from IoTivity or a device. They lock deterministic
# output for the selected mappings, while the table test below states the
# key-block-length contract explicitly.
def test_fixed_synthetic_gcm_regression_vector():
assert derive_mfg_certificate_owner_psk(**_VALID_INPUTS).hex() == (
"ccd6c618a91290dee8c106544ed79a33"
)
def test_owner_then_device_uuid_order_matches_iotivity_callers():
reversed_context = derive_mfg_certificate_owner_psk(
**{
**_VALID_INPUTS,
"owner_uuid": _VALID_INPUTS["device_uuid"],
"device_uuid": _VALID_INPUTS["owner_uuid"],
}
)
assert reversed_context.hex() == "8f0f2416c483546dc1806db769b21b68"
assert reversed_context != derive_mfg_certificate_owner_psk(**_VALID_INPUTS)
def test_fixed_synthetic_ccm8_regression_vector():
inputs = {
**_VALID_INPUTS,
"cipher_name": "ECDHE-ECDSA-AES128-CCM8",
}
assert derive_mfg_certificate_owner_psk(**inputs).hex() == (
"ddd3d945e266ee3dc27ff3a2c4321d32"
)
def test_standard_and_confirmed_labels_derive_distinct_keys():
confirmed = derive_mfg_certificate_owner_psk(**_VALID_INPUTS)
standard = derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "oxm_label": STANDARD_MFG_CERTIFICATE_OXM_LABEL}
)
assert standard.hex() == "26ee1fe4c3e74509a2f5db5ab41b1e47"
assert standard != confirmed
def test_iotivity_cipher_key_block_lengths_are_immutable():
assert dict(MFG_CERTIFICATE_KEY_BLOCK_LENGTHS) == {
"ECDHE-ECDSA-AES128-SHA256": 96,
"ECDHE-ECDSA-AES128-CCM": 40,
"ECDHE-ECDSA-AES128-CCM8": 40,
"ECDHE-ECDSA-AES128-GCM-SHA256": 120,
"AES256-SHA256": 128,
"ECDHE-ECDSA-AES256-SHA384": 160,
"ECDHE-ECDSA-AES256-GCM-SHA384": 184,
"AES128-GCM-SHA256": 120,
}
with pytest.raises(TypeError):
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS["new-cipher"] = 1
@pytest.mark.parametrize(
("field", "length"),
[
("master_secret", 48),
("client_random", 32),
("server_random", 32),
("owner_uuid", 16),
("device_uuid", 16),
],
)
def test_binary_inputs_require_exact_bytes_and_lengths(field, length):
with pytest.raises(TypeError, match=f"{field} must be bytes"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, field: bytearray(length)}
)
for invalid_length in (length - 1, length + 1):
with pytest.raises(ValueError, match=f"exactly {length} bytes"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, field: b"x" * invalid_length}
)
@pytest.mark.parametrize("field", ["owner_uuid", "device_uuid"])
def test_nil_uuid_is_rejected(field):
with pytest.raises(ValueError, match="must not be the nil UUID"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, field: bytes(16)}
)
def test_cipher_and_label_must_be_explicit_supported_values():
with pytest.raises(TypeError, match="cipher_name must be a string"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "cipher_name": b"cipher"}
)
with pytest.raises(ValueError, match="unexpected.*cipher"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "cipher_name": "ECDHE-RSA-AES128-GCM-SHA256"}
)
with pytest.raises(TypeError, match="oxm_label must be bytes"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "oxm_label": "oic.sec.doxm.mfgcert"}
)
with pytest.raises(ValueError, match="unexpected.*label"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "oxm_label": b"unsupported"}
)
def test_failures_do_not_include_key_material():
key_material = b"private-master-secret"
with pytest.raises(ValueError) as raised:
derive_mfg_certificate_owner_psk(
**{
**_VALID_INPUTS,
"master_secret": key_material,
"cipher_name": "unsupported",
}
)
assert key_material.hex() not in str(raised.value)
assert "private-master-secret" not in str(raised.value)
+380
View File
@@ -0,0 +1,380 @@
from __future__ import annotations
import gc
import traceback
import weakref
from concurrent.futures import ThreadPoolExecutor
from dataclasses import asdict
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionError
from smartthings_local.protocol import auth as auth_module
from smartthings_local.protocol import dtls_session as session_module
from smartthings_local.protocol.auth import PskAuth
from smartthings_local.protocol.dtls_session import DtlsCoapSession
_IDENTITY = b"i" * 16
_KEY = b"k" * 16
_OTHER_IDENTITY = b"j" * 16
_OTHER_KEY = b"l" * 32
class _BytesSubclass(bytes):
pass
def _fake_openssl_util(setter):
return SimpleNamespace(
ffi=auth_module._util.ffi,
lib=SimpleNamespace(SSL_CTX_set_psk_client_callback=setter),
)
def _invoke_callback(callback, identity_size: int, key_size: int):
ffi = auth_module._util.ffi
identity_buffer = ffi.new("char[]", max(identity_size, 1))
key_buffer = ffi.new("unsigned char[]", max(key_size, 1))
copied = callback(
ffi.NULL,
ffi.NULL,
identity_buffer,
identity_size,
key_buffer,
key_size,
)
return (
copied,
bytes(ffi.buffer(identity_buffer, max(identity_size, 1))),
bytes(ffi.buffer(key_buffer, max(key_size, 1))),
)
@pytest.mark.parametrize("key_length", [16, 32])
def test_psk_auth_accepts_exact_supported_credential_lengths(key_length):
provider = PskAuth(identity=_IDENTITY, key=b"k" * key_length)
assert repr(provider) == "PskAuth()"
@pytest.mark.parametrize(
("identity", "key"),
[
("i" * 16, _KEY),
(bytearray(_IDENTITY), _KEY),
(memoryview(_IDENTITY), _KEY),
(_BytesSubclass(_IDENTITY), _KEY),
(_IDENTITY, "k" * 16),
(_IDENTITY, bytearray(_KEY)),
(_IDENTITY, memoryview(_KEY)),
(_IDENTITY, _BytesSubclass(_KEY)),
],
)
def test_psk_auth_rejects_non_bytes_credentials(identity, key):
with pytest.raises(TypeError, match="identity and key must be bytes"):
PskAuth(identity=identity, key=key)
@pytest.mark.parametrize("identity_length", [0, 15, 17])
def test_psk_auth_rejects_invalid_identity_lengths(identity_length):
with pytest.raises(ValueError, match="raw 16-byte OCF UUID"):
PskAuth(identity=b"i" * identity_length, key=_KEY)
def test_psk_auth_rejects_identity_with_nul_byte():
with pytest.raises(ValueError, match="cannot contain a NUL"):
PskAuth(identity=b"i" * 15 + b"\x00", key=_KEY)
@pytest.mark.parametrize("key_length", [0, 15, 17, 31, 33])
def test_psk_auth_rejects_invalid_key_lengths(key_length):
with pytest.raises(ValueError, match="16 or 32 bytes"):
PskAuth(identity=_IDENTITY, key=b"k" * key_length)
def test_psk_auth_is_immutable_and_has_no_public_credential_surface():
provider = PskAuth(identity=_IDENTITY, key=_KEY)
rendered = repr(provider)
assert rendered == "PskAuth()"
assert str(provider) == rendered
assert _IDENTITY.decode() not in rendered
assert _KEY.decode() not in rendered
assert not hasattr(provider, "identity")
assert not hasattr(provider, "key")
assert not hasattr(provider, "_identity")
assert not hasattr(provider, "_key")
with pytest.raises(TypeError):
vars(provider)
with pytest.raises(TypeError):
asdict(provider)
with pytest.raises(AttributeError, match="immutable"):
provider.identity = _OTHER_IDENTITY
with pytest.raises(AttributeError, match="immutable"):
del provider._callback
def test_psk_auth_identity_equality_does_not_compare_credentials():
first = PskAuth(identity=_IDENTITY, key=_KEY)
second = PskAuth(identity=_IDENTITY, key=_KEY)
assert first != second
assert len({first, second}) == 2
def test_psk_callback_copies_exact_identity_and_key():
installed = {}
def setter(context_handle, callback):
installed["context"] = context_handle
installed["callback"] = callback
context_handle = object()
context = MagicMock()
context._context = context_handle
provider = PskAuth(identity=_IDENTITY, key=_KEY)
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
provider.configure_context(context)
assert installed["context"] is context_handle
callback = installed["callback"]
copied, identity_bytes, key_bytes = _invoke_callback(callback, 17, 16)
assert copied == 16
assert identity_bytes == _IDENTITY + b"\x00"
assert key_bytes == _KEY
context.set_cipher_list.assert_called_once_with(
b"ECDHE-PSK-AES128-CBC-SHA256:@SECLEVEL=0"
)
context.load_verify_locations.assert_not_called()
context.set_verify.assert_not_called()
@pytest.mark.parametrize(
("identity_size", "key_size"),
[(16, 16), (17, 15)],
)
def test_psk_callback_rejects_short_buffers_without_partial_copy(
identity_size,
key_size,
):
installed = {}
provider = PskAuth(identity=_IDENTITY, key=_KEY)
context = MagicMock()
context._context = object()
with patch.object(
auth_module,
"_util",
_fake_openssl_util(
lambda _context, callback: installed.setdefault(
"callback", callback
)
),
):
provider.configure_context(context)
ffi = auth_module._util.ffi
identity_buffer = ffi.new("char[]", 17)
key_buffer = ffi.new("unsigned char[]", 16)
ffi.memmove(identity_buffer, b"I" * 17, 17)
ffi.memmove(key_buffer, b"K" * 16, 16)
copied = installed["callback"](
ffi.NULL,
ffi.NULL,
identity_buffer,
identity_size,
key_buffer,
key_size,
)
assert copied == 0
assert bytes(ffi.buffer(identity_buffer, 17)) == b"I" * 17
assert bytes(ffi.buffer(key_buffer, 16)) == b"K" * 16
@pytest.mark.parametrize("null_buffer", ["identity", "key"])
def test_psk_callback_rejects_null_buffers(null_buffer):
installed = {}
provider = PskAuth(identity=_IDENTITY, key=_KEY)
context = MagicMock()
context._context = object()
with patch.object(
auth_module,
"_util",
_fake_openssl_util(
lambda _context, callback: installed.setdefault(
"callback", callback
)
),
):
provider.configure_context(context)
ffi = auth_module._util.ffi
identity_buffer = ffi.new("char[]", 17)
key_buffer = ffi.new("unsigned char[]", 16)
ffi.memmove(identity_buffer, b"I" * 17, 17)
ffi.memmove(key_buffer, b"K" * 16, 16)
if null_buffer == "identity":
identity_buffer = ffi.NULL
else:
key_buffer = ffi.NULL
copied = installed["callback"](
ffi.NULL,
ffi.NULL,
identity_buffer,
17,
key_buffer,
16,
)
assert copied == 0
if null_buffer == "identity":
assert bytes(ffi.buffer(key_buffer, 16)) == b"K" * 16
else:
assert bytes(ffi.buffer(identity_buffer, 17)) == b"I" * 17
def test_psk_auth_unsupported_binding_error_contains_no_credentials():
provider = PskAuth(identity=_IDENTITY, key=_KEY)
context = MagicMock()
context._context = object()
unsupported_util = SimpleNamespace(
ffi=auth_module._util.ffi,
lib=SimpleNamespace(),
)
with (
patch.object(auth_module, "_util", unsupported_util),
pytest.raises(RuntimeError) as captured,
):
provider.configure_context(context)
rendered = (
str(captured.value)
+ repr(captured.value)
+ "".join(traceback.format_exception(captured.value))
)
assert _IDENTITY.decode() not in rendered
assert _KEY.decode() not in rendered
context.set_cipher_list.assert_not_called()
def test_psk_auth_configures_real_openssl_context():
context = SSL.Context(SSL.DTLS_METHOD)
provider = PskAuth(identity=_IDENTITY, key=_KEY)
assert provider.configure_context(context) is None
def test_distinct_psk_providers_do_not_share_callback_credentials():
callbacks = []
def setter(_context, callback):
callbacks.append(callback)
first = PskAuth(identity=_IDENTITY, key=_KEY)
second = PskAuth(identity=_OTHER_IDENTITY, key=_OTHER_KEY)
first_context = MagicMock()
first_context._context = object()
second_context = MagicMock()
second_context._context = object()
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
first.configure_context(first_context)
second.configure_context(second_context)
assert callbacks[0] is not callbacks[1]
with ThreadPoolExecutor(max_workers=2) as executor:
first_future = executor.submit(_invoke_callback, callbacks[0], 17, 16)
second_future = executor.submit(
_invoke_callback,
callbacks[1],
17,
32,
)
first_result = first_future.result()
second_result = second_future.result()
assert first_result == (16, _IDENTITY + b"\x00", _KEY)
assert second_result == (32, _OTHER_IDENTITY + b"\x00", _OTHER_KEY)
def test_session_retains_psk_callback_only_with_provider_lifetime():
callback_reference = None
def setter(_context, callback):
nonlocal callback_reference
callback_reference = weakref.ref(callback)
provider = PskAuth(identity=_IDENTITY, key=_KEY)
session = DtlsCoapSession(
"appliance.invalid",
49154,
auth=provider,
)
context = MagicMock()
context._context = object()
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
session.auth.configure_context(context)
del provider
gc.collect()
assert callback_reference is not None
assert callback_reference() is not None
assert _invoke_callback(callback_reference(), 17, 16) == (
16,
_IDENTITY + b"\x00",
_KEY,
)
del session
gc.collect()
assert callback_reference() is None
def test_session_accepts_psk_provider_without_legacy_certificate_material():
provider = PskAuth(identity=_IDENTITY, key=_KEY)
session = DtlsCoapSession("appliance.invalid", 49154, auth=provider)
assert session.auth is provider
assert session.cert_path is None
assert session.key_path is None
assert session.cert_pem is None
assert session.key_pem is None
def test_psk_handshake_rejection_does_not_expose_credentials():
provider = PskAuth(identity=_IDENTITY, key=_KEY)
session = DtlsCoapSession("appliance.invalid", 49154, auth=provider)
context = MagicMock()
context._context = object()
connection = MagicMock()
connection.do_handshake.side_effect = SSL.Error()
udp_socket = MagicMock()
endpoint = SimpleNamespace(sockaddr=("192.0.2.100", 49154))
with (
patch.object(auth_module, "_util", _fake_openssl_util(lambda *_: None)),
patch.object(session_module.SSL, "Context", return_value=context),
patch.object(session_module.SSL, "Connection", return_value=connection),
patch.object(
session_module,
"open_connected_udp_socket",
return_value=(udp_socket, endpoint),
),
pytest.raises(SessionError) as captured,
):
session.connect()
rendered = (
str(captured.value)
+ repr(captured.value)
+ "".join(traceback.format_exception(captured.value))
)
assert _IDENTITY.decode() not in rendered
assert _KEY.decode() not in rendered
udp_socket.close.assert_called_once_with()
+130 -2
View File
@@ -6,8 +6,23 @@ import inspect
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
from smartthings_local.ocf.state_cache import StateCache
from smartthings_local.protocol.auth import AuthenticationProvider, CertificateAuth
from smartthings_local.protocol.dtls_session import DtlsCoapSession
from smartthings_local.protocol.auth import (
AuthenticationProvider,
CertificateAuth,
PskAuth,
SamsungServerProfile,
SamsungServerRole,
ServerCertificateAuth,
)
from smartthings_local.protocol.dtls_session import (
ConnectCancellation,
DtlsCoapSession,
)
from smartthings_local.protocol.ocf_multicast import (
OcfResponderPortDiscoveryResult,
discover_ocf_responder_ports,
)
from smartthings_local.protocol.owner_psk import derive_mfg_certificate_owner_psk
def _assert_compatible_signature(callable_object, expected: list[str]) -> None:
@@ -46,10 +61,109 @@ def test_dtls_session_constructor_keeps_file_memory_and_local_port_inputs():
assert auth_parameter.default is None
def test_known_host_multicast_discovery_has_a_bounded_explicit_interface_api():
parameters = inspect.signature(discover_ocf_responder_ports).parameters
assert list(parameters) == [
"target_address",
"interface_address",
"discovery_port",
"timeout",
"rounds",
]
assert parameters["target_address"].default is inspect.Parameter.empty
for name in ("interface_address", "discovery_port", "timeout", "rounds"):
assert parameters[name].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["interface_address"].default is inspect.Parameter.empty
assert parameters["discovery_port"].default == 5683
assert parameters["timeout"].default == 3.0
assert parameters["rounds"].default == 2
result = OcfResponderPortDiscoveryResult(
ports=(43123,),
attempts=2,
responses=1,
)
assert result.found is True
assert result.ports == (43123,)
def test_certificate_auth_is_a_public_authentication_provider():
provider = CertificateAuth.from_files("/synthetic/cert.pem", "/synthetic/key")
assert isinstance(provider, AuthenticationProvider)
for factory in (CertificateAuth.from_files, CertificateAuth.from_memory):
profile_parameter = inspect.signature(factory).parameters["server_profile"]
assert profile_parameter.kind is inspect.Parameter.KEYWORD_ONLY
assert profile_parameter.default is None
def test_samsung_server_profile_is_public_and_explicitly_bound():
parameters = inspect.signature(SamsungServerProfile.bound_device).parameters
assert list(parameters) == [
"expected_certificate_identity",
"role",
"additional_ca_pem",
]
assert (
parameters["expected_certificate_identity"].default
is inspect.Parameter.empty
)
assert parameters["role"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["role"].default is SamsungServerRole.HOME_APPLIANCE
assert parameters["additional_ca_pem"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["additional_ca_pem"].default is None
def test_server_certificate_auth_is_a_public_authentication_provider():
profile = SamsungServerProfile.bound_device(
"abababab-abab-abab-abab-abababababab",
role=SamsungServerRole.VD_DEVICE,
)
provider = ServerCertificateAuth(server_profile=profile)
assert isinstance(provider, AuthenticationProvider)
session = DtlsCoapSession("device.example", 5684, auth=provider)
assert session.auth is provider
assert session.cert_path is None
assert session.key_path is None
assert session.cert_pem is None
assert session.key_pem is None
parameters = inspect.signature(ServerCertificateAuth).parameters
assert list(parameters) == ["server_profile"]
assert parameters["server_profile"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["server_profile"].default is inspect.Parameter.empty
def test_psk_auth_is_a_public_authentication_provider():
provider = PskAuth(identity=b"i" * 16, key=b"k" * 16)
assert isinstance(provider, AuthenticationProvider)
parameters = inspect.signature(PskAuth).parameters
assert list(parameters) == ["identity", "key"]
assert all(
parameter.kind is inspect.Parameter.KEYWORD_ONLY
and parameter.default is inspect.Parameter.empty
for parameter in parameters.values()
)
def test_owner_psk_derivation_keeps_every_security_input_explicit():
parameters = inspect.signature(
derive_mfg_certificate_owner_psk
).parameters
assert list(parameters) == [
"master_secret",
"client_random",
"server_random",
"owner_uuid",
"device_uuid",
"cipher_name",
"oxm_label",
]
assert all(
parameter.kind is inspect.Parameter.KEYWORD_ONLY
and parameter.default is inspect.Parameter.empty
for parameter in parameters.values()
)
def test_dtls_session_keeps_current_consumer_methods():
expected = {
@@ -65,6 +179,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,
[
+169
View File
@@ -0,0 +1,169 @@
"""Session-owned pacing for request sends."""
from __future__ import annotations
from unittest.mock import Mock
import pytest
from smartthings_local.errors import SessionClosedError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.coap import (
METHOD_GET,
METHOD_POST,
OBSERVE,
TYPE_ACK,
TYPE_CON,
build_coap,
parse_coap,
)
from smartthings_local.protocol.dtls_session import DtlsCoapSession
class _NullAuth:
def configure_context(self, _context):
return None
def _session():
session = DtlsCoapSession(
"device.example",
5684,
auth=_NullAuth(),
rate_limit_rps=1_000_000,
)
session.conn = object()
return session
def test_first_get_post_and_subscribe_are_paced_before_send():
session = _session()
order = []
requests = []
def pace():
order.append("pace")
def send(datagram):
order.append("send")
request = parse_coap(datagram)
requests.append(request)
_mtype, _code, mid, token, options, _payload = request
if any(number == OBSERVE for number, _value in options):
assert session._observe_tokens[token] == "/mode/vs/0"
return
session._dispatch_coap(
build_coap(TYPE_ACK, 0x45, mid, token, [], b"ok")
)
session.pace = pace
session._send_dgram = send
assert session.get(["device", "0"]) == (0x45, b"ok")
assert session.post(["mode", "vs", "0"], b"payload") == (0x45, b"ok")
observe_token = session.subscribe(["mode", "vs", "0"])
assert session._observe_tokens[observe_token] == "/mode/vs/0"
assert order == ["pace", "send", "pace", "send", "pace", "send"]
assert [request[1] for request in requests] == [
METHOD_GET,
METHOD_POST,
METHOD_GET,
]
def test_every_subscribe_in_registration_burst_honors_rate_limit(monkeypatch):
session = _session()
now = [100.0]
waits = []
sends = []
class StopEvent:
def wait(self, delay):
waits.append(delay)
now[0] += delay
def send(datagram):
sends.append(datagram)
session._last_send_ts = now[0]
monkeypatch.setattr(dtls_session.time, "monotonic", lambda: now[0])
session._stop = StopEvent()
session._min_req_interval = 0.2
session._last_send_ts = 0.0
session._send_dgram = send
for index in range(11):
session.subscribe(["resource", "vs", str(index)])
assert len(sends) == 11
assert waits == [pytest.approx(0.2)] * 10
def test_existing_caller_pacing_before_subscribe_does_not_wait_twice(
monkeypatch,
):
session = _session()
now = [100.05]
waits = []
class StopEvent:
def wait(self, delay):
waits.append(delay)
now[0] += delay
monkeypatch.setattr(dtls_session.time, "monotonic", lambda: now[0])
session._stop = StopEvent()
session._min_req_interval = 0.2
session._last_send_ts = 100.0
session._send_dgram = Mock()
session.pace()
session.subscribe(["mode", "vs", "0"])
assert waits == [pytest.approx(0.15)]
session._send_dgram.assert_called_once()
def test_subscribe_rechecks_liveness_after_pacing_before_registering():
session = _session()
session._send_dgram = Mock()
def close_during_pacing():
session.conn = None
session.pace = close_during_pacing
with pytest.raises(SessionClosedError):
session.subscribe(["mode", "vs", "0"])
assert session._observe_tokens == {}
session._send_dgram.assert_not_called()
def test_ack_ping_and_observe_deregister_are_not_paced():
session = _session()
session.pace = Mock(side_effect=AssertionError("control send was paced"))
class Connection:
def __init__(self):
self.sent = []
def send(self, datagram):
self.sent.append(datagram)
def bio_read(self, _size):
return b""
connection = Connection()
session.conn = connection
session.ping()
session._send_observe_dereg(b"\x40", ["mode", "vs", "0"])
session._dispatch_coap(
build_coap(TYPE_CON, 0x45, 0x1234, b"unknown", [], b"state")
)
session.pace.assert_not_called()
assert len(connection.sent) == 3
assert parse_coap(connection.sent[-1])[:4] == (TYPE_ACK, 0, 0x1234, b"")
+348
View File
@@ -0,0 +1,348 @@
"""Deterministic tests for bounded DTLS handshake timing."""
from __future__ import annotations
import socket
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionTimeoutError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.dtls_session import DtlsCoapSession
from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint
class _Clock:
def __init__(self):
self.now = 100.0
def monotonic(self):
return self.now
def advance(self, seconds):
self.now += seconds
class _Auth:
def __init__(self, clock=None, configure_delay=0.0):
self.clock = clock
self.configure_delay = configure_delay
def configure_context(self, _context):
if self.clock is not None:
self.clock.advance(self.configure_delay)
class _Connection:
def __init__(self, outcomes=None, outputs=None, timer=None):
self.outcomes = list(outcomes or ())
self.outputs = list(outputs or ())
self.timer = timer
self.bio_writes = []
self.timeout_calls = 0
def set_connect_state(self):
return None
def set_ciphertext_mtu(self, _mtu):
return None
def do_handshake(self):
outcome = self.outcomes.pop(0) if self.outcomes else "want-read"
if outcome == "want-read":
raise SSL.WantReadError()
if isinstance(outcome, Exception):
raise outcome
def bio_read(self, _size):
if self.outputs:
return self.outputs.pop(0)
raise SSL.WantReadError()
def bio_write(self, data):
self.bio_writes.append(data)
def DTLSv1_get_timeout(self):
return self.timer
def DTLSv1_handle_timeout(self):
self.timeout_calls += 1
class _Socket:
def __init__(self, clock, inbound=()):
self.clock = clock
self.inbound = list(inbound)
self.timeouts = []
self.sent = []
self.closed = False
def settimeout(self, timeout):
self.timeouts.append(timeout)
def send(self, data):
self.sent.append(data)
return len(data)
def recv(self, _size):
if self.inbound:
result = self.inbound.pop(0)
if isinstance(result, Exception):
self.clock.advance(self.timeouts[-1])
raise result
return result
self.clock.advance(self.timeouts[-1])
raise TimeoutError()
def close(self):
self.closed = True
def _session(auth=None):
return DtlsCoapSession(
"device.example",
5684,
auth=auth or _Auth(),
)
def _install_handshake(
monkeypatch,
clock,
*,
outcomes=(),
outputs=(),
inbound=(),
timer=None,
):
connection = _Connection(outcomes, outputs, timer)
sock = _Socket(clock, inbound)
endpoint = ResolvedUdpEndpoint(
socket.AF_INET,
("192.0.2.10", 5684),
)
open_calls = []
def open_socket(*args, **kwargs):
open_calls.append((args, kwargs))
sock.settimeout(kwargs["timeout"])
return sock, 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,
)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
monkeypatch.setattr(dtls_session.time, "sleep", clock.advance)
monkeypatch.setattr(
dtls_session.time,
"time",
lambda: pytest.fail("wall clock must not control handshake deadlines"),
)
return connection, sock, endpoint, open_calls
@pytest.mark.parametrize("timeout", (True, "1", object()))
def test_connect_timeout_type_is_explicit(timeout):
with pytest.raises(TypeError, match="number or None"):
_session().connect(timeout=timeout)
@pytest.mark.parametrize(
"timeout",
(
0,
-1,
float("nan"),
float("inf"),
float("-inf"),
10**1000,
),
)
def test_connect_timeout_must_be_positive_and_finite(timeout):
with pytest.raises(ValueError, match="positive finite"):
_session().connect(timeout=timeout)
def test_connect_timeout_caps_every_blocking_poll(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, open_calls = _install_handshake(
monkeypatch,
clock,
)
session = _session()
with pytest.raises(SessionTimeoutError):
session.connect(timeout=4.75)
assert clock.now == pytest.approx(104.75)
assert sock.closed
assert open_calls == [
(
("device.example", 5684),
{
"family": socket.AF_UNSPEC,
"local_port": None,
"timeout": 0.5,
},
)
]
assert max(sock.timeouts) <= 0.5
assert sock.timeouts[-1] == pytest.approx(0.25)
def test_short_timeout_is_not_rounded_up_to_poll_interval(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, open_calls = _install_handshake(
monkeypatch,
clock,
)
with pytest.raises(SessionTimeoutError):
_session().connect(timeout=0.125)
assert clock.now == pytest.approx(100.125)
assert open_calls[0][1]["timeout"] == pytest.approx(0.125)
assert sock.timeouts == pytest.approx([0.125, 0.125])
def test_default_timeout_uses_session_constant(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
)
session = _session()
session.HANDSHAKE_TIMEOUT_S = 0.2
with pytest.raises(SessionTimeoutError):
session.connect()
assert clock.now == pytest.approx(100.2)
assert sock.closed
def test_context_setup_consumes_the_same_deadline(monkeypatch):
clock = _Clock()
socket_opened = False
def open_socket(*_args, **_kwargs):
nonlocal socket_opened
socket_opened = True
raise AssertionError("expired setup must not open a socket")
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)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
session = _session(_Auth(clock, configure_delay=0.2))
with pytest.raises(SessionTimeoutError):
session.connect(timeout=0.1)
assert not socket_opened
def test_socket_setup_consumes_the_same_deadline(monkeypatch):
clock = _Clock()
connection = _Connection()
connection.do_handshake = lambda: pytest.fail(
"expired socket setup must not start a handshake"
)
sock = _Socket(clock)
endpoint = ResolvedUdpEndpoint(
socket.AF_INET,
("192.0.2.10", 5684),
)
def open_socket(*_args, **kwargs):
sock.settimeout(kwargs["timeout"])
clock.advance(0.2)
return sock, 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)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
with pytest.raises(SessionTimeoutError):
_session().connect(timeout=0.1)
assert sock.closed
def test_successful_handshake_preserves_connected_session_state(monkeypatch):
clock = _Clock()
connection, sock, endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
outcomes=("want-read", "success"),
inbound=(b"synthetic server flight",),
)
session = _session()
session.connect(timeout=1.0)
assert connection.bio_writes == [b"synthetic server flight"]
assert session.conn is connection
assert session.sock is sock
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert not sock.closed
def test_connect_services_openssl_retransmit_timer(monkeypatch):
clock = _Clock()
outbound = b"\x16\xfe\xfd" + b"\x00" * 8 + b"\x00\x01x"
connection, sock, _endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
outcomes=("want-read", "want-read", "success"),
outputs=(outbound, outbound),
inbound=(TimeoutError(), b"synthetic server flight"),
timer=0.0,
)
_session().connect(timeout=2.0)
assert connection.timeout_calls == 1
assert sock.sent == [outbound, outbound]
assert connection.bio_writes == [b"synthetic server flight"]
@pytest.mark.parametrize("success_delay", (0.1, 0.2))
def test_handshake_success_at_or_after_deadline_is_retained(
monkeypatch,
success_delay,
):
clock = _Clock()
connection, sock, endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
)
def late_success():
clock.advance(success_delay)
connection.do_handshake = late_success
session = _session()
session.connect(timeout=0.1)
assert session.conn is connection
assert session.sock is sock
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert not sock.closed
+316
View File
@@ -0,0 +1,316 @@
"""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_reported_handshake_success_wins_cancel_during_unsubscribe(
monkeypatch,
):
class CancelDuringUnsubscribe(ConnectCancellation):
def _unsubscribe(self, reader, writer):
self.set()
return super()._unsubscribe(reader, writer)
cancel = CancelDuringUnsubscribe()
connection = _Connection(succeed=True)
data_socket, peer = socket.socketpair()
endpoint = _install_connection(monkeypatch, connection, data_socket)
session = _session()
try:
session.connect(cancel=cancel)
assert cancel.is_set()
assert session.sock is data_socket
assert session.conn is connection
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert not cancel._writers
finally:
session.close()
peer.close()
def test_cancel_during_backend_failure_does_not_read_unset_completion(
monkeypatch,
):
cancel = ConnectCancellation()
connection = _Connection()
def fail_after_cancel():
cancel.set()
raise SSL.Error("synthetic backend failure")
connection.do_handshake = fail_after_cancel
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
assert not cancel._writers
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()
+240
View File
@@ -3,7 +3,12 @@
from __future__ import annotations
import threading
from types import SimpleNamespace
import pytest
from mqtt_demo import bridge as bridge_module
from mqtt_demo.bridge import PushBridge
from smartthings_local.ocf.keepalive import KeepaliveTask
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
from smartthings_local.ocf.poll_scheduler import PollScheduler, PollTier
@@ -73,3 +78,238 @@ def test_poll_scheduler_worker_stops_without_leaking_thread():
[PollTier("idle", interval_s=3600.0, paths=())],
)
_assert_worker_stops(scheduler.run_forever, "test-poll-scheduler")
class _SessionWorker:
def __init__(self, *args, **kwargs):
self.stop = None
self.started = threading.Event()
self.exited = threading.Event()
self.on_reachable = kwargs.get("on_reachable")
self.on_unreachable = kwargs.get("on_unreachable")
self.last_success_ts = 0.0
def run_forever(self, stop):
self.stop = stop
self.started.set()
stop.wait()
self.exited.set()
class _JoinedSession:
def __init__(self, workers, start_index, error=None):
self.workers = workers
self.start_index = start_index
self.error = error
def join(self):
assert all(
worker.started.wait(_THREAD_DEADLINE_S)
for worker in self.workers[self.start_index:]
)
if self.error is not None:
raise self.error
class _BlockingJoinedSession(_JoinedSession):
def __init__(self, workers):
super().__init__(workers, 0)
self.joined = threading.Event()
self.release = threading.Event()
def join(self):
super().join()
self.joined.set()
assert self.release.wait(_THREAD_DEADLINE_S)
def _bridge():
bridge = object.__new__(PushBridge)
bridge.descriptor = SimpleNamespace(
observe_paths=(),
poll_tiers=[],
is_active=lambda _state: False,
)
bridge.shared = SimpleNamespace(PING_INTERVAL_S=3600.0)
bridge.app = SimpleNamespace(klass="test")
bridge.log = SimpleNamespace(info=lambda *args: None, warning=lambda *args: None)
bridge.cache = SimpleNamespace(links={})
bridge.stop = threading.Event()
bridge._session_stop_lock = threading.Lock()
bridge._session_stop = None
bridge.scheduler = None
bridge.keepalive = None
bridge.observe_refresh = None
bridge._seed_from_device0 = lambda _session: None
bridge._retag_logger_with_serial = lambda: None
bridge.maybe_publish_state = lambda **kwargs: None
bridge.set_availability = lambda _online: None
return bridge
@pytest.mark.parametrize(
("session_count", "join_error"),
[
pytest.param(2, None, id="reconnect"),
pytest.param(1, RuntimeError("reader failed"), id="reader-error"),
],
)
def test_bridge_retires_session_workers_before_returning(
monkeypatch, session_count, join_error
):
workers = []
def worker_factory(*args, **kwargs):
worker = _SessionWorker(*args, **kwargs)
workers.append(worker)
return worker
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
bridge = _bridge()
session_stops = []
for session_index in range(session_count):
start_index = len(workers)
session = _JoinedSession(
workers,
start_index,
error=join_error if session_index == session_count - 1 else None,
)
if session.error is None:
bridge._run_session_inner(session)
else:
with pytest.raises(RuntimeError, match="reader failed"):
bridge._run_session_inner(session)
session_workers = workers[start_index:]
assert len(session_workers) == 3
assert len({id(worker.stop) for worker in session_workers}) == 1
session_stops.append(session_workers[0].stop)
assert len(workers) == session_count * 3
assert len({id(stop) for stop in session_stops}) == session_count
assert all(stop is not bridge.stop for stop in session_stops)
assert all(stop.is_set() for stop in session_stops)
assert not bridge.stop.is_set()
assert all(worker.exited.is_set() for worker in workers)
keepalive_workers = tuple(
workers[index] for index in range(1, len(workers), 3)
)
assert all(worker.on_reachable is None for worker in keepalive_workers)
assert all(worker.on_unreachable is None for worker in keepalive_workers)
assert bridge.scheduler is None
assert bridge.keepalive is None
assert bridge.observe_refresh is None
assert bridge._session_stop is None
def test_bridge_retires_started_worker_when_later_thread_fails_to_start(
monkeypatch,
):
workers = []
def worker_factory(*args, **kwargs):
worker = _SessionWorker(*args, **kwargs)
workers.append(worker)
return worker
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
bridge = _bridge()
original_start = threading.Thread.start
start_count = 0
def fail_second_start(thread):
nonlocal start_count
start_count += 1
if start_count == 2:
raise RuntimeError("synthetic thread start failure")
original_start(thread)
monkeypatch.setattr(threading.Thread, "start", fail_second_start)
with pytest.raises(RuntimeError, match="synthetic thread start failure"):
bridge._run_session_inner(SimpleNamespace(join=lambda: None))
assert len(workers) == 3
assert workers[0].started.wait(_THREAD_DEADLINE_S)
assert workers[0].exited.wait(_THREAD_DEADLINE_S)
assert workers[0].stop is not bridge.stop
assert workers[0].stop.is_set()
assert workers[1].stop is None
assert workers[2].stop is None
assert bridge.scheduler is None
assert bridge.keepalive is None
assert bridge.observe_refresh is None
assert bridge._session_stop is None
def test_request_stop_wakes_current_session_workers(monkeypatch):
workers = []
def worker_factory(*args, **kwargs):
worker = _SessionWorker(*args, **kwargs)
workers.append(worker)
return worker
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
bridge = _bridge()
session = _BlockingJoinedSession(workers)
session_thread = threading.Thread(
target=bridge._run_session_inner,
args=(session,),
daemon=True,
)
session_thread.start()
try:
assert session.joined.wait(_THREAD_DEADLINE_S)
session_stop = bridge._session_stop
assert session_stop is not None
assert not session_stop.is_set()
bridge.request_stop()
assert bridge.stop.is_set()
assert session_stop.is_set()
assert all(
worker.exited.wait(_THREAD_DEADLINE_S) for worker in workers
)
assert session_thread.is_alive()
finally:
session.release.set()
session_thread.join(_THREAD_DEADLINE_S)
assert not session_thread.is_alive()
assert bridge._session_stop is None
def test_session_workers_observe_stop_requested_before_handoff(monkeypatch):
workers = []
def worker_factory(*args, **kwargs):
worker = _SessionWorker(*args, **kwargs)
workers.append(worker)
return worker
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
bridge = _bridge()
bridge.request_stop()
bridge._run_session_inner(_JoinedSession(workers, 0))
assert len(workers) == 3
assert all(worker.stop.is_set() for worker in workers)
assert all(worker.exited.wait(_THREAD_DEADLINE_S) for worker in workers)
assert bridge._session_stop is None