14 Commits
Author SHA1 Message Date
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
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
14 changed files with 939 additions and 24 deletions
+32 -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
@@ -605,6 +635,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 +662,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
+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,
)
+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)
+21
View File
@@ -18,6 +18,7 @@ from smartthings_local.protocol.dtls_session import (
ConnectCancellation,
DtlsCoapSession,
)
from smartthings_local.protocol.owner_psk import derive_mfg_certificate_owner_psk
def _assert_compatible_signature(callable_object, expected: list[str]) -> None:
@@ -114,6 +115,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