16 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
17 changed files with 1712 additions and 24 deletions
+60 -1
View File
@@ -57,6 +57,10 @@ retransmissions within that same deadline:
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:
@@ -171,6 +175,32 @@ the key must be exactly 16 or 32 bytes. `PskAuth` selects only
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
@@ -232,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
@@ -605,6 +663,7 @@ smartthings_local/ The installable library — `pip install sm
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
@@ -631,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, certificate profiles, connect deadline, session interruption)
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 []
+5 -3
View File
@@ -34,14 +34,16 @@ def _drive_dtls_handshake(
bounded only by its deadline, while the diagnostic probe retains its
explicit retry budget.
Return ``True`` only when the handshake completes before the deadline.
TLS and socket failures are left to the caller to classify.
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 time.monotonic() < deadline
return True
except SSL.WantReadError:
pass
+13 -6
View File
@@ -322,13 +322,15 @@ class DtlsCoapSession:
timeout: float | None = None,
cancel: ConnectCancellation | None = None,
):
"""Perform a cancellable DTLS handshake within a monotonic deadline.
"""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.
A ``ConnectCancellation`` wakes the network wait immediately and does
not alter an already established session.
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)
@@ -380,6 +382,7 @@ class DtlsCoapSession:
io_failed = False
cancelled = False
interrupted = False
completed = False
try:
try:
completed = _drive_dtls_handshake(
@@ -401,7 +404,7 @@ class DtlsCoapSession:
finally:
if wake_subscription is not None:
interrupted = cancel._unsubscribe(*wake_subscription)
if cancelled or interrupted:
if cancelled or (interrupted and not completed):
sock.close()
raise SessionClosedError()
if backend_failed:
@@ -876,8 +879,8 @@ class DtlsCoapSession:
deadline = time.time() + timeout
szx = BLOCK_SZX # server may negotiate down; track per-transfer
while True:
if num > 0:
self.pace()
self.pace()
self._check_live()
container = self._exchange_block(
tok, path_segs, query, num, szx, deadline)
if 'err' in container:
@@ -1023,6 +1026,8 @@ class DtlsCoapSession:
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()
@@ -1086,6 +1091,8 @@ class DtlsCoapSession:
Returns the token used (in case the caller wants to deregister
later)."""
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,
)
+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",
+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)
+51
View File
@@ -18,6 +18,11 @@ 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:
@@ -56,6 +61,32 @@ 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)
@@ -114,6 +145,26 @@ def test_psk_auth_is_a_public_authentication_provider():
)
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 = {
"close",
+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"")
+9 -5
View File
@@ -323,12 +323,12 @@ def test_connect_services_openssl_retransmit_timer(monkeypatch):
@pytest.mark.parametrize("success_delay", (0.1, 0.2))
def test_handshake_success_at_or_after_deadline_is_rejected(
def test_handshake_success_at_or_after_deadline_is_retained(
monkeypatch,
success_delay,
):
clock = _Clock()
connection, sock, _endpoint, _open_calls = _install_handshake(
connection, sock, endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
)
@@ -338,7 +338,11 @@ def test_handshake_success_at_or_after_deadline_is_rejected(
connection.do_handshake = late_success
with pytest.raises(SessionTimeoutError):
_session().connect(timeout=0.1)
session = _session()
session.connect(timeout=0.1)
assert sock.closed
assert session.conn is connection
assert session.sock is sock
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert not sock.closed
+38 -2
View File
@@ -220,9 +220,44 @@ def test_cancel_wakes_blocked_connect_without_poll_latency(monkeypatch):
peer.close()
def test_cancel_wins_race_with_reported_handshake_success(monkeypatch):
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(on_success=cancel.set, succeed=True)
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()
@@ -233,6 +268,7 @@ def test_cancel_wins_race_with_reported_handshake_success(monkeypatch):
assert data_socket.fileno() == -1
assert session.sock is None
assert session.conn is None
assert not cancel._writers
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