Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4999e8beac | ||
|
|
b3045f5ddb | ||
|
|
0d9d13b8fb | ||
|
|
0722c552c9 | ||
|
|
a45ee4c004 | ||
|
|
512df7ff36 | ||
|
|
6b9a508fd9 | ||
|
|
ab0ec4961c | ||
|
|
6db846563a | ||
|
|
8c374e17b4 | ||
|
|
cd86424ca0 | ||
|
|
6da7e95731 | ||
|
|
a3bec9470d | ||
|
|
79493fbd48 | ||
|
|
627fcb19da | ||
|
|
31be87061a |
@@ -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
|
||||
```
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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 []
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
|
||||
@@ -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"")
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user