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

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

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

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

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

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

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

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

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

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

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

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

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

Refs QuiteYellow/SmartThings-Local#37
2026-08-14 17:33:13 +01:00
Quite Yellow d4aebebca4 Merge pull request #32 from Moballo-LLC/codex/py-06-psk-auth
feat(protocol): add PSK authentication provider
2026-08-09 14:21:54 +01:00
Quite Yellow 7f0f1e5531 Merge pull request #31 from Moballo-LLC/codex/py-05-certificate-auth
refactor(protocol): add certificate authentication provider
2026-08-09 10:41:16 +01:00
Jason Morcos 2a6fc627f2 feat(protocol): add PSK authentication provider 2026-08-08 16:06:53 -07:00
Jason Morcos 8fb37ca2ed refactor(protocol): add certificate authentication provider 2026-08-08 14:29:30 -07:00
Jack Nagy 74d78543e1 chore: allow NOTICE in distribution check allowlist
The distribution checker enforces an exact-contents allowlist; add the
NOTICE file to the wheel dist-info/licenses expectation and the sdist
required set so the packaged trademark notice passes verification.
2026-08-07 16:50:26 +01:00
Quite Yellow da8f917541 Merge pull request #29 from QuiteYellow/docs/trademark-notice
docs: add trademark non-affiliation notice (README + NOTICE)
2026-08-07 16:47:27 +01:00
Jack Nagy a63bc0bb01 docs: add trademark non-affiliation notice (README + NOTICE)
Add a Trademarks & disclaimer section to the README and a root NOTICE
file stating this is an independent, unofficial project not affiliated
with Samsung, and that Samsung/SmartThings marks are used nominatively.

Ship NOTICE inside the distributed artifacts by adding it to
license-files (wheel .dist-info/licenses/) and the sdist include list.
2026-08-07 16:46:47 +01:00
30 changed files with 6231 additions and 275 deletions
+23
View File
@@ -0,0 +1,23 @@
SmartThings-Local
Copyright (c) 2026 Jack Nagy
This software is licensed under the MIT License. See the LICENSE file for
the full terms.
------------------------------------------------------------------------
Trademarks & disclaimer
------------------------------------------------------------------------
This is an independent, unofficial project. It is NOT affiliated with,
authorised, endorsed, or sponsored by Samsung Electronics Co., Ltd. or any
of its subsidiaries.
"Samsung", "SmartThings", and any related names, marks, and logos are
trademarks of Samsung Electronics Co., Ltd. They are used in this project
only nominatively -- to identify the hardware and protocols this software
interoperates with -- and no claim is made to any right in them. Use of
these marks does not imply any affiliation with or endorsement by their
owner.
The software is provided for interoperability with hardware you own,
without warranty of any kind.
+199 -7
View File
@@ -27,12 +27,16 @@ session directly:
```python
import cbor2
from smartthings_local.protocol.auth import CertificateAuth
from smartthings_local.protocol.dtls_session import DtlsCoapSession
auth = CertificateAuth.from_files(
"certs/client_fullchain.pem",
"certs/client.key",
)
sess = DtlsCoapSession(
"192.0.2.100", 49154,
cert_path="certs/client_fullchain.pem",
key_path="certs/client.key",
auth=auth,
)
sess.connect()
sess.start_reader()
@@ -44,12 +48,159 @@ sess.subscribe(["operational", "state", "vs", "0"], # OBSERVE
sess.close()
```
If the cert/key are minted at runtime and never written to disk (e.g. inside an HA config flow), pass them in memory instead of by path:
`connect()` uses a 12-second monotonic DTLS handshake deadline by default. A
caller that needs a shorter bounded attempt can pass a positive finite value
without changing later reader timeouts. OpenSSL's DTLS timer schedules flight
retransmissions within that same deadline:
```python
sess = DtlsCoapSession("192.0.2.100", 49154, cert_pem=cert_pem, key_pem=key_pem)
sess.connect(timeout=4.0)
```
The deadline stops further setup, retries, and network waits. If OpenSSL
reports that the handshake completed at the deadline boundary, the completed
session is retained rather than torn down as a timeout.
Connection attempts can also use a one-way cancellation signal. The signal is
backed by a socketpair, so setting it wakes the network wait immediately while
OpenSSL retains control of DTLS retransmission timing:
```python
from smartthings_local.protocol.dtls_session import ConnectCancellation
cancel_connect = ConnectCancellation()
# Another thread may call cancel_connect.set().
sess.connect(timeout=8.0, cancel=cancel_connect)
```
Setting the signal stops subscribed connection attempts and closes their
temporary UDP sockets. It does not alter an already established session or add
new session lifecycle methods. Interrupted attempts raise `SessionClosedError`.
If the cert/key are minted at runtime and never written to disk (e.g. inside
an HA config flow), create the provider from memory instead:
```python
auth = CertificateAuth.from_memory(cert_pem, key_pem)
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
```
Some newer OCF-PKI devices require an exact Samsung DTLS offer and present a
hardware certificate whose subject contains a certificate UUID. That UUID can
be distinct from the runtime OCF device UUID reported by `/oic/d`, so callers
must obtain and verify the certificate identity independently. When the caller
already has an authorized client certificate and a previously verified
hardware-certificate UUID, opt in to both requirements explicitly:
```python
from smartthings_local.protocol.auth import (
CertificateAuth,
SamsungServerProfile,
)
server_profile = SamsungServerProfile.bound_device(
expected_certificate_uuid,
additional_ca_pem=additional_samsung_ca_pem,
)
auth = CertificateAuth.from_memory(
cert_pem,
key_pem,
server_profile=server_profile,
)
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
```
The default profile is restricted to Samsung home-appliance leaves with
`OU=OCF HA Device`. The profile limits the ClientHello to P-256,
`ECDHE-ECDSA-AES128-GCM-SHA256`, and the observed SHA-256/SHA-1 RSA/ECDSA
signature set, disables session tickets, preserves certificate-chain
verification, and requires the exact subject role
`C=KR, O=Samsung Electronics, OU=OCF HA Device` with a common name ending in
the expected certificate UUID. `additional_ca_pem` is optional and accepts
only a bounded PEM CA-certificate chain; it is applied only to this profiled
context. Without a profile, `CertificateAuth` retains its existing verification
behavior.
Samsung VD-family devices can present the same wire profile with the distinct
`OU=OCF VD Device` role. Select that role explicitly; profiles never fall back
between device classes:
```python
from smartthings_local.protocol.auth import (
SamsungServerProfile,
SamsungServerRole,
ServerCertificateAuth,
)
server_profile = SamsungServerProfile.bound_device(
expected_certificate_uuid,
role=SamsungServerRole.VD_DEVICE,
)
auth = ServerCertificateAuth(server_profile=server_profile)
sess = DtlsCoapSession("192.0.2.100", 5684, auth=auth)
```
`ServerCertificateAuth` is for a server-authenticated channel that does not
send a client certificate, such as the initial DTLS carrier used by
manufacturer-certificate OTM. It still verifies the CA chain, exact selected
subject role, and pinned certificate UUID. It does not learn an identity from
the first endpoint it reaches and cannot be combined with client credentials.
This API deliberately does not discover, mint, authorize, provision, rotate,
or persist credentials, and it performs no ownership transfer or OCF security
resource writes. In particular, the server-only provider can authenticate the
initial manufacturer-certificate channel, but it does not implement the OTM
that follows. The already-owned new-PKI case in
[issue #16](https://github.com/QuiteYellow/SmartThings-Local/issues/16) still
requires an authorized client identity before ordinary protected resources
can be used.
For compatibility, the existing `cert_path` / `key_path` and `cert_pem` /
`key_pem` session arguments remain supported without a deprecation warning.
They are routed through `CertificateAuth` internally. Do not combine `auth`
with those legacy arguments.
An existing OCF PSK credential can be supplied through `PskAuth`:
```python
from smartthings_local.protocol.auth import PskAuth
auth = PskAuth(identity=psk_identity, key=psk_key)
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
```
The identity must be the raw 16-byte OCF UUID and cannot contain a NUL byte;
the key must be exactly 16 or 32 bytes. `PskAuth` selects only
`ECDHE-PSK-AES128-CBC-SHA256` and does not acquire, derive, provision, rotate,
or persist credentials. Ownership transfer and credential discovery are
outside this package.
Code that has already completed an authenticated manufacturer-certificate
session can derive IoTivity's 128-bit OwnerPSK from the resulting TLS state:
```python
from smartthings_local.protocol.owner_psk import derive_mfg_certificate_owner_psk
owner_psk = derive_mfg_certificate_owner_psk(
master_secret=master_secret,
client_random=client_random,
server_random=server_random,
owner_uuid=owner_uuid,
device_uuid=device_uuid,
cipher_name=cipher_name,
oxm_label=selected_oxm_label,
)
```
The caller must supply the exact authenticated TLS values, non-nil raw OCF
UUIDs, negotiated cipher name, and label for the selected OXM. Use
`STANDARD_MFG_CERTIFICATE_OXM_LABEL` for `oic.sec.doxm.mfgcert` and
`CONFIRMED_MFG_CERTIFICATE_OXM_LABEL` for
`x.org.iotivity.conmfgcert`; do not infer the label from the appliance model.
The helper performs deterministic key derivation only: it does not access a
session, discover credentials, choose an ownership method, write security
resources, run OTM, or persist the result.
### Classified errors
Runtime transport failures use the public types in
@@ -111,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
@@ -129,7 +308,7 @@ For a full worked integration, the higher-level `smartthings_local.ocf` layer (`
Each appliance runs an independent bridge built around three coordinated pieces over one persistent DTLS session: a `StateCache` (single source of truth for all reps), a `PollScheduler` (tiered adaptive polling: hot/warm/cold plus a periodic `/device/0` sweep), and a `KeepaliveTask` (CoAP empty-CON ping for DTLS-layer liveness, with consecutive-failure detection for MQTT availability). Tier cadences are descriptor-declared and were calibrated against the empirically-measured per-firmware ceilings: dryer ~14 req/s, oven ~8 req/s. OBSERVE registrations (RFC 7641) are kept as an opportunistic freshness accelerator: when the appliance has internet and emits notifications, the cache absorbs them and the next-poll timer is reset for that resource; when it's air-gapped, polling alone carries the UX with no other code change. Token-stable Block2 (RFC 7959) handles multi-block reads. Writes are optimistically merged into the cache the moment the device 2.04-confirms, with the scheduler deferring that resource's next poll past the fetchback-revert window. Reconnect with exponential backoff on session errors, gated by a stateless DTLS ClientHello pre-flight (`smartthings_local/protocol/dtls_probe.py`) so a silent/rebooting device or wrong port drops into backoff in ~1 RTT instead of eating the full handshake timeout; when `OCF_PORT` is unset the same probe auto-discovers the live port across the OCF band.
On the currently supported firmware families, authentication uses a client cert keyed to the UUID published in Samsung's own wildcard cloud TLS cert. Their factory ACL grants that UUID `perm=31` (full CRUDN) on `href=*`. That certificate path is not universal: the WD53 profile in issue #16 and the washer in issue #20 reject it and need separate authentication work.
On the currently supported firmware families, authentication uses a client cert keyed to the UUID published in Samsung's own wildcard cloud TLS cert. Their factory ACL grants that UUID `perm=31` (full CRUDN) on `href=*`. That certificate path is not universal: the WD53 profile in issue #16 and the washer in issue #20 reject it. For those newer OCF-PKI devices, the `SamsungServerProfile` and `ServerCertificateAuth` providers (see Quick start) pin and verify the device's hardware certificate, but getting an authorized client credential to reach protected resources is still an open problem.
---
@@ -144,7 +323,7 @@ nmap -Pn -sU -p 5683,5684,49152-49160 "$APPLIANCE_IP"
Read the result:
- **`5684/udp` or a 4915x port with a DTLS first-flight response** → an OCF DTLS listener. Standard-port OCF-PKI firmware may still require an unsupported authentication profile.
- **`5684/udp` or a 4915x port with a DTLS first-flight response** → an OCF DTLS listener. Standard-port OCF-PKI firmware needs the Samsung server-certificate profile (`SamsungServerProfile` / `ServerCertificateAuth`, see Quick start), and no working client credential for it exists yet.
- **`5683/udp` responds to public OCF security/resource GETs** → use `/oic/res` to learn the device's advertised secure endpoint; do not assume that endpoint is fixed.
- **Only `8888/tcp` open (token-based HTTPS)** → older firmware (~2018–2022). **Not supported here.**
@@ -479,9 +658,12 @@ smartthings_local/ The installable library — `pip install sm
__init__.py
protocol/ DTLS-CoAP transport (reusable by any consumer, not just MQTT)
__init__.py
auth.py Immutable DTLS authentication providers
coap.py CoAP wire protocol: message encode/decode, token handling
dtls_session.py DTLS session: handshake, client-cert auth (file or in-memory PEM), Block2, liveness
dtls_probe.py Stateless DTLS liveness + opt-in stateful diagnostic
dtls_handshake.py Shared memory-BIO handshake driver, bounded by a monotonic deadline (used by session + probe)
owner_psk.py Pure manufacturer-certificate OwnerPSK derivation
ocf_root_ca.pem Samsung OCF root CA, bundled for handshake verification
ocf/ OCF resource + state layer (reusable)
__init__.py
@@ -508,7 +690,7 @@ mqtt_demo/ MQTT bridge demo (consumes smartthings_loca
.env.example Template — copy to .env, fill in
setup_cert.py One-shot cert minting script (live-fetches AC14K_M + UUID)
pyproject.toml Packaging — PyPI dist `smartthings-local`, hatch-vcs versioning
tests/ pytest suite (CoAP wire, state cache, import isolation, cert loading, DTLS probe, bridge port resolution, cert signing)
tests/ pytest suite (CoAP wire, state cache, import isolation, cert loading, DTLS probe, bridge port resolution, cert signing, certificate profiles, OwnerPSK derivation, connect deadline, session interruption)
.github/workflows/publish.yml Build + PyPI Trusted Publishing on `v*` tags
```
@@ -560,3 +742,13 @@ If reconnects become persistent (e.g. >10 in a minute) something's wrong: check
## Contributing
If you submit a PR, please don't include real device UUIDs, MACs, serials, IPs, or bearer tokens. Use the placeholders from `.env.example`.
---
## Trademarks & disclaimer
This is an independent, unofficial project. It is **not affiliated with, authorised, endorsed, or sponsored by Samsung Electronics Co., Ltd.** or any of its subsidiaries.
"Samsung", "SmartThings", and any related names, marks, and logos are trademarks of Samsung Electronics Co., Ltd. They are used in this project **only nominatively** — to identify the hardware and protocols this software interoperates with — and no claim is made to any right in them. Use of these marks does not imply any affiliation with or endorsement by their owner.
The software is provided under the [MIT License](LICENSE) for interoperability with hardware you own, without warranty of any kind.
+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 []
+2 -1
View File
@@ -9,7 +9,7 @@ description = "Local CoAP-over-DTLS (OCF) transport + polling layer for Samsung
readme = "README.md"
requires-python = ">=3.11"
license = "MIT"
license-files = ["LICENSE"]
license-files = ["LICENSE", "NOTICE"]
authors = [{ name = "Jack Nagy" }]
keywords = ["smartthings", "samsung", "ocf", "coap", "dtls", "home-assistant", "iot"]
classifiers = [
@@ -60,5 +60,6 @@ include = [
"tests",
"README.md",
"LICENSE",
"NOTICE",
"pyproject.toml",
]
+571
View File
@@ -0,0 +1,571 @@
"""Immutable authentication providers for DTLS sessions."""
from __future__ import annotations
import logging
import re
import warnings
from enum import Enum
from os import PathLike
from pathlib import Path
from typing import Protocol, runtime_checkable
from uuid import UUID
from cryptography.x509.oid import ExtensionOID
from OpenSSL import SSL, _util, crypto
logger = logging.getLogger(__name__)
_OCF_ROOT_CA = str(Path(__file__).with_name("ocf_root_ca.pem"))
_DTLS_CIPHERS = b"ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0"
_DTLS_PSK_CIPHERS = b"ECDHE-PSK-AES128-CBC-SHA256:@SECLEVEL=0"
_SAMSUNG_SERVER_CURVES = b"prime256v1"
_SAMSUNG_SERVER_SIGNATURE_ALGORITHMS = (
b"RSA+SHA256:ECDSA+SHA256:RSA+SHA1:ECDSA+SHA1"
)
_SAMSUNG_SERVER_CN_RE = re.compile(
r"\AOCF Device: [^()\r\n]{1,96} "
r"\((?P<device_identity>[0-9a-f]{8}-[0-9a-f]{4}-"
r"[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12})\)\Z",
re.IGNORECASE,
)
_PSK_CLIENT_CALLBACK_CDEF = (
"unsigned int (*)(SSL *, char *, char *, unsigned int, "
"unsigned char *, unsigned int)"
)
_PEM_CERT_RE = re.compile(
rb"-----BEGIN CERTIFICATE-----.*?-----END CERTIFICATE-----",
re.DOTALL,
)
def _verify_peer(_connection, _certificate, _error, _depth, ok):
"""Keep pyOpenSSL's existing verification result unchanged."""
return ok
def _load_pem_chain(ctx: SSL.Context, cert_pem: str, key_pem: str) -> None:
"""Load a PEM certificate chain and private key into a context in memory."""
certificates = _PEM_CERT_RE.findall(cert_pem.encode())
if not certificates:
raise ValueError("No certificates found in cert_pem")
ctx.use_certificate(
crypto.load_certificate(crypto.FILETYPE_PEM, certificates[0])
)
for extra in certificates[1:]:
ctx.add_extra_chain_cert(
crypto.load_certificate(crypto.FILETYPE_PEM, extra)
)
ctx.use_privatekey(
crypto.load_privatekey(crypto.FILETYPE_PEM, key_pem.encode())
)
ctx.check_privatekey()
@runtime_checkable
class AuthenticationProvider(Protocol):
"""Configure authentication for a newly created DTLS context."""
def configure_context(self, context: SSL.Context) -> None:
"""Configure a context while this provider remains session-owned."""
class SamsungServerRole(Enum):
"""Known Samsung OCF hardware-certificate subject roles."""
HOME_APPLIANCE = "OCF HA Device"
VD_DEVICE = "OCF VD Device"
class SamsungServerProfile:
"""Opt-in Samsung hardware-certificate verification profile."""
__slots__ = (
"_additional_ca_certificates",
"_expected_certificate_identity",
"_role",
)
def __init__(
self,
*,
expected_certificate_identity: UUID | str,
role: SamsungServerRole = SamsungServerRole.HOME_APPLIANCE,
additional_ca_pem: str | None = None,
) -> None:
if type(expected_certificate_identity) is UUID:
parsed_identity = expected_certificate_identity
elif type(expected_certificate_identity) is str:
try:
parsed_identity = UUID(expected_certificate_identity)
except ValueError:
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
) from None
if expected_certificate_identity != str(parsed_identity):
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
)
else:
raise TypeError(
"expected_certificate_identity must be a UUID or string"
)
if parsed_identity.int == 0:
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
)
if type(role) is not SamsungServerRole:
raise TypeError("role must be a SamsungServerRole")
certificates: tuple[bytes, ...] = ()
if additional_ca_pem is not None:
if type(additional_ca_pem) is not str:
raise TypeError("additional_ca_pem must be a string")
try:
raw_ca_pem = additional_ca_pem.encode("ascii")
except UnicodeEncodeError:
raise ValueError(
"additional_ca_pem must contain ASCII PEM certificates"
) from None
parsed_certificates = tuple(_PEM_CERT_RE.findall(raw_ca_pem))
if (
not 1 <= len(parsed_certificates) <= 4
or len(raw_ca_pem) > 32 * 1024
or _PEM_CERT_RE.sub(b"", raw_ca_pem).strip()
):
raise ValueError(
"additional_ca_pem must contain one to four PEM certificates"
)
try:
loaded_certificates = [
crypto.load_certificate(crypto.FILETYPE_PEM, certificate)
for certificate in parsed_certificates
]
basic_constraints = [
[
extension
for extension in certificate.to_cryptography().extensions
if extension.oid == ExtensionOID.BASIC_CONSTRAINTS
]
for certificate in loaded_certificates
]
except (crypto.Error, ValueError):
raise ValueError(
"additional_ca_pem contains an invalid certificate"
) from None
if any(
len(constraints) != 1 or not constraints[0].value.ca
for constraints in basic_constraints
):
raise ValueError(
"additional_ca_pem must contain only CA certificates"
)
fingerprints = {
crypto.dump_certificate(crypto.FILETYPE_ASN1, certificate)
for certificate in loaded_certificates
}
if len(fingerprints) != len(loaded_certificates):
raise ValueError(
"additional_ca_pem must not contain duplicate certificates"
)
certificates = parsed_certificates
object.__setattr__(
self,
"_expected_certificate_identity",
parsed_identity,
)
object.__setattr__(self, "_role", role)
object.__setattr__(self, "_additional_ca_certificates", certificates)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("SamsungServerProfile is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("SamsungServerProfile is immutable")
@classmethod
def bound_device(
cls,
expected_certificate_identity: UUID | str,
*,
role: SamsungServerRole = SamsungServerRole.HOME_APPLIANCE,
additional_ca_pem: str | None = None,
) -> SamsungServerProfile:
"""Bind a verified Samsung hardware leaf to its certificate UUID."""
return cls(
expected_certificate_identity=expected_certificate_identity,
role=role,
additional_ca_pem=additional_ca_pem,
)
def __repr__(self) -> str:
"""Return a representation without device or trust-chain details."""
return "SamsungServerProfile()"
def _configure_context(self, context: SSL.Context) -> None:
curve_setter = getattr(_util.lib, "SSL_CTX_set1_curves_list", None)
if curve_setter is not None:
if curve_setter(context._context, _SAMSUNG_SERVER_CURVES) != 1:
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
)
else:
# pyOpenSSL 23.1 does not expose SSL_CTX_set1_curves_list. Its
# public set_tmp_ecdh fallback produces the same single P-256
# supported-groups ClientHello extension; wire-level tests protect
# that compatibility path. Newer pyOpenSSL uses the exact setter.
try:
with warnings.catch_warnings():
warnings.simplefilter("ignore", DeprecationWarning)
curve = crypto.get_elliptic_curve("prime256v1")
context.set_tmp_ecdh(curve)
except (AttributeError, TypeError, ValueError, SSL.Error):
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
) from None
signature_setter = getattr(
_util.lib,
"SSL_CTX_set1_sigalgs_list",
None,
)
if (
signature_setter is None
or signature_setter(
context._context,
_SAMSUNG_SERVER_SIGNATURE_ALGORITHMS,
)
!= 1
):
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
)
context.set_options(SSL.OP_NO_TICKET)
if self._additional_ca_certificates:
store = context.get_cert_store()
try:
for certificate in self._additional_ca_certificates:
store.add_cert(
crypto.load_certificate(
crypto.FILETYPE_PEM,
certificate,
)
)
except crypto.Error:
raise RuntimeError(
"OpenSSL rejected the Samsung server trust profile"
) from None
def _verify_peer(
self,
_connection,
certificate,
_error,
depth,
ok,
) -> bool:
if not ok or certificate is None or depth < 0:
return False
if depth > 0:
return True
try:
with warnings.catch_warnings():
# pyOpenSSL deprecates this API in favor of cryptography, but
# reparsing Samsung's non-DER factory leaves with cryptography
# rejects certificates that OpenSSL has already verified.
warnings.simplefilter("ignore", DeprecationWarning)
components = [
(name.decode("ascii"), value.decode("ascii"))
for name, value in certificate.get_subject().get_components()
]
except (
AttributeError,
TypeError,
UnicodeDecodeError,
ValueError,
crypto.Error,
):
logger.warning("Unable to parse Samsung server certificate subject")
return False
common_names = [value for name, value in components if name == "CN"]
organizational_units = [
value for name, value in components if name == "OU"
]
organizations = [value for name, value in components if name == "O"]
countries = [value for name, value in components if name == "C"]
# This deliberately pins the complete Samsung subject role. The OCF
# reference implementation reads only the UUID-bearing CN, but
# relaxing C/O/OU here could accept a different certificate cohort.
if (
len(common_names) != 1
or organizational_units != [self._role.value]
or organizations != ["Samsung Electronics"]
or countries != ["KR"]
):
return False
match = _SAMSUNG_SERVER_CN_RE.fullmatch(common_names[0])
return (
match is not None
and UUID(match.group("device_identity"))
== self._expected_certificate_identity
)
def _configure_certificate_server(
context: SSL.Context,
server_profile: SamsungServerProfile | None,
) -> None:
"""Configure certificate-server verification for one DTLS context."""
context.load_verify_locations(_OCF_ROOT_CA)
if server_profile is None:
context.set_verify(SSL.VERIFY_PEER, _verify_peer)
else:
server_profile._configure_context(context)
context.set_verify(
SSL.VERIFY_PEER,
server_profile._verify_peer,
)
# @SECLEVEL=0 permits SHA-1 in Samsung's server cert chain (AC14K_M
# intermediate is SHA-1 signed). This is the only channel that reaches
# the OpenSSL instance cryptography bundles; ctypes and cffi bindings
# do not expose SSL_CTX_set_security_level on this build.
context.set_cipher_list(_DTLS_CIPHERS)
class ServerCertificateAuth:
"""Verify a pinned Samsung server without a client certificate."""
__slots__ = ("_server_profile",)
def __init__(self, *, server_profile: SamsungServerProfile) -> None:
if type(server_profile) is not SamsungServerProfile:
raise TypeError("server_profile must be a SamsungServerProfile")
object.__setattr__(self, "_server_profile", server_profile)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("ServerCertificateAuth is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("ServerCertificateAuth is immutable")
def __repr__(self) -> str:
"""Return a representation without server identity or trust details."""
return "ServerCertificateAuth()"
def configure_context(self, context: SSL.Context) -> None:
"""Verify the selected server profile without loading client material."""
_configure_certificate_server(context, self._server_profile)
class CertificateAuth:
"""Certificate authentication loaded from files or in-memory PEM data.
Use :meth:`from_files` or :meth:`from_memory` to create an instance.
Credential sources are intentionally not exposed as public attributes.
"""
__slots__ = (
"_certificate_path",
"_certificate_pem",
"_private_key_path",
"_private_key_pem",
"_server_profile",
)
def __init__(
self,
*,
certificate_path: str | PathLike[str] | None = None,
private_key_path: str | PathLike[str] | None = None,
certificate_pem: str | None = None,
private_key_pem: str | None = None,
server_profile: SamsungServerProfile | None = None,
) -> None:
file_supplied = (
certificate_path is not None or private_key_path is not None
)
memory_supplied = (
certificate_pem is not None or private_key_pem is not None
)
if file_supplied and memory_supplied:
raise ValueError(
"pass either certificate_path/private_key_path or "
"certificate_pem/private_key_pem, not both"
)
if file_supplied:
if certificate_path is None or private_key_path is None:
raise ValueError(
"certificate_path and private_key_path must be passed together"
)
elif memory_supplied:
if certificate_pem is None or private_key_pem is None:
raise ValueError(
"certificate_pem and private_key_pem must be passed together"
)
else:
raise ValueError(
"must pass either certificate_path/private_key_path or "
"certificate_pem/private_key_pem"
)
if (
server_profile is not None
and type(server_profile) is not SamsungServerProfile
):
raise TypeError("server_profile must be a SamsungServerProfile")
object.__setattr__(
self,
"_certificate_path",
str(certificate_path) if certificate_path is not None else None,
)
object.__setattr__(
self,
"_private_key_path",
str(private_key_path) if private_key_path is not None else None,
)
object.__setattr__(self, "_certificate_pem", certificate_pem)
object.__setattr__(self, "_private_key_pem", private_key_pem)
object.__setattr__(self, "_server_profile", server_profile)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("CertificateAuth is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("CertificateAuth is immutable")
@classmethod
def from_files(
cls,
certificate_path: str | PathLike[str],
private_key_path: str | PathLike[str],
*,
server_profile: SamsungServerProfile | None = None,
) -> CertificateAuth:
"""Create a provider backed by certificate-chain and key files."""
return cls(
certificate_path=certificate_path,
private_key_path=private_key_path,
server_profile=server_profile,
)
@classmethod
def from_memory(
cls,
certificate_pem: str,
private_key_pem: str,
*,
server_profile: SamsungServerProfile | None = None,
) -> CertificateAuth:
"""Create a provider backed by an in-memory PEM chain and key."""
return cls(
certificate_pem=certificate_pem,
private_key_pem=private_key_pem,
server_profile=server_profile,
)
def __repr__(self) -> str:
"""Return a representation that never includes credential material."""
return "CertificateAuth()"
def configure_context(self, context: SSL.Context) -> None:
"""Apply the existing certificate authentication profile to a context."""
_configure_certificate_server(context, self._server_profile)
if self._certificate_pem is not None:
_load_pem_chain(
context,
self._certificate_pem,
self._private_key_pem,
)
else:
context.use_certificate_chain_file(self._certificate_path)
context.use_privatekey_file(self._private_key_path)
context.check_privatekey()
class PskAuth:
"""DTLS authentication using an existing OCF PSK credential.
The identity must be a raw 16-byte OCF UUID. The key must contain 16 or
32 bytes. Credential material is intentionally not exposed as public
attributes and is never included in this provider's representation. A
configured context must not outlive this provider; ``DtlsCoapSession``
enforces that lifetime by retaining its provider.
"""
__slots__ = ("_callback",)
def __init__(self, *, identity: bytes, key: bytes) -> None:
if type(identity) is not bytes or type(key) is not bytes:
raise TypeError("identity and key must be bytes")
if len(identity) != 16:
raise ValueError("identity must be a raw 16-byte OCF UUID")
if b"\x00" in identity:
raise ValueError("identity cannot contain a NUL byte")
if len(key) not in (16, 32):
raise ValueError("key must be 16 or 32 bytes")
ffi = _util.ffi
@ffi.callback(_PSK_CLIENT_CALLBACK_CDEF)
def client_callback(
_ssl,
_identity_hint,
identity_buffer,
max_identity_length,
key_buffer,
max_key_length,
):
# OpenSSL callbacks cannot propagate Python exceptions. Fail
# before touching either destination when a buffer is unavailable
# or too small for the complete credential.
if (
identity_buffer == ffi.NULL
or key_buffer == ffi.NULL
or len(identity) + 1 > max_identity_length
or len(key) > max_key_length
):
return 0
ffi.memmove(
identity_buffer,
identity + b"\x00",
len(identity) + 1,
)
ffi.memmove(key_buffer, key, len(key))
return len(key)
object.__setattr__(self, "_callback", client_callback)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("PskAuth is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("PskAuth is immutable")
def __repr__(self) -> str:
"""Return a representation that never includes credential material."""
return "PskAuth()"
def configure_context(self, context: SSL.Context) -> None:
"""Configure one context for the narrow Samsung OCF PSK profile."""
setter = getattr(_util.lib, "SSL_CTX_set_psk_client_callback", None)
if setter is None:
raise RuntimeError(
"the installed OpenSSL binding does not support DTLS PSK"
)
context.set_cipher_list(_DTLS_PSK_CIPHERS)
setter(context._context, self._callback)
__all__ = [
"AuthenticationProvider",
"CertificateAuth",
"PskAuth",
"SamsungServerProfile",
"SamsungServerRole",
"ServerCertificateAuth",
]
+9
View File
@@ -13,6 +13,7 @@ from ..errors import MalformedMessageError
URI_PATH = 11
URI_QUERY = 15
OBSERVE = 6
ETAG = 4
CONTENT_FORMAT = 12
ACCEPT = 17
BLOCK2 = 23
@@ -121,6 +122,14 @@ def block_value(num, more, szx):
return struct.pack('>I', v)[1:]
def block_fields(value):
"""Decode a CoAP Block-N option value. Inverse of block_value().
Returns (num, more, szx). An empty value means block 0, no more,
SZX=0 — RFC 7959 §2.2 allows a zero-length option to elide it."""
v = int.from_bytes(value, 'big')
return v >> 4, (v >> 3) & 1, v & 0x07
def fmt_code(c):
"""0x45 → '2.05', 0x84 → '4.04'. Used in log lines."""
return f"{c >> 5}.{c & 0x1F:02d}"
@@ -0,0 +1,99 @@
"""Shared memory-BIO driver for bounded DTLS handshakes."""
from __future__ import annotations
import select
import time
from collections.abc import Callable
from OpenSSL import SSL
from .coap import split_dtls
_HANDSHAKE_POLL_S = 0.5
_MAX_DATAGRAM_SIZE = 65535
class _HandshakeCancelled(Exception):
"""Internal signal that a handshake wake socket became readable."""
def _drive_dtls_handshake(
connection,
sock,
*,
deadline: float,
retries: int | None = None,
on_datagram: Callable[[bytes], None] | None = None,
wake_socket=None,
) -> bool:
"""Drive one memory-BIO DTLS handshake up to a monotonic deadline.
OpenSSL owns the retransmission schedule. ``retries`` optionally limits
how many expired retransmission timers are serviced; a normal session is
bounded only by its deadline, while the diagnostic probe retains its
explicit retry budget.
Return ``True`` once OpenSSL reports the handshake complete. The deadline
prevents another setup, retry, or network-wait iteration; it does not tear
down a session that completed while ``do_handshake()`` was running. TLS and
socket failures are left to the caller to classify.
"""
retransmits = 0
while time.monotonic() < deadline:
try:
connection.do_handshake()
return True
except SSL.WantReadError:
pass
try:
output = connection.bio_read(_MAX_DATAGRAM_SIZE)
except SSL.WantReadError:
output = None
if output:
for record in split_dtls(output):
if sock.send(record) != len(record):
raise OSError("incomplete UDP send")
remaining = deadline - time.monotonic()
if remaining <= 0:
break
wait = min(_HANDSHAKE_POLL_S, remaining)
timed_out = False
if wake_socket is None:
sock.settimeout(wait)
try:
datagram = sock.recv(_MAX_DATAGRAM_SIZE)
except TimeoutError:
timed_out = True
else:
readable, _, _ = select.select(
(sock, wake_socket),
(),
(),
wait,
)
if wake_socket in readable:
raise _HandshakeCancelled()
if sock not in readable:
timed_out = True
else:
datagram = sock.recv(_MAX_DATAGRAM_SIZE)
if timed_out:
timer = connection.DTLSv1_get_timeout()
if timer is not None and timer <= 0:
if retries is not None and retransmits >= retries:
break
connection.DTLSv1_handle_timeout()
retransmits += 1
continue
if not datagram:
continue
if on_datagram is not None:
on_datagram(datagram)
connection.bio_write(datagram)
return False
+33 -65
View File
@@ -39,8 +39,9 @@ from dataclasses import dataclass
from OpenSSL import SSL
from ..errors import ProbeError
from .auth import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
from .coap import split_dtls
from .dtls_session import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
from .dtls_handshake import _drive_dtls_handshake
from .endpoint import open_connected_udp_socket
# DTLS record content types (RFC 6347 §4.1)
@@ -603,73 +604,40 @@ def diagnose_dtls_handshake(
started = time.monotonic()
deadline = started + timeout
seen = set()
retransmits = 0
def record_datagram(datagram):
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
result.datagrams.append(datagram)
for content_type, detail in classify_datagram(datagram):
if content_type == _CT_HANDSHAKE:
if detail not in seen:
seen.add(detail)
result.handshake_msgs.append(detail)
if result.outcome == DEAD:
result.outcome = LIVE
elif content_type == _CT_ALERT and detail is not None:
level, name = detail
result.alert = (level, name)
if level == 2: # fatal
result.outcome = REJECTED
try:
while time.monotonic() < deadline:
try:
conn.do_handshake()
result.outcome = COMPLETED
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
break
except SSL.WantReadError:
pass
except SSL.Error:
# A fatal Alert lands here; the alert record was already
# captured below, so classification still works.
result.error = ProbeError()
break
try:
o = conn.bio_read(65535)
if o:
for r in split_dtls(o):
if sock.send(r) != len(r):
raise OSError('short UDP send')
except SSL.WantReadError:
pass
remaining = deadline - time.monotonic()
if remaining <= 0:
break
sock.settimeout(min(0.5, remaining))
try:
d = sock.recv(65535)
except TimeoutError:
# No answer to the last flight. Service OpenSSL's DTLS
# retransmit timer: once it has counted down to 0,
# handle_timeout() re-queues the previous flight into the
# write BIO for the next iteration to flush. A live server
# answers within a flight or two; a silent/non-DTLS port
# never does, so we give up only after `retries`
# retransmits — one dropped ClientHello no longer reads as
# a false DEAD.
to = conn.DTLSv1_get_timeout()
if to is not None and to <= 0:
if retransmits >= retries:
break
conn.DTLSv1_handle_timeout()
retransmits += 1
continue
if not d:
continue
completed = _drive_dtls_handshake(
conn,
sock,
deadline=deadline,
retries=retries,
on_datagram=record_datagram,
)
if completed:
result.outcome = COMPLETED
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
result.datagrams.append(d)
for ct, detail in classify_datagram(d):
if ct == _CT_HANDSHAKE:
if detail not in seen:
seen.add(detail)
result.handshake_msgs.append(detail)
if result.outcome == DEAD:
result.outcome = LIVE
elif ct == _CT_ALERT and detail is not None:
level, name = detail
result.alert = (level, name)
if level == 2: # fatal
result.outcome = REJECTED
conn.bio_write(d)
except SSL.Error:
# A fatal Alert lands here; record_datagram() has already classified
# the alert record before it is fed back into OpenSSL.
result.error = ProbeError()
except OSError:
result.error = ProbeError()
finally:
File diff suppressed because it is too large Load Diff
+322
View File
@@ -0,0 +1,322 @@
"""Bounded discovery of a known host's plaintext OCF response port.
Some OCF devices receive multicast discovery on UDP 5683 but reply from an
ephemeral port. This module records only token-correlated response ports from
the caller's expected IPv4 address. The results are candidates: callers still
need directory parsing, a DTLS probe, and authenticated identity validation.
"""
from __future__ import annotations
import ipaddress
import math
import secrets
import selectors
import socket
import time
from dataclasses import dataclass
from ..errors import MalformedMessageError
from .coap import (
ACCEPT,
CF_CBOR,
METHOD_GET,
TYPE_ACK,
TYPE_CON,
TYPE_NON,
URI_PATH,
URI_QUERY,
build_coap,
parse_coap,
)
__all__ = [
"OcfResponderPortDiscoveryResult",
"discover_ocf_responder_ports",
]
_OCF_MULTICAST_GROUP = socket.inet_ntoa(bytes((224, 0, 1, 187)))
_OCF_DISCOVERY_PORT = 5683
_OCF_CBOR = (10_000).to_bytes(2, "big")
_OCF_CONTENT_FORMAT_VERSION = 2049
_OCF_VERSION_1_0 = (2048).to_bytes(2, "big")
_CONTENT = 0x45
_MAX_DATAGRAM_BYTES = 8192
_MAX_DATAGRAMS_PER_ROUND = 64
_MAX_PORTS = 8
@dataclass(frozen=True, slots=True, repr=False)
class OcfResponderPortDiscoveryResult:
"""Redacted result of one known-host multicast discovery operation."""
ports: tuple[int, ...]
attempts: int
responses: int
error_code: str | None = None
@property
def found(self) -> bool:
"""Return whether at least one response port was discovered."""
return bool(self.ports)
def __repr__(self) -> str:
return (
"OcfResponderPortDiscoveryResult("
f"found={self.found!r}, port_count={len(self.ports)}, "
f"attempts={self.attempts}, responses={self.responses}, "
f"error_code={self.error_code!r})"
)
def _validate_address(value: object, name: str) -> tuple[str, bytes]:
if not isinstance(value, str):
raise TypeError(f"{name} must be an IPv4 address string")
try:
address = ipaddress.IPv4Address(value)
except ipaddress.AddressValueError as exc:
raise ValueError(f"{name} must be a valid IPv4 address") from exc
if address.is_multicast or address.is_unspecified or address.is_reserved:
raise ValueError(f"{name} must be a unicast IPv4 address")
return str(address), address.packed
def _validate_options(
*,
discovery_port: object,
timeout: object,
rounds: object,
) -> tuple[int, float, int]:
if isinstance(discovery_port, bool) or not isinstance(discovery_port, int):
raise TypeError("discovery_port must be an integer")
if not 1 <= discovery_port <= 65535:
raise ValueError("discovery_port must be between 1 and 65535")
if isinstance(timeout, bool) or not isinstance(timeout, (int, float)):
raise TypeError("timeout must be a number")
timeout_value = float(timeout)
if not math.isfinite(timeout_value) or not 0 < timeout_value <= 30:
raise ValueError("timeout must be greater than zero and at most 30")
if isinstance(rounds, bool) or not isinstance(rounds, int):
raise TypeError("rounds must be an integer")
if not 1 <= rounds <= 4:
raise ValueError("rounds must be between one and four")
return discovery_port, timeout_value, rounds
def _request(
token: bytes,
message_id: int,
*,
versioned: bool,
filtered: bool,
) -> bytes:
options = [
(URI_PATH, b"oic"),
(URI_PATH, b"res"),
(ACCEPT, _OCF_CBOR if versioned else CF_CBOR),
]
if filtered:
options.append((URI_QUERY, b"rt=oic.r.doxm"))
if versioned:
options.append((_OCF_CONTENT_FORMAT_VERSION, _OCF_VERSION_1_0))
return build_coap(TYPE_NON, METHOD_GET, message_id, token, options)
def _result(
ports: tuple[int, ...],
attempts: int,
responses: int,
error_code: str | None = None,
) -> OcfResponderPortDiscoveryResult:
return OcfResponderPortDiscoveryResult(
ports=ports,
attempts=attempts,
responses=responses,
error_code=error_code,
)
def discover_ocf_responder_ports(
target_address: str,
*,
interface_address: str,
discovery_port: int = _OCF_DISCOVERY_PORT,
timeout: float = 3.0,
rounds: int = 2,
) -> OcfResponderPortDiscoveryResult:
"""Find plaintext OCF response ports for one known IPv4 host.
Each round sends modern OCF and legacy IoTivity NON requests to the
link-local multicast group. Only a 2.05 response with a request token and
the exact target source address contributes a candidate. One monotonic
deadline bounds all rounds, and every socket is closed before return.
"""
_target_address, target_key = _validate_address(target_address, "target_address")
interface_address, interface_key = _validate_address(
interface_address, "interface_address"
)
discovery_port, timeout, rounds = _validate_options(
discovery_port=discovery_port,
timeout=timeout,
rounds=rounds,
)
try:
selector = selectors.DefaultSelector()
except (OSError, ValueError):
return _result((), 0, 0, "interface_unavailable")
active = None
try:
active = socket.socket(socket.AF_INET, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_IF, interface_key)
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_TTL, 1)
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_LOOP, 0)
active.bind((interface_address, 0))
active.setblocking(False)
selector.register(active, selectors.EVENT_READ)
except (OSError, ValueError):
if active is not None:
try:
active.close()
except OSError:
pass
selector.close()
return _result((), 0, 0, "interface_unavailable")
started = time.monotonic()
deadline = started + timeout
accepted_tokens: set[bytes] = set()
observations: set[tuple[bytes, int]] = set()
ports: list[int] = []
seen_ports: set[int] = set()
attempts = 0
responses = 0
too_many_ports = False
try:
for round_number in range(rounds):
if time.monotonic() >= deadline:
break
# Preserve the unfiltered modern and legacy requests used by the
# installed appliance generations. Older media firmware can omit
# usable endpoint policy from its large unfiltered directory but
# answer the smaller legacy DOXM-filtered lookup, so send that as
# a third bounded fallback rather than narrowing every request.
for versioned, filtered in (
(True, False),
(False, False),
(False, True),
):
token = secrets.token_bytes(8)
while token in accepted_tokens:
token = secrets.token_bytes(8)
accepted_tokens.add(token)
request = _request(
token,
secrets.randbits(16),
versioned=versioned,
filtered=filtered,
)
try:
sent = active.sendto(
request,
(_OCF_MULTICAST_GROUP, discovery_port),
)
except OSError:
continue
attempts += 1
if sent != len(request):
continue
round_deadline = started + timeout * (round_number + 1) / rounds
datagrams = 0
while datagrams < _MAX_DATAGRAMS_PER_ROUND:
remaining = min(deadline, round_deadline) - time.monotonic()
if remaining <= 0:
break
try:
events = selector.select(remaining)
except (OSError, ValueError):
break
if not events:
break
try:
datagram, source = active.recvfrom(_MAX_DATAGRAM_BYTES + 1)
except (BlockingIOError, OSError):
continue
datagrams += 1
if len(datagram) > _MAX_DATAGRAM_BYTES:
continue
if (
len(datagram) < 4
or datagram[0] >> 6 != 1
or datagram[0] & 0x0F > 8
or 4 + (datagram[0] & 0x0F) > len(datagram)
):
continue
if not isinstance(source, tuple) or len(source) != 2:
continue
source_host, source_port = source
if not isinstance(source_host, str):
continue
try:
source_key = socket.inet_pton(socket.AF_INET, source_host)
except OSError:
continue
if source_key != target_key:
continue
if (
isinstance(source_port, bool)
or not isinstance(source_port, int)
or not 1 <= source_port <= 65535
):
continue
try:
message_type, code, mid, token, _options, payload = parse_coap(
datagram
)
except (IndexError, ValueError, MalformedMessageError):
continue
if (
token not in accepted_tokens
or code != _CONTENT
or message_type not in (TYPE_NON, TYPE_CON)
or not payload
):
continue
if message_type == TYPE_CON:
try:
active.sendto(build_coap(TYPE_ACK, 0, mid, b"", []), source)
except OSError:
pass
observation = (token, source_port)
if observation in observations:
continue
observations.add(observation)
responses += 1
if source_port in seen_ports:
continue
if len(ports) >= _MAX_PORTS:
too_many_ports = True
continue
seen_ports.add(source_port)
ports.append(source_port)
if too_many_ports:
return _result((), attempts, responses, "ambiguous_response")
if ports:
return _result(tuple(ports), attempts, responses)
if attempts == 0:
return _result((), attempts, responses, "interface_unavailable")
return _result((), attempts, responses, "no_response")
finally:
try:
selector.unregister(active)
except (KeyError, OSError, ValueError):
pass
try:
active.close()
except OSError:
pass
selector.close()
+128
View File
@@ -0,0 +1,128 @@
"""Pure IoTivity manufacturer-certificate OwnerPSK derivation."""
from __future__ import annotations
from collections.abc import Mapping
import hashlib
import hmac
from types import MappingProxyType
from typing import Final
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL: Final = b"x.org.iotivity.conmfgcert"
STANDARD_MFG_CERTIFICATE_OXM_LABEL: Final = b"oic.sec.doxm.mfgcert"
# OpenSSL cipher names mapped to the key-block lengths used by IoTivity's
# CAGenerateOwnerPSK implementation.
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS: Final[Mapping[str, int]] = MappingProxyType(
{
"ECDHE-ECDSA-AES128-SHA256": 96,
"ECDHE-ECDSA-AES128-CCM": 40,
"ECDHE-ECDSA-AES128-CCM8": 40,
"ECDHE-ECDSA-AES128-GCM-SHA256": 120,
"AES256-SHA256": 128,
"ECDHE-ECDSA-AES256-SHA384": 160,
"ECDHE-ECDSA-AES256-GCM-SHA384": 184,
"AES128-GCM-SHA256": 120,
}
)
_TLS_MASTER_SECRET_BYTES: Final = 48
_TLS_RANDOM_BYTES: Final = 32
_OCF_UUID_BYTES: Final = 16
_OWNER_PSK_BYTES: Final = 16
def _require_bytes(name: str, value: bytes, length: int) -> bytes:
if not isinstance(value, bytes):
raise TypeError(f"{name} must be bytes")
if len(value) != length:
raise ValueError(f"{name} must be exactly {length} bytes")
return value
def _require_uuid(name: str, value: bytes) -> bytes:
value = _require_bytes(name, value, _OCF_UUID_BYTES)
if not any(value):
raise ValueError(f"{name} must not be the nil UUID")
return value
def _tls12_p_hash_sha256(
key: bytes,
label: bytes,
random1: bytes,
random2: bytes,
length: int,
) -> bytes:
seed = label + random1 + random2
a_value = hmac.new(key, seed, hashlib.sha256).digest()
output = bytearray()
while len(output) < length:
output.extend(hmac.new(key, a_value + seed, hashlib.sha256).digest())
a_value = hmac.new(key, a_value, hashlib.sha256).digest()
return bytes(output[:length])
def derive_mfg_certificate_owner_psk(
*,
master_secret: bytes,
client_random: bytes,
server_random: bytes,
owner_uuid: bytes,
device_uuid: bytes,
cipher_name: str,
oxm_label: bytes,
) -> bytes:
"""Derive a 128-bit OwnerPSK from caller-supplied DTLS state.
This implements IoTivity's two-stage TLS 1.2 SHA-256 P_hash operation.
It performs no session access, network I/O, ownership writes, or storage.
The caller must supply state from an authenticated manufacturer-certificate
session and explicitly select the OXM label used by that transaction.
IoTivity's other 96-byte ECDH_ANON, ECDHE_PSK, and ECDHE_RSA mappings are
intentionally outside this helper's manufacturer-certificate allowlist.
"""
if not isinstance(cipher_name, str):
raise TypeError("cipher_name must be a string")
key_block_bytes = MFG_CERTIFICATE_KEY_BLOCK_LENGTHS.get(cipher_name)
if key_block_bytes is None:
raise ValueError("unexpected manufacturer-certificate DTLS cipher")
master_secret = _require_bytes(
"master_secret", master_secret, _TLS_MASTER_SECRET_BYTES
)
client_random = _require_bytes(
"client_random", client_random, _TLS_RANDOM_BYTES
)
server_random = _require_bytes(
"server_random", server_random, _TLS_RANDOM_BYTES
)
owner_uuid = _require_uuid("owner_uuid", owner_uuid)
device_uuid = _require_uuid("device_uuid", device_uuid)
if not isinstance(oxm_label, bytes):
raise TypeError("oxm_label must be bytes")
if oxm_label not in {
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
STANDARD_MFG_CERTIFICATE_OXM_LABEL,
}:
raise ValueError("unexpected manufacturer-certificate OXM label")
key_block = _tls12_p_hash_sha256(
master_secret,
b"key expansion",
server_random,
client_random,
key_block_bytes,
)
# IoTivity's OTM callers pass the owner UUID first and the target device
# UUID second. The lower adapter's historical rsrc/prov parameter names
# describe those arguments inconsistently, so preserve the caller order.
return _tls12_p_hash_sha256(
key_block,
oxm_label,
owner_uuid,
device_uuid,
_OWNER_PSK_BYTES,
)
File diff suppressed because it is too large Load Diff
+19 -1
View File
@@ -2,7 +2,8 @@ import pytest
from smartthings_local.errors import MalformedMessageError
from smartthings_local.protocol.coap import (
build_coap, parse_coap, encode_options, block_value, fmt_code,
build_coap, parse_coap, encode_options, block_value, block_fields,
fmt_code,
TYPE_CON, METHOD_GET, URI_PATH, ACCEPT, CF_CBOR, BLOCK2,
)
@@ -43,6 +44,23 @@ def test_block_value_promotes_to_two_bytes_when_num_is_large():
assert len(v) == 2
@pytest.mark.parametrize('num, more, szx', [
(0, 0, 0),
(0, 1, 6),
(2, 1, 6),
(1, 0, 4),
(0xFFF, 0, 0),
(0xFFFFF, 1, 7),
])
def test_block_fields_inverts_block_value(num, more, szx):
assert block_fields(block_value(num, more, szx)) == (num, more, szx)
def test_block_fields_treats_empty_value_as_block_zero():
# RFC 7959 §2.2: a zero-length Block option means num=0, m=0, szx=0.
assert block_fields(b'') == (0, 0, 0)
def test_fmt_code_formats_class_dot_detail():
assert fmt_code(0x45) == '2.05'
assert fmt_code(0x84) == '4.04'
+272 -18
View File
@@ -1,31 +1,69 @@
import gc
import traceback
import weakref
from dataclasses import asdict
from datetime import datetime, timedelta, timezone
import pytest
from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.x509.oid import NameOID
from OpenSSL import SSL, crypto
from smartthings_local.protocol.auth import CertificateAuth
from smartthings_local.protocol.dtls_session import DtlsCoapSession, _load_pem_chain
def _make_self_signed_pem_pair():
"""A throwaway self-signed cert + key, just to exercise PEM loading —
not meant to resemble a real Samsung client cert."""
key = crypto.PKey()
key.generate_key(crypto.TYPE_RSA, 2048)
def _make_generated_pem_chain():
"""Create a throwaway leaf + root chain unrelated to Samsung devices."""
now = datetime.now(timezone.utc)
root_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
root_name = x509.Name(
[x509.NameAttribute(NameOID.COMMON_NAME, "Synthetic test root")]
)
root_cert = (
x509.CertificateBuilder()
.subject_name(root_name)
.issuer_name(root_name)
.public_key(root_key.public_key())
.serial_number(1)
.not_valid_before(now - timedelta(minutes=1))
.not_valid_after(now + timedelta(hours=1))
.add_extension(x509.BasicConstraints(ca=True, path_length=None), True)
.sign(root_key, hashes.SHA256())
)
cert = crypto.X509()
cert.get_subject().CN = "test"
cert.set_serial_number(1)
cert.gmtime_adj_notBefore(0)
cert.gmtime_adj_notAfter(3600)
cert.set_issuer(cert.get_subject())
cert.set_pubkey(key)
cert.sign(key, "sha256")
leaf_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
leaf_name = x509.Name(
[x509.NameAttribute(NameOID.COMMON_NAME, "Synthetic test client")]
)
leaf_cert = (
x509.CertificateBuilder()
.subject_name(leaf_name)
.issuer_name(root_name)
.public_key(leaf_key.public_key())
.serial_number(2)
.not_valid_before(now - timedelta(minutes=1))
.not_valid_after(now + timedelta(hours=1))
.add_extension(x509.BasicConstraints(ca=False, path_length=None), True)
.sign(root_key, hashes.SHA256())
)
cert_pem = crypto.dump_certificate(crypto.FILETYPE_PEM, cert).decode()
key_pem = crypto.dump_privatekey(crypto.FILETYPE_PEM, key).decode()
cert_pem = (
leaf_cert.public_bytes(serialization.Encoding.PEM)
+ root_cert.public_bytes(serialization.Encoding.PEM)
).decode()
key_pem = leaf_key.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.PKCS8,
serialization.NoEncryption(),
).decode()
return cert_pem, key_pem
def test_load_pem_chain_loads_cert_and_key_in_memory():
cert_pem, key_pem = _make_self_signed_pem_pair()
cert_pem, key_pem = _make_generated_pem_chain()
ctx = SSL.Context(SSL.DTLS_METHOD)
_load_pem_chain(ctx, cert_pem, key_pem)
ctx.check_privatekey() # raises if cert/key don't match
@@ -37,7 +75,7 @@ def test_load_pem_chain_rejects_cert_pem_with_no_certificates():
def test_session_requires_exactly_one_cert_source():
cert_pem, key_pem = _make_self_signed_pem_pair()
cert_pem, key_pem = _make_generated_pem_chain()
with pytest.raises(ValueError):
DtlsCoapSession("host", 1234) # neither pair given
@@ -49,10 +87,226 @@ def test_session_requires_exactly_one_cert_source():
with pytest.raises(ValueError):
DtlsCoapSession("host", 1234, cert_pem=cert_pem) # key_pem missing
with pytest.raises(ValueError):
DtlsCoapSession("host", 1234, cert_path="/a") # key_path missing
def test_session_rejects_provider_with_legacy_certificate_arguments():
cert_pem, key_pem = _make_generated_pem_chain()
auth = CertificateAuth.from_memory(cert_pem, key_pem)
with pytest.raises(ValueError, match="auth or legacy certificate"):
DtlsCoapSession(
"host",
1234,
cert_pem=cert_pem,
key_pem=key_pem,
auth=auth,
)
def test_session_rejects_object_that_is_not_an_authentication_provider():
with pytest.raises(TypeError, match="AuthenticationProvider"):
DtlsCoapSession("host", 1234, auth=object())
def test_session_accepts_explicit_certificate_provider():
cert_pem, key_pem = _make_generated_pem_chain()
auth = CertificateAuth.from_memory(cert_pem, key_pem)
session = DtlsCoapSession("host", 1234, auth=auth)
assert session.auth is auth
assert session.cert_path is None
assert session.key_path is None
assert session.cert_pem is None
assert session.key_pem is None
def test_session_retains_authentication_provider_for_its_lifetime():
class RetainedProvider:
def configure_context(self, _context):
return None
auth = RetainedProvider()
reference = weakref.ref(auth)
session = DtlsCoapSession("host", 1234, auth=auth)
del auth
gc.collect()
assert reference() is session.auth
def test_session_accepts_pem_pair():
cert_pem, key_pem = _make_self_signed_pem_pair()
cert_pem, key_pem = _make_generated_pem_chain()
sess = DtlsCoapSession("host", 1234, cert_pem=cert_pem, key_pem=key_pem)
assert sess.cert_path is None
assert sess.key_path is None
assert sess.cert_pem == cert_pem
assert isinstance(sess.auth, CertificateAuth)
def test_session_routes_legacy_file_pair_through_certificate_auth(tmp_path):
cert_path = tmp_path / "client.pem"
key_path = tmp_path / "client-key.pem"
session = DtlsCoapSession(
"host",
1234,
cert_path=cert_path,
key_path=key_path,
)
assert isinstance(session.auth, CertificateAuth)
assert session.cert_path == str(cert_path)
assert session.key_path == str(key_path)
def test_certificate_auth_loads_generated_chain_from_memory_and_files(tmp_path):
cert_pem, key_pem = _make_generated_pem_chain()
memory_context = SSL.Context(SSL.DTLS_METHOD)
CertificateAuth.from_memory(cert_pem, key_pem).configure_context(
memory_context
)
memory_context.check_privatekey()
cert_path = tmp_path / "client.pem"
key_path = tmp_path / "client-key.pem"
cert_path.write_text(cert_pem)
key_path.write_text(key_pem)
file_context = SSL.Context(SSL.DTLS_METHOD)
CertificateAuth.from_files(cert_path, key_path).configure_context(file_context)
file_context.check_privatekey()
def test_certificate_auth_rejects_invalid_memory_material():
auth = CertificateAuth.from_memory("not a certificate", "not a key")
with pytest.raises(ValueError, match="No certificates found"):
auth.configure_context(SSL.Context(SSL.DTLS_METHOD))
def test_invalid_certificate_error_does_not_include_credential_material():
marker = "credential" + "-marker"
certificate_blob = (
"-----BEGIN CERTIFICATE-----\n"
f"{marker}\n"
"-----END CERTIFICATE-----\n"
)
key_blob = "invalid-" + marker
auth = CertificateAuth.from_memory(certificate_blob, key_blob)
with pytest.raises(crypto.Error) as captured:
auth.configure_context(SSL.Context(SSL.DTLS_METHOD))
rendered = (
str(captured.value)
+ repr(captured.value)
+ "".join(traceback.format_exception(captured.value))
)
assert marker not in rendered
def test_certificate_auth_rejects_invalid_file_material(tmp_path):
cert_path = tmp_path / "invalid.pem"
key_path = tmp_path / "invalid-key.pem"
cert_path.write_text("invalid")
key_path.write_text("invalid")
with pytest.raises(SSL.Error):
CertificateAuth.from_files(cert_path, key_path).configure_context(
SSL.Context(SSL.DTLS_METHOD)
)
def test_certificate_auth_rejects_incomplete_or_mixed_sources():
with pytest.raises(ValueError):
CertificateAuth()
with pytest.raises(ValueError):
CertificateAuth(certificate_path="/synthetic/client.pem")
with pytest.raises(ValueError):
CertificateAuth(certificate_pem="certificate")
certificate_path = "/synthetic/client.pem"
key_path = "/synthetic/client-key.pem"
certificate_data = "certificate"
key_data = "key"
with pytest.raises(ValueError):
CertificateAuth(
certificate_path=certificate_path,
private_key_path=key_path,
certificate_pem=certificate_data,
private_key_pem=key_data,
)
def test_certificate_auth_is_immutable_and_has_secret_safe_repr():
cert_pem, key_pem = _make_generated_pem_chain()
auth = CertificateAuth.from_memory(cert_pem, key_pem)
rendered = repr(auth)
assert rendered == "CertificateAuth()"
assert cert_pem not in rendered
assert key_pem not in rendered
with pytest.raises(AttributeError, match="immutable"):
auth.certificate_pem = None
with pytest.raises(AttributeError, match="immutable"):
del auth._certificate_pem
def test_certificate_auth_has_no_public_or_dataclass_credential_surface():
cert_pem, key_pem = _make_generated_pem_chain()
auth = CertificateAuth.from_memory(cert_pem, key_pem)
assert not hasattr(auth, "certificate_pem")
assert not hasattr(auth, "private_key_pem")
with pytest.raises(TypeError):
vars(auth)
with pytest.raises(TypeError):
asdict(auth)
def test_certificate_auth_context_setup_matches_legacy_happy_path():
class RecordingContext:
def __init__(self):
self.calls = []
self.verify_callback = None
def load_verify_locations(self, path):
self.calls.append(("load_verify_locations", path))
def set_verify(self, mode, callback):
self.calls.append(("set_verify", mode))
self.verify_callback = callback
def set_cipher_list(self, ciphers):
self.calls.append(("set_cipher_list", ciphers))
def use_certificate_chain_file(self, path):
self.calls.append(("use_certificate_chain_file", path))
def use_privatekey_file(self, path):
self.calls.append(("use_privatekey_file", path))
def check_privatekey(self):
self.calls.append(("check_privatekey",))
context = RecordingContext()
CertificateAuth.from_files(
"/synthetic/client.pem",
"/synthetic/client-key.pem",
).configure_context(context)
assert [call[0] for call in context.calls] == [
"load_verify_locations",
"set_verify",
"set_cipher_list",
"use_certificate_chain_file",
"use_privatekey_file",
"check_privatekey",
]
assert context.calls[1] == ("set_verify", SSL.VERIFY_PEER)
assert context.calls[2] == (
"set_cipher_list",
b"ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0",
)
callback = context.verify_callback
assert callback(None, None, 0, 0, True) is True
assert callback(None, None, 1, 0, False) is False
+173
View File
@@ -0,0 +1,173 @@
"""Reader-thread death visibility (QuiteYellow/SmartThings-Local#37).
A connected UDP socket surfaces ICMP errors on recv; before this the
reader exited silently on the first one and every later request waited
out its full timeout against a session nobody was reading. These tests
pin the three behaviours that fixed it: advisory ICMP errnos keep the
reader alive, a real socket error exits with a WARNING and clears
_reader_running, and callers then fail fast with SessionClosedError.
"""
import errno
import logging
import socket
import threading
import time
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionClosedError
from smartthings_local.protocol.dtls_session import DtlsCoapSession
_LOGGER_NAME = "smartthings_local.protocol.dtls_session"
class _NullAuth:
"""Structural AuthenticationProvider — never configured, we skip connect()."""
def configure_context(self, _context):
return None
class _FakeConn:
"""Minimal stand-in for SSL.Connection: each datagram written to the
BIO surfaces as one decrypted packet on the next recv(), then
WantReadError like a drained DTLS record buffer."""
def __init__(self):
self._decrypted = []
def bio_write(self, datagram):
self._decrypted.append(datagram)
def recv(self, _n):
if self._decrypted:
return self._decrypted.pop(0)
raise SSL.WantReadError()
def bio_read(self, _n):
return b""
def send(self, _datagram):
return None
def shutdown(self):
return None
class _FakeSock:
"""Scripted UDP socket. Each step is bytes to return, an exception to
raise, or a callable to run (then a timeout, so the loop re-checks
_stop). An exhausted script blocks like a real recv timeout; once
close()d it raises EBADF the way a closed fd does."""
def __init__(self, steps=()):
self._steps = list(steps)
self.closed = False
self.timeout = None
def settimeout(self, value):
self.timeout = value
def recv(self, _n):
if self.closed:
raise OSError(errno.EBADF, "bad file descriptor")
if self._steps:
step = self._steps.pop(0)
if callable(step):
step()
raise socket.timeout()
if isinstance(step, BaseException):
raise step
return step
time.sleep(0.01)
raise socket.timeout()
def send(self, data):
return len(data)
def close(self):
self.closed = True
def _make_session():
sess = DtlsCoapSession("host", 1234, auth=_NullAuth())
sess.conn = _FakeConn()
return sess
def _run_reader(sess, steps, timeout=2.0):
sess.sock = _FakeSock(steps)
sess.start_reader()
sess._reader_thread.join(timeout)
assert not sess._reader_thread.is_alive(), "reader thread did not exit"
def test_advisory_icmp_error_does_not_kill_reader(caplog):
sess = _make_session()
dispatched = []
sess._dispatch_coap = dispatched.append
with caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME):
_run_reader(sess, [
OSError(errno.ECONNREFUSED, "connection refused"),
b"\x60\x00\x00\x00", # survives, gets dispatched
lambda: sess._stop.set(), # end the loop cleanly
])
assert dispatched == [b"\x60\x00\x00\x00"]
assert not sess._reader_running.is_set()
assert any(r.levelno == logging.DEBUG and "advisory" in r.getMessage()
for r in caplog.records)
# An advisory errno is not a real exit — no WARNING.
assert not any(r.levelno >= logging.WARNING for r in caplog.records)
def test_fatal_socket_error_exits_with_warning(caplog):
sess = _make_session()
with caplog.at_level(logging.WARNING, logger=_LOGGER_NAME):
_run_reader(sess, [OSError(errno.EBADF, "bad file descriptor")])
assert not sess._reader_running.is_set()
warnings = [r for r in caplog.records if r.levelno == logging.WARNING]
assert len(warnings) == 1
assert "reader exiting" in warnings[0].getMessage()
def test_request_fails_fast_after_reader_death():
sess = _make_session()
_run_reader(sess, [OSError(errno.EBADF, "bad file descriptor")])
assert not sess._reader_running.is_set()
start = time.monotonic()
with pytest.raises(SessionClosedError):
sess.get(["oic", "d"], timeout=10.0)
elapsed = time.monotonic() - start
# The whole point: no waiting out the request timeout.
assert elapsed < 1.0, f"get() waited {elapsed:.2f}s instead of failing fast"
def test_close_does_not_log_warning_on_teardown(caplog):
sess = _make_session()
sess.sock = _FakeSock() # empty script: blocks on recv
sess.start_reader()
time.sleep(0.05) # let the reader reach recv
with caplog.at_level(logging.WARNING, logger=_LOGGER_NAME):
sess.close()
sess._reader_thread.join(2.0)
assert not sess._reader_thread.is_alive()
assert not sess._reader_running.is_set()
assert not any(r.levelno >= logging.WARNING for r in caplog.records)
def test_check_live_without_reader_matches_old_conn_guard():
sess = _make_session() # conn set, reader never started
assert sess._reader_thread is None
sess._check_live() # must not raise — config-flow behaviour
sess.conn = None
with pytest.raises(SessionClosedError):
sess._check_live()
+1 -1
View File
@@ -376,7 +376,7 @@ def test_session_uses_connected_socket_send_and_recv(monkeypatch):
assert open_calls == [(('device.example', 5684), {
'family': socket.AF_INET6,
'local_port': None,
'timeout': 2.0,
'timeout': 0.5,
})]
session.close()
+1
View File
@@ -17,6 +17,7 @@ def test_smartthings_local_imports_without_mqtt_demo_present(tmp_path):
import_lines = [
"import smartthings_local.protocol.coap",
"import smartthings_local.protocol.ocf_multicast",
"import smartthings_local.protocol.dtls_session",
"import smartthings_local.ocf.state_cache",
"import smartthings_local.ocf.poll_scheduler",
+508
View File
@@ -0,0 +1,508 @@
"""Blockwise OBSERVE notifications (QuiteYellow/SmartThings-Local#39).
A notification carries only the first block of the representation
(RFC 7959 §2.6). Before this, _dispatch_coap handed that first block
straight to on_notification and the consumer decoded a truncated CBOR
buffer. These tests pin the replacement: a truncated notification is
withheld, the resource is re-read from block 0 under a fresh one-shot
token, and only the reassembled representation reaches the callback.
The re-read starts at block 0 rather than continuing at NUM=1 for two
reasons, both recorded on #39: RFC 7959 §3.4 forbids continuing on the
observation's token, and Samsung's RT-OCF drops a transfer that opens
at NUM>0 under a token it has not seen.
"""
import logging
import socket
import threading
import time
import pytest
from OpenSSL import SSL
from smartthings_local.errors import BlockwiseError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.coap import (
BLOCK2, ETAG, METHOD_GET, OBSERVE, TYPE_ACK, TYPE_NON,
block_fields, block_value, build_coap, parse_coap,
)
from smartthings_local.protocol.dtls_session import DtlsCoapSession
SZX = 6 # 1024-byte blocks, the only size these appliances honour
_LOGGER_NAME = "smartthings_local.protocol.dtls_session"
class _NullAuth:
"""Structural AuthenticationProvider — never configured, we skip connect()."""
def configure_context(self, _context):
return None
class _LoopbackConn:
"""SSL.Connection stand-in that answers requests from a script.
`responder` is called with each parsed request and returns a list of
CoAP datagrams to hand back (possibly empty, to model a silent
device). Responses surface on recv() the way decrypted records do.
"""
def __init__(self, responder):
self._responder = responder
self._inbox = []
self._lock = threading.Lock()
self.sent = []
# -- client -> device
def send(self, datagram):
self.sent.append(parse_coap(datagram))
for reply in self._responder(parse_coap(datagram)):
with self._lock:
self._inbox.append(reply)
return len(datagram)
def bio_read(self, _n):
return b""
# -- device -> client
def inject(self, datagram):
"""Push a device-initiated frame (an OBSERVE notification)."""
with self._lock:
self._inbox.append(datagram)
def bio_write(self, _datagram):
return None
def recv(self, _n):
with self._lock:
if self._inbox:
return self._inbox.pop(0)
raise SSL.WantReadError()
def shutdown(self):
return None
def pending(self):
with self._lock:
return len(self._inbox)
class _PumpSock:
"""UDP socket stand-in. recv() returns a dummy datagram whenever the
connection has something decrypted waiting, so the reader loop keeps
pumping; otherwise it times out like a real socket."""
def __init__(self, conn):
self._conn = conn
self.closed = False
def settimeout(self, _value):
return None
def recv(self, _n):
for _ in range(20):
if self.closed:
raise OSError("closed")
if self._conn.pending():
return b"\x00"
time.sleep(0.005)
raise socket.timeout()
def send(self, data):
return len(data)
def close(self):
self.closed = True
def _make_session(responder, **kwargs):
calls = []
sess = DtlsCoapSession(
"host", 1234, auth=_NullAuth(),
on_notification=lambda href, payload: calls.append((href, payload)),
**kwargs)
sess.conn = _LoopbackConn(responder)
sess.sock = _PumpSock(sess.conn)
sess.start_reader()
return sess, calls
def _notification(tok, payload, *, block2=None, mtype=TYPE_NON, obs=1):
opts = [(OBSERVE, bytes([obs]))]
if block2 is not None:
opts.append((BLOCK2, block2))
return build_coap(mtype, 0x45, 0x1234, tok, opts, payload)
def _content(tok, mid, payload, *, block2=None, etag=None):
opts = []
if etag is not None:
opts.append((ETAG, etag))
if block2 is not None:
opts.append((BLOCK2, block2))
return build_coap(TYPE_ACK, 0x45, mid, tok, opts, payload)
def _requested_block(request):
"""(num, szx) the request asked for, or (0, None) with no Block2."""
_, _, _, _, opts, _ = request
b2 = [v for n, v in opts if n == BLOCK2]
if not b2:
return 0, None
num, _, szx = block_fields(b2[0])
return num, szx
def _wait_for(predicate, timeout=3.0):
deadline = time.time() + timeout
while time.time() < deadline:
if predicate():
return True
time.sleep(0.01)
return False
def _close(sess):
sess.close()
sess.join()
# --------------------------------------------------------------------
# The no-regression case
def test_single_block_notification_is_delivered_inline():
sess, calls = _make_session(lambda request: [])
try:
tok = sess.subscribe(["oven", "vs", "0"])
sess.conn.inject(_notification(tok, b"\xa1\x01\x02"))
assert _wait_for(lambda: calls)
assert calls == [("/oven/vs/0", b"\xa1\x01\x02")]
# The subscribe GET is the only thing we sent — no refetch.
assert len(sess.conn.sent) == 1
finally:
_close(sess)
def test_complete_block_zero_notification_is_delivered_inline():
"""Block2 present but M=0 and NUM=0 means the whole representation
fit in one block. Nothing to fetch back."""
sess, calls = _make_session(lambda request: [])
try:
tok = sess.subscribe(["oven", "vs", "0"])
sess.conn.inject(
_notification(tok, b"\xa1\x01\x02",
block2=block_value(0, 0, SZX)))
assert _wait_for(lambda: calls)
assert calls == [("/oven/vs/0", b"\xa1\x01\x02")]
assert len(sess.conn.sent) == 1
finally:
_close(sess)
# --------------------------------------------------------------------
# The fix
def test_truncated_notification_is_refetched_and_reassembled():
blocks = [b"A" * 1024, b"B" * 40]
def responder(request):
_mtype, code, mid, tok, opts, _ = request
if code != METHOD_GET or any(n == OBSERVE for n, _ in opts):
return [] # the subscribe registration itself
num, _ = _requested_block(request)
more = 1 if num + 1 < len(blocks) else 0
return [_content(tok, mid, blocks[num],
block2=block_value(num, more, SZX))]
sess, calls = _make_session(responder)
try:
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, blocks[0],
block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: calls), "callback never fired"
assert calls == [("/mode/vs/0", b"".join(blocks))]
finally:
_close(sess)
def test_refetch_uses_a_fresh_one_shot_token_not_the_observe_token():
"""RFC 7959 §3.4: the requests for additional blocks cannot use the
token of the observation relationship."""
blocks = [b"A" * 1024, b"B" * 40]
def responder(request):
_mtype, _code, mid, tok, opts, _ = request
if any(n == OBSERVE for n, _ in opts):
return []
num, _ = _requested_block(request)
more = 1 if num + 1 < len(blocks) else 0
return [_content(tok, mid, blocks[num],
block2=block_value(num, more, SZX))]
sess, calls = _make_session(responder)
try:
observe_tok = sess.subscribe(["mode", "vs", "0"])
assert len(observe_tok) == 1, "OBSERVE registrations use 1-byte tokens"
sess.conn.inject(
_notification(observe_tok, blocks[0],
block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: calls)
refetch = [r for r in sess.conn.sent
if not any(n == OBSERVE for n, _ in r[4])]
assert refetch, "no refetch request was sent"
tokens = {r[3] for r in refetch}
assert observe_tok not in tokens
assert all(len(t) == 4 for t in tokens), "one-shot tokens are 4-byte"
assert len(tokens) == 1, "the transfer must hold one token throughout"
# And the transfer restarts at block 0 rather than continuing at 1.
assert _requested_block(refetch[0])[0] == 0
assert [_requested_block(r)[0] for r in refetch] == [0, 1]
# No Observe option on a continuation request.
assert not any(n == OBSERVE for r in refetch for n, _ in r[4])
finally:
_close(sess)
def test_silent_device_drops_the_notification_without_delivering_a_partial():
sess, calls = _make_session(lambda request: [])
try:
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, b"A" * 1024,
block2=block_value(0, 1, SZX)))
# Give the worker a chance to try and fail. _BLOCK_ACK_TIMEOUT is
# 4s per attempt, so we only need to see that nothing partial got
# through in the meantime.
assert not _wait_for(lambda: calls, timeout=0.6)
assert sess._reader_thread.is_alive(), "reader must survive"
finally:
_close(sess)
def test_non_2xx_refetch_is_dropped_rather_than_delivered():
def responder(request):
_mtype, _code, mid, tok, opts, _ = request
if any(n == OBSERVE for n, _ in opts):
return []
return [build_coap(TYPE_ACK, 0x84, mid, tok, [], b"")]
sess, calls = _make_session(responder)
try:
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, b"A" * 1024,
block2=block_value(0, 1, SZX)))
assert not _wait_for(lambda: calls, timeout=0.6)
finally:
_close(sess)
def test_notification_burst_collapses_to_one_refetch_per_resource():
blocks = [b"A" * 1024, b"B" * 40]
gate = threading.Event()
def responder(request):
_mtype, _code, mid, tok, opts, _ = request
if any(n == OBSERVE for n, _ in opts):
return []
gate.wait(2.0) # hold the first transfer open
num, _ = _requested_block(request)
more = 1 if num + 1 < len(blocks) else 0
return [_content(tok, mid, blocks[num],
block2=block_value(num, more, SZX))]
sess, calls = _make_session(responder)
try:
tok = sess.subscribe(["mode", "vs", "0"])
for seq in range(5):
sess.conn.inject(
_notification(tok, blocks[0], obs=seq + 1,
block2=block_value(0, 1, SZX)))
# Five notifications, one queue entry: latest wins per href.
assert _wait_for(lambda: sess._refetch_pending or sess.conn.sent[1:])
assert len(sess._refetch_pending) <= 1
gate.set()
assert _wait_for(lambda: calls)
assert _wait_for(
lambda: not sess._refetch_pending and len(calls) >= 1)
time.sleep(0.2)
# Two transfers at most: the one in flight when the burst landed,
# plus one for the final state.
starts = [r for r in sess.conn.sent
if not any(n == OBSERVE for n, _ in r[4])
and _requested_block(r)[0] == 0]
assert len(starts) <= 2, f"{len(starts)} refetches for one burst"
assert calls[-1] == ("/mode/vs/0", b"".join(blocks))
finally:
gate.set()
_close(sess)
def test_close_during_a_queued_refetch_stops_the_worker():
sess, _calls = _make_session(lambda request: [])
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, b"A" * 1024, block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: sess._refetch_thread is not None)
sess.close()
sess.join() # hangs if the worker outlives the session
assert not sess._refetch_thread.is_alive()
def test_refetch_worker_exits_when_the_reader_dies():
sess, _calls = _make_session(lambda request: [])
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, b"A" * 1024, block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: sess._refetch_thread is not None)
# Kill the reader the way a socket error does, without close().
sess.sock.closed = True
assert _wait_for(lambda: not sess._reader_running.is_set(), timeout=5.0)
sess._refetch_thread.join(6.0)
assert not sess._refetch_thread.is_alive()
sess.close()
def test_debug_bridge_promotes_the_refetch_outcome_to_info(monkeypatch, caplog):
"""The hardware validation for #39 reads this line to confirm which
token the re-read used, so it has to survive the bridge's INFO
default. Without DEBUG_BRIDGE it stays at debug."""
monkeypatch.setattr(dtls_session, "DEBUG_BRIDGE", True)
blocks = [b"A" * 1024, b"B" * 40]
def responder(request):
_mtype, _code, mid, tok, opts, _ = request
if any(n == OBSERVE for n, _ in opts):
return []
num, _ = _requested_block(request)
more = 1 if num + 1 < len(blocks) else 0
return [_content(tok, mid, blocks[num],
block2=block_value(num, more, SZX))]
sess, calls = _make_session(responder)
try:
with caplog.at_level(logging.INFO, logger=_LOGGER_NAME):
tok = sess.subscribe(["mode", "vs", "0"])
sess.conn.inject(
_notification(tok, blocks[0], block2=block_value(0, 1, SZX)))
assert _wait_for(lambda: calls)
line = next((r.getMessage() for r in caplog.records
if r.getMessage().startswith("refetch /mode/vs/0")), None)
assert line is not None, "no refetch line at INFO"
assert "blocks=2" in line
assert f"bytes={sum(len(b) for b in blocks)}" in line
assert line.endswith("ok")
# The token in the line is the one-shot token, not the observe one.
assert f"tok={tok.hex()} " not in line
finally:
_close(sess)
# --------------------------------------------------------------------
# Shared Block2 loop hardening
def test_stale_block_number_is_not_concatenated():
"""A retransmit of block 0 arriving while we wait for block 1 must
not be appended as if it were block 1."""
served = []
def responder(request):
_mtype, _code, mid, tok, _opts, _ = request
num, _ = _requested_block(request)
served.append(num)
if num == 0:
return [_content(tok, mid, b"A" * 1024,
block2=block_value(0, 1, SZX))]
# Answer the block-1 request with a duplicate of block 0 first.
return [
_content(tok, mid, b"A" * 1024, block2=block_value(0, 1, SZX)),
_content(tok, mid, b"B" * 40, block2=block_value(1, 0, SZX)),
]
sess, _calls = _make_session(responder)
try:
code, payload = sess.get(["mode", "vs", "0"], timeout=5.0)
assert code == 0x45
assert payload == b"A" * 1024 + b"B" * 40
finally:
_close(sess)
def test_etag_change_mid_transfer_restarts_then_fails():
etags = [b"\x01", b"\x02", b"\x03", b"\x04"]
def responder(request):
_mtype, _code, mid, tok, _opts, _ = request
num, _ = _requested_block(request)
# A different ETag on every single response: the representation
# never settles, so reassembly can never be consistent.
etag = etags.pop(0) if etags else b"\xff"
payload = b"A" * 1024 if num == 0 else b"B" * 40
more = 1 if num == 0 else 0
return [_content(tok, mid, payload, etag=etag,
block2=block_value(num, more, SZX))]
sess, _calls = _make_session(responder)
try:
with pytest.raises(BlockwiseError):
sess.get(["mode", "vs", "0"], timeout=5.0)
finally:
_close(sess)
def test_stable_etag_across_blocks_reassembles():
def responder(request):
_mtype, _code, mid, tok, _opts, _ = request
num, _ = _requested_block(request)
payload = b"A" * 1024 if num == 0 else b"B" * 40
more = 1 if num == 0 else 0
return [_content(tok, mid, payload, etag=b"\x77",
block2=block_value(num, more, SZX))]
sess, _calls = _make_session(responder)
try:
code, payload = sess.get(["mode", "vs", "0"], timeout=5.0)
assert code == 0x45
assert payload == b"A" * 1024 + b"B" * 40
finally:
_close(sess)
def test_szx_downshift_asks_for_the_block_after_what_we_have():
"""Server answers block 0 at SZX=6 (1024B) then drops to SZX=4
(256B). Block numbers index the new size, so the next request is
block 4, not block 1."""
requested = []
def responder(request):
_mtype, _code, mid, tok, _opts, _ = request
num, szx = _requested_block(request)
requested.append((num, szx))
if num == 0:
return [_content(tok, mid, b"A" * 1024,
block2=block_value(0, 1, 4))]
return [_content(tok, mid, b"B" * 100,
block2=block_value(num, 0, 4))]
sess, _calls = _make_session(responder)
try:
code, payload = sess.get(["mode", "vs", "0"], timeout=5.0)
assert code == 0x45
assert payload == b"A" * 1024 + b"B" * 100
# 1024 bytes in hand at 256B blocks = blocks 0..3 done, ask for 4.
assert requested[1] == (4, 4)
finally:
_close(sess)
+392
View File
@@ -0,0 +1,392 @@
"""Known-host OCF multicast responder-port discovery tests."""
from __future__ import annotations
from types import SimpleNamespace
import pytest
from smartthings_local.protocol.coap import (
ACCEPT,
TYPE_CON,
TYPE_NON,
URI_PATH,
URI_QUERY,
build_coap,
parse_coap,
)
from smartthings_local.protocol.ocf_multicast import (
_OCF_MULTICAST_GROUP,
OcfResponderPortDiscoveryResult,
discover_ocf_responder_ports,
)
_TARGET = "192.0.2.20"
_OTHER = "192.0.2.21"
_INTERFACE = "192.0.2.10"
class _FakeSocket:
def __init__(self, responder=None):
self.responder = responder
self.incoming = []
self.sent = []
self.options = []
self.bound = None
self.blocking = None
self.closed = False
def setsockopt(self, level, option, value):
self.options.append((level, option, value))
def bind(self, address):
self.bound = address
def setblocking(self, enabled):
self.blocking = enabled
def sendto(self, datagram, destination):
self.sent.append((datagram, destination))
if self.responder is not None:
self.incoming.extend(self.responder(datagram, destination))
return len(datagram)
def recvfrom(self, _size):
return self.incoming.pop(0)
def close(self):
self.closed = True
class _FakeSelector:
def __init__(self, active):
self.active = active
self.closed = False
def register(self, *_args):
return None
def unregister(self, *_args):
return None
def select(self, _timeout):
if not self.active.incoming:
return []
return [(SimpleNamespace(fileobj=self.active), selectors_event_read())]
def close(self):
self.closed = True
def selectors_event_read():
return 1
def _response(
datagram, _destination=None, *, host=_TARGET, port=43123, message_type=TYPE_NON
):
_mtype, _code, mid, token, _options, _payload = parse_coap(datagram)
response_mid = mid if message_type != TYPE_CON else (mid + 1) & 0xFFFF
return [
(
build_coap(message_type, 0x45, response_mid, token, [], b"directory"),
(host, port),
)
]
@pytest.fixture
def patch_socket(monkeypatch):
created = []
def install(responder=None):
active = _FakeSocket(responder)
selector = _FakeSelector(active)
created.append((active, selector))
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.socket.socket",
lambda *_args: active,
)
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.selectors.DefaultSelector",
lambda: selector,
)
return active, selector
return install
def test_sends_proven_directory_requests_and_filtered_fallback_on_interface(
patch_socket,
):
active, selector = patch_socket(_response)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.ports == (43123,)
assert result.attempts == 3
assert result.responses == 3
assert active.bound == (_INTERFACE, 0)
assert active.blocking is False
assert active.closed
assert selector.closed
assert all(
destination == (_OCF_MULTICAST_GROUP, 5683) for _, destination in active.sent
)
requests = [parse_coap(datagram) for datagram, _ in active.sent]
assert all(
[value for number, value in item[4] if number == URI_PATH] == [b"oic", b"res"]
for item in requests
)
queries = [
[value for number, value in item[4] if number == URI_QUERY] for item in requests
]
assert queries == [[], [], [b"rt=oic.r.doxm"]]
accepts = [
[value for number, value in item[4] if number == ACCEPT] for item in requests
]
assert accepts == [[b"\x27\x10"], [b"\x3c"], [b"\x3c"]]
assert [value for number, value in requests[0][4] if number == 2049] == [
b"\x08\x00"
]
assert [value for number, value in requests[1][4] if number == 2049] == []
assert [value for number, value in requests[2][4] if number == 2049] == []
@pytest.mark.parametrize(
("accepted_accept", "accepted_query"),
(
(b"\x27\x10", ()),
(b"\x3c", ()),
(b"\x3c", (b"rt=oic.r.doxm",)),
),
)
def test_each_directory_request_profile_can_find_the_responder(
patch_socket, accepted_accept, accepted_query
):
def responder(datagram, destination):
_mtype, _code, _mid, _token, options, _payload = parse_coap(datagram)
accept = next(value for number, value in options if number == ACCEPT)
query = tuple(value for number, value in options if number == URI_QUERY)
if (accept, query) != (accepted_accept, accepted_query):
return []
return _response(datagram, destination)
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.ports == (43123,)
assert result.responses == 1
def test_accepts_only_token_correlated_content_from_the_target(patch_socket):
def responder(datagram, _destination):
responses = _response(datagram)
_mtype, _code, mid, token, _options, _payload = parse_coap(datagram)
responses.extend(
[
(build_coap(TYPE_NON, 0x45, mid, b"wrong", [], b"x"), (_TARGET, 49999)),
(build_coap(TYPE_NON, 0x44, mid, token, [], b"x"), (_TARGET, 49998)),
(build_coap(TYPE_NON, 0x45, mid, token, [], b"x"), (_OTHER, 49997)),
(build_coap(TYPE_NON, 0x45, mid, token, [], b""), (_TARGET, 49996)),
]
)
return responses
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.ports == (43123,)
assert result.responses == 3
def test_ignores_malformed_and_oversized_datagrams(patch_socket):
def responder(datagram, _destination):
valid = _response(datagram)
invalid_version = bytes([valid[0][0][0] & 0x3F]) + valid[0][0][1:]
invalid_token_length = bytes([0x49]) + valid[0][0][1:]
return [
(b"\x40", (_TARGET, 49999)),
(invalid_version, (_TARGET, 49998)),
(invalid_token_length, (_TARGET, 49997)),
(b"x" * 8193, (_TARGET, 49996)),
*valid,
]
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.ports == (43123,)
assert result.responses == 3
def test_acknowledges_confirmable_responses(patch_socket):
active, _selector = patch_socket(
lambda datagram, _destination: _response(datagram, message_type=TYPE_CON)
)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert result.found
acknowledgements = [
(parse_coap(datagram), destination)
for datagram, destination in active.sent
if parse_coap(datagram)[0] == 2
]
assert len(acknowledgements) == 3
assert all(item[0][1] == 0 and item[0][3] == b"" for item in acknowledgements)
assert all(
destination == (_TARGET, 43123) for _item, destination in acknowledgements
)
def test_collects_a_small_bounded_candidate_set(patch_socket):
response_number = 0
def responder(datagram, _destination):
nonlocal response_number
port = 40000 + response_number % 8
response_number += 1
return _response(datagram, port=port)
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=4,
)
assert result.ports == tuple(range(40000, 40008))
assert result.error_code is None
def test_fails_closed_when_too_many_distinct_ports_answer(patch_socket):
next_port = 40000
def responder(datagram, _destination):
nonlocal next_port
responses = [
*_response(datagram, port=next_port),
*_response(datagram, port=next_port + 1),
]
next_port += 2
return responses
patch_socket(responder)
result = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=4,
)
assert result.ports == ()
assert result.responses == 24
assert result.error_code == "ambiguous_response"
def test_no_response_and_interface_failures_are_fixed_results(
patch_socket, monkeypatch
):
active, _selector = patch_socket()
no_response = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
rounds=1,
)
assert no_response == OcfResponderPortDiscoveryResult(
ports=(),
attempts=3,
responses=0,
error_code="no_response",
)
assert active.closed
class BrokenSocket:
def setsockopt(self, *_args):
raise OSError("synthetic")
def close(self):
return None
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.socket.socket",
lambda *_args: BrokenSocket(),
)
unavailable = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
)
assert unavailable.error_code == "interface_unavailable"
assert unavailable.attempts == 0
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.selectors.DefaultSelector",
lambda: (_ for _ in ()).throw(OSError("synthetic")),
)
unavailable = discover_ocf_responder_ports(
_TARGET,
interface_address=_INTERFACE,
)
assert unavailable.error_code == "interface_unavailable"
assert unavailable.attempts == 0
@pytest.mark.parametrize(
("kwargs", "exception"),
[
({"target_address": 123}, TypeError),
({"target_address": "not-an-address"}, ValueError),
({"target_address": _OCF_MULTICAST_GROUP}, ValueError),
({"interface_address": "0.0.0.0"}, ValueError),
({"discovery_port": True}, TypeError),
({"discovery_port": 0}, ValueError),
({"timeout": float("nan")}, ValueError),
({"rounds": 0}, ValueError),
({"rounds": 5}, ValueError),
],
)
def test_rejects_invalid_options_without_opening_a_socket(
monkeypatch, kwargs, exception
):
values = {
"target_address": _TARGET,
"interface_address": _INTERFACE,
**kwargs,
}
socket_factory = pytest.fail
monkeypatch.setattr(
"smartthings_local.protocol.ocf_multicast.socket.socket",
socket_factory,
)
with pytest.raises(exception):
discover_ocf_responder_ports(**values)
def test_result_repr_omits_ports_and_addresses():
result = OcfResponderPortDiscoveryResult(ports=(43123,), attempts=2, responses=1)
rendered = repr(result)
assert "43123" not in rendered
assert _TARGET not in rendered
assert "port_count=1" in rendered
+86
View File
@@ -0,0 +1,86 @@
"""Oven descriptor flatten() contracts for the HA Number entity's range.
The oven reports ``x.com.samsung.da.desired = 0`` whenever no cycle is
set. That is "no setpoint", not a 0 °C target, and publishing it as one
makes Home Assistant reject every state message against the Number
entity's declared 30-270 range.
"""
from __future__ import annotations
import pytest
from mqtt_demo.samples import oven
def _links(desired, current=180):
"""A /temperatures/vs/0 link tree carrying one desired/current pair."""
return {
'/temperatures/vs/0': {
'x.com.samsung.da.items': [{
'x.com.samsung.da.current': str(current),
'x.com.samsung.da.desired': str(desired),
}],
},
}
@pytest.mark.parametrize('desired', [
oven.SETPOINT_MIN_C,
oven.SETPOINT_MIN_C + oven.SETPOINT_STEP_C,
180,
oven.SETPOINT_MAX_C,
])
def test_settable_setpoints_are_published_unchanged(desired):
assert oven.flatten(_links(desired))['target_temp_c'] == desired
@pytest.mark.parametrize('desired', [
0, # the idle oven; see module docstring
oven.SETPOINT_MIN_C - 1,
oven.SETPOINT_MAX_C + 1,
])
def test_unsettable_setpoints_are_published_as_absent(desired):
assert oven.flatten(_links(desired))['target_temp_c'] is None
def test_out_of_range_setpoint_does_not_suppress_current_temperature():
"""The guard applies to the setpoint alone. A cooling oven still
reports its cavity temperature after the cycle ends."""
sensors = oven.flatten(_links(0, current=210))
assert sensors['target_temp_c'] is None
assert sensors['current_temp_c'] == 210
def test_missing_temperature_resource_leaves_both_absent():
sensors = oven.flatten({})
assert sensors['target_temp_c'] is None
assert sensors['current_temp_c'] is None
def test_every_committed_write_is_a_value_flatten_will_publish():
"""The write path snaps to the step grid *before* bounds-checking, so
it accepts more than flatten() publishes: 29 commits as 30, and 271 as
270. That is fine for a slider, but it means the two range checks are
not symmetric. What has to hold is the weaker invariant: any setpoint
the oven is actually told to adopt is one flatten() will show back,
otherwise a write appears to succeed and then reads as unknown."""
handler = oven.command_handlers()[oven.CMD_SETPOINT]
for requested in range(-20, oven.SETPOINT_MAX_C + 40):
write = handler(str(requested), _links(180))
if write is None:
continue
_path, body = write
committed = int(body['x.com.samsung.da.items'][0][
'x.com.samsung.da.desired'])
assert oven.flatten(_links(committed))['target_temp_c'] == committed
def test_zero_is_rejected_on_the_write_path_too():
"""0 is the one value that neither snaps into range nor publishes."""
handler = oven.command_handlers()[oven.CMD_SETPOINT]
assert handler('0', _links(180)) is None
+144
View File
@@ -0,0 +1,144 @@
"""IoTivity manufacturer-certificate OwnerPSK derivation contracts."""
from __future__ import annotations
import pytest
from smartthings_local.protocol.owner_psk import (
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS,
STANDARD_MFG_CERTIFICATE_OXM_LABEL,
derive_mfg_certificate_owner_psk,
)
_VALID_INPUTS = {
"master_secret": bytes(range(48)),
"client_random": bytes(range(32)),
"server_random": bytes(range(32, 64)),
"owner_uuid": bytes.fromhex("00112233445566778899aabbccddeeff"),
"device_uuid": bytes.fromhex("ffeeddccbbaa99887766554433221100"),
"cipher_name": "ECDHE-ECDSA-AES128-GCM-SHA256",
"oxm_label": CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
}
# These synthetic expected values were generated from the fixed inputs above;
# they were not captured from IoTivity or a device. They lock deterministic
# output for the selected mappings, while the table test below states the
# key-block-length contract explicitly.
def test_fixed_synthetic_gcm_regression_vector():
assert derive_mfg_certificate_owner_psk(**_VALID_INPUTS).hex() == (
"ccd6c618a91290dee8c106544ed79a33"
)
def test_owner_then_device_uuid_order_matches_iotivity_callers():
reversed_context = derive_mfg_certificate_owner_psk(
**{
**_VALID_INPUTS,
"owner_uuid": _VALID_INPUTS["device_uuid"],
"device_uuid": _VALID_INPUTS["owner_uuid"],
}
)
assert reversed_context.hex() == "8f0f2416c483546dc1806db769b21b68"
assert reversed_context != derive_mfg_certificate_owner_psk(**_VALID_INPUTS)
def test_fixed_synthetic_ccm8_regression_vector():
inputs = {
**_VALID_INPUTS,
"cipher_name": "ECDHE-ECDSA-AES128-CCM8",
}
assert derive_mfg_certificate_owner_psk(**inputs).hex() == (
"ddd3d945e266ee3dc27ff3a2c4321d32"
)
def test_standard_and_confirmed_labels_derive_distinct_keys():
confirmed = derive_mfg_certificate_owner_psk(**_VALID_INPUTS)
standard = derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "oxm_label": STANDARD_MFG_CERTIFICATE_OXM_LABEL}
)
assert standard.hex() == "26ee1fe4c3e74509a2f5db5ab41b1e47"
assert standard != confirmed
def test_iotivity_cipher_key_block_lengths_are_immutable():
assert dict(MFG_CERTIFICATE_KEY_BLOCK_LENGTHS) == {
"ECDHE-ECDSA-AES128-SHA256": 96,
"ECDHE-ECDSA-AES128-CCM": 40,
"ECDHE-ECDSA-AES128-CCM8": 40,
"ECDHE-ECDSA-AES128-GCM-SHA256": 120,
"AES256-SHA256": 128,
"ECDHE-ECDSA-AES256-SHA384": 160,
"ECDHE-ECDSA-AES256-GCM-SHA384": 184,
"AES128-GCM-SHA256": 120,
}
with pytest.raises(TypeError):
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS["new-cipher"] = 1
@pytest.mark.parametrize(
("field", "length"),
[
("master_secret", 48),
("client_random", 32),
("server_random", 32),
("owner_uuid", 16),
("device_uuid", 16),
],
)
def test_binary_inputs_require_exact_bytes_and_lengths(field, length):
with pytest.raises(TypeError, match=f"{field} must be bytes"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, field: bytearray(length)}
)
for invalid_length in (length - 1, length + 1):
with pytest.raises(ValueError, match=f"exactly {length} bytes"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, field: b"x" * invalid_length}
)
@pytest.mark.parametrize("field", ["owner_uuid", "device_uuid"])
def test_nil_uuid_is_rejected(field):
with pytest.raises(ValueError, match="must not be the nil UUID"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, field: bytes(16)}
)
def test_cipher_and_label_must_be_explicit_supported_values():
with pytest.raises(TypeError, match="cipher_name must be a string"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "cipher_name": b"cipher"}
)
with pytest.raises(ValueError, match="unexpected.*cipher"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "cipher_name": "ECDHE-RSA-AES128-GCM-SHA256"}
)
with pytest.raises(TypeError, match="oxm_label must be bytes"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "oxm_label": "oic.sec.doxm.mfgcert"}
)
with pytest.raises(ValueError, match="unexpected.*label"):
derive_mfg_certificate_owner_psk(
**{**_VALID_INPUTS, "oxm_label": b"unsupported"}
)
def test_failures_do_not_include_key_material():
key_material = b"private-master-secret"
with pytest.raises(ValueError) as raised:
derive_mfg_certificate_owner_psk(
**{
**_VALID_INPUTS,
"master_secret": key_material,
"cipher_name": "unsupported",
}
)
assert key_material.hex() not in str(raised.value)
assert "private-master-secret" not in str(raised.value)
+380
View File
@@ -0,0 +1,380 @@
from __future__ import annotations
import gc
import traceback
import weakref
from concurrent.futures import ThreadPoolExecutor
from dataclasses import asdict
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionError
from smartthings_local.protocol import auth as auth_module
from smartthings_local.protocol import dtls_session as session_module
from smartthings_local.protocol.auth import PskAuth
from smartthings_local.protocol.dtls_session import DtlsCoapSession
_IDENTITY = b"i" * 16
_KEY = b"k" * 16
_OTHER_IDENTITY = b"j" * 16
_OTHER_KEY = b"l" * 32
class _BytesSubclass(bytes):
pass
def _fake_openssl_util(setter):
return SimpleNamespace(
ffi=auth_module._util.ffi,
lib=SimpleNamespace(SSL_CTX_set_psk_client_callback=setter),
)
def _invoke_callback(callback, identity_size: int, key_size: int):
ffi = auth_module._util.ffi
identity_buffer = ffi.new("char[]", max(identity_size, 1))
key_buffer = ffi.new("unsigned char[]", max(key_size, 1))
copied = callback(
ffi.NULL,
ffi.NULL,
identity_buffer,
identity_size,
key_buffer,
key_size,
)
return (
copied,
bytes(ffi.buffer(identity_buffer, max(identity_size, 1))),
bytes(ffi.buffer(key_buffer, max(key_size, 1))),
)
@pytest.mark.parametrize("key_length", [16, 32])
def test_psk_auth_accepts_exact_supported_credential_lengths(key_length):
provider = PskAuth(identity=_IDENTITY, key=b"k" * key_length)
assert repr(provider) == "PskAuth()"
@pytest.mark.parametrize(
("identity", "key"),
[
("i" * 16, _KEY),
(bytearray(_IDENTITY), _KEY),
(memoryview(_IDENTITY), _KEY),
(_BytesSubclass(_IDENTITY), _KEY),
(_IDENTITY, "k" * 16),
(_IDENTITY, bytearray(_KEY)),
(_IDENTITY, memoryview(_KEY)),
(_IDENTITY, _BytesSubclass(_KEY)),
],
)
def test_psk_auth_rejects_non_bytes_credentials(identity, key):
with pytest.raises(TypeError, match="identity and key must be bytes"):
PskAuth(identity=identity, key=key)
@pytest.mark.parametrize("identity_length", [0, 15, 17])
def test_psk_auth_rejects_invalid_identity_lengths(identity_length):
with pytest.raises(ValueError, match="raw 16-byte OCF UUID"):
PskAuth(identity=b"i" * identity_length, key=_KEY)
def test_psk_auth_rejects_identity_with_nul_byte():
with pytest.raises(ValueError, match="cannot contain a NUL"):
PskAuth(identity=b"i" * 15 + b"\x00", key=_KEY)
@pytest.mark.parametrize("key_length", [0, 15, 17, 31, 33])
def test_psk_auth_rejects_invalid_key_lengths(key_length):
with pytest.raises(ValueError, match="16 or 32 bytes"):
PskAuth(identity=_IDENTITY, key=b"k" * key_length)
def test_psk_auth_is_immutable_and_has_no_public_credential_surface():
provider = PskAuth(identity=_IDENTITY, key=_KEY)
rendered = repr(provider)
assert rendered == "PskAuth()"
assert str(provider) == rendered
assert _IDENTITY.decode() not in rendered
assert _KEY.decode() not in rendered
assert not hasattr(provider, "identity")
assert not hasattr(provider, "key")
assert not hasattr(provider, "_identity")
assert not hasattr(provider, "_key")
with pytest.raises(TypeError):
vars(provider)
with pytest.raises(TypeError):
asdict(provider)
with pytest.raises(AttributeError, match="immutable"):
provider.identity = _OTHER_IDENTITY
with pytest.raises(AttributeError, match="immutable"):
del provider._callback
def test_psk_auth_identity_equality_does_not_compare_credentials():
first = PskAuth(identity=_IDENTITY, key=_KEY)
second = PskAuth(identity=_IDENTITY, key=_KEY)
assert first != second
assert len({first, second}) == 2
def test_psk_callback_copies_exact_identity_and_key():
installed = {}
def setter(context_handle, callback):
installed["context"] = context_handle
installed["callback"] = callback
context_handle = object()
context = MagicMock()
context._context = context_handle
provider = PskAuth(identity=_IDENTITY, key=_KEY)
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
provider.configure_context(context)
assert installed["context"] is context_handle
callback = installed["callback"]
copied, identity_bytes, key_bytes = _invoke_callback(callback, 17, 16)
assert copied == 16
assert identity_bytes == _IDENTITY + b"\x00"
assert key_bytes == _KEY
context.set_cipher_list.assert_called_once_with(
b"ECDHE-PSK-AES128-CBC-SHA256:@SECLEVEL=0"
)
context.load_verify_locations.assert_not_called()
context.set_verify.assert_not_called()
@pytest.mark.parametrize(
("identity_size", "key_size"),
[(16, 16), (17, 15)],
)
def test_psk_callback_rejects_short_buffers_without_partial_copy(
identity_size,
key_size,
):
installed = {}
provider = PskAuth(identity=_IDENTITY, key=_KEY)
context = MagicMock()
context._context = object()
with patch.object(
auth_module,
"_util",
_fake_openssl_util(
lambda _context, callback: installed.setdefault(
"callback", callback
)
),
):
provider.configure_context(context)
ffi = auth_module._util.ffi
identity_buffer = ffi.new("char[]", 17)
key_buffer = ffi.new("unsigned char[]", 16)
ffi.memmove(identity_buffer, b"I" * 17, 17)
ffi.memmove(key_buffer, b"K" * 16, 16)
copied = installed["callback"](
ffi.NULL,
ffi.NULL,
identity_buffer,
identity_size,
key_buffer,
key_size,
)
assert copied == 0
assert bytes(ffi.buffer(identity_buffer, 17)) == b"I" * 17
assert bytes(ffi.buffer(key_buffer, 16)) == b"K" * 16
@pytest.mark.parametrize("null_buffer", ["identity", "key"])
def test_psk_callback_rejects_null_buffers(null_buffer):
installed = {}
provider = PskAuth(identity=_IDENTITY, key=_KEY)
context = MagicMock()
context._context = object()
with patch.object(
auth_module,
"_util",
_fake_openssl_util(
lambda _context, callback: installed.setdefault(
"callback", callback
)
),
):
provider.configure_context(context)
ffi = auth_module._util.ffi
identity_buffer = ffi.new("char[]", 17)
key_buffer = ffi.new("unsigned char[]", 16)
ffi.memmove(identity_buffer, b"I" * 17, 17)
ffi.memmove(key_buffer, b"K" * 16, 16)
if null_buffer == "identity":
identity_buffer = ffi.NULL
else:
key_buffer = ffi.NULL
copied = installed["callback"](
ffi.NULL,
ffi.NULL,
identity_buffer,
17,
key_buffer,
16,
)
assert copied == 0
if null_buffer == "identity":
assert bytes(ffi.buffer(key_buffer, 16)) == b"K" * 16
else:
assert bytes(ffi.buffer(identity_buffer, 17)) == b"I" * 17
def test_psk_auth_unsupported_binding_error_contains_no_credentials():
provider = PskAuth(identity=_IDENTITY, key=_KEY)
context = MagicMock()
context._context = object()
unsupported_util = SimpleNamespace(
ffi=auth_module._util.ffi,
lib=SimpleNamespace(),
)
with (
patch.object(auth_module, "_util", unsupported_util),
pytest.raises(RuntimeError) as captured,
):
provider.configure_context(context)
rendered = (
str(captured.value)
+ repr(captured.value)
+ "".join(traceback.format_exception(captured.value))
)
assert _IDENTITY.decode() not in rendered
assert _KEY.decode() not in rendered
context.set_cipher_list.assert_not_called()
def test_psk_auth_configures_real_openssl_context():
context = SSL.Context(SSL.DTLS_METHOD)
provider = PskAuth(identity=_IDENTITY, key=_KEY)
assert provider.configure_context(context) is None
def test_distinct_psk_providers_do_not_share_callback_credentials():
callbacks = []
def setter(_context, callback):
callbacks.append(callback)
first = PskAuth(identity=_IDENTITY, key=_KEY)
second = PskAuth(identity=_OTHER_IDENTITY, key=_OTHER_KEY)
first_context = MagicMock()
first_context._context = object()
second_context = MagicMock()
second_context._context = object()
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
first.configure_context(first_context)
second.configure_context(second_context)
assert callbacks[0] is not callbacks[1]
with ThreadPoolExecutor(max_workers=2) as executor:
first_future = executor.submit(_invoke_callback, callbacks[0], 17, 16)
second_future = executor.submit(
_invoke_callback,
callbacks[1],
17,
32,
)
first_result = first_future.result()
second_result = second_future.result()
assert first_result == (16, _IDENTITY + b"\x00", _KEY)
assert second_result == (32, _OTHER_IDENTITY + b"\x00", _OTHER_KEY)
def test_session_retains_psk_callback_only_with_provider_lifetime():
callback_reference = None
def setter(_context, callback):
nonlocal callback_reference
callback_reference = weakref.ref(callback)
provider = PskAuth(identity=_IDENTITY, key=_KEY)
session = DtlsCoapSession(
"appliance.invalid",
49154,
auth=provider,
)
context = MagicMock()
context._context = object()
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
session.auth.configure_context(context)
del provider
gc.collect()
assert callback_reference is not None
assert callback_reference() is not None
assert _invoke_callback(callback_reference(), 17, 16) == (
16,
_IDENTITY + b"\x00",
_KEY,
)
del session
gc.collect()
assert callback_reference() is None
def test_session_accepts_psk_provider_without_legacy_certificate_material():
provider = PskAuth(identity=_IDENTITY, key=_KEY)
session = DtlsCoapSession("appliance.invalid", 49154, auth=provider)
assert session.auth is provider
assert session.cert_path is None
assert session.key_path is None
assert session.cert_pem is None
assert session.key_pem is None
def test_psk_handshake_rejection_does_not_expose_credentials():
provider = PskAuth(identity=_IDENTITY, key=_KEY)
session = DtlsCoapSession("appliance.invalid", 49154, auth=provider)
context = MagicMock()
context._context = object()
connection = MagicMock()
connection.do_handshake.side_effect = SSL.Error()
udp_socket = MagicMock()
endpoint = SimpleNamespace(sockaddr=("192.0.2.100", 49154))
with (
patch.object(auth_module, "_util", _fake_openssl_util(lambda *_: None)),
patch.object(session_module.SSL, "Context", return_value=context),
patch.object(session_module.SSL, "Connection", return_value=connection),
patch.object(
session_module,
"open_connected_udp_socket",
return_value=(udp_socket, endpoint),
),
pytest.raises(SessionError) as captured,
):
session.connect()
rendered = (
str(captured.value)
+ repr(captured.value)
+ "".join(traceback.format_exception(captured.value))
)
assert _IDENTITY.decode() not in rendered
assert _KEY.decode() not in rendered
udp_socket.close.assert_called_once_with()
+138 -1
View File
@@ -6,7 +6,23 @@ import inspect
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
from smartthings_local.ocf.state_cache import StateCache
from smartthings_local.protocol.dtls_session import DtlsCoapSession
from smartthings_local.protocol.auth import (
AuthenticationProvider,
CertificateAuth,
PskAuth,
SamsungServerProfile,
SamsungServerRole,
ServerCertificateAuth,
)
from smartthings_local.protocol.dtls_session import (
ConnectCancellation,
DtlsCoapSession,
)
from smartthings_local.protocol.ocf_multicast import (
OcfResponderPortDiscoveryResult,
discover_ocf_responder_ports,
)
from smartthings_local.protocol.owner_psk import derive_mfg_certificate_owner_psk
def _assert_compatible_signature(callable_object, expected: list[str]) -> None:
@@ -40,6 +56,113 @@ def test_dtls_session_constructor_keeps_file_memory_and_local_port_inputs():
"local_port",
],
)
auth_parameter = inspect.signature(DtlsCoapSession).parameters["auth"]
assert auth_parameter.kind is inspect.Parameter.KEYWORD_ONLY
assert auth_parameter.default is None
def test_known_host_multicast_discovery_has_a_bounded_explicit_interface_api():
parameters = inspect.signature(discover_ocf_responder_ports).parameters
assert list(parameters) == [
"target_address",
"interface_address",
"discovery_port",
"timeout",
"rounds",
]
assert parameters["target_address"].default is inspect.Parameter.empty
for name in ("interface_address", "discovery_port", "timeout", "rounds"):
assert parameters[name].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["interface_address"].default is inspect.Parameter.empty
assert parameters["discovery_port"].default == 5683
assert parameters["timeout"].default == 3.0
assert parameters["rounds"].default == 2
result = OcfResponderPortDiscoveryResult(
ports=(43123,),
attempts=2,
responses=1,
)
assert result.found is True
assert result.ports == (43123,)
def test_certificate_auth_is_a_public_authentication_provider():
provider = CertificateAuth.from_files("/synthetic/cert.pem", "/synthetic/key")
assert isinstance(provider, AuthenticationProvider)
for factory in (CertificateAuth.from_files, CertificateAuth.from_memory):
profile_parameter = inspect.signature(factory).parameters["server_profile"]
assert profile_parameter.kind is inspect.Parameter.KEYWORD_ONLY
assert profile_parameter.default is None
def test_samsung_server_profile_is_public_and_explicitly_bound():
parameters = inspect.signature(SamsungServerProfile.bound_device).parameters
assert list(parameters) == [
"expected_certificate_identity",
"role",
"additional_ca_pem",
]
assert (
parameters["expected_certificate_identity"].default
is inspect.Parameter.empty
)
assert parameters["role"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["role"].default is SamsungServerRole.HOME_APPLIANCE
assert parameters["additional_ca_pem"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["additional_ca_pem"].default is None
def test_server_certificate_auth_is_a_public_authentication_provider():
profile = SamsungServerProfile.bound_device(
"abababab-abab-abab-abab-abababababab",
role=SamsungServerRole.VD_DEVICE,
)
provider = ServerCertificateAuth(server_profile=profile)
assert isinstance(provider, AuthenticationProvider)
session = DtlsCoapSession("device.example", 5684, auth=provider)
assert session.auth is provider
assert session.cert_path is None
assert session.key_path is None
assert session.cert_pem is None
assert session.key_pem is None
parameters = inspect.signature(ServerCertificateAuth).parameters
assert list(parameters) == ["server_profile"]
assert parameters["server_profile"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["server_profile"].default is inspect.Parameter.empty
def test_psk_auth_is_a_public_authentication_provider():
provider = PskAuth(identity=b"i" * 16, key=b"k" * 16)
assert isinstance(provider, AuthenticationProvider)
parameters = inspect.signature(PskAuth).parameters
assert list(parameters) == ["identity", "key"]
assert all(
parameter.kind is inspect.Parameter.KEYWORD_ONLY
and parameter.default is inspect.Parameter.empty
for parameter in parameters.values()
)
def test_owner_psk_derivation_keeps_every_security_input_explicit():
parameters = inspect.signature(
derive_mfg_certificate_owner_psk
).parameters
assert list(parameters) == [
"master_secret",
"client_random",
"server_random",
"owner_uuid",
"device_uuid",
"cipher_name",
"oxm_label",
]
assert all(
parameter.kind is inspect.Parameter.KEYWORD_ONLY
and parameter.default is inspect.Parameter.empty
for parameter in parameters.values()
)
def test_dtls_session_keeps_current_consumer_methods():
@@ -56,6 +179,20 @@ def test_dtls_session_keeps_current_consumer_methods():
"subscribe",
}
assert expected <= set(dir(DtlsCoapSession))
assert "abort" not in DtlsCoapSession.__dict__
assert "quiesce_for_close" not in DtlsCoapSession.__dict__
_assert_compatible_signature(DtlsCoapSession.connect, ["self"])
connect_timeout = inspect.signature(DtlsCoapSession.connect).parameters[
"timeout"
]
assert connect_timeout.kind is inspect.Parameter.KEYWORD_ONLY
assert connect_timeout.default is None
connect_cancel = inspect.signature(DtlsCoapSession.connect).parameters[
"cancel"
]
assert connect_cancel.kind is inspect.Parameter.KEYWORD_ONLY
assert connect_cancel.default is None
assert callable(ConnectCancellation().set)
_assert_compatible_signature(
DtlsCoapSession.get,
[
+169
View File
@@ -0,0 +1,169 @@
"""Session-owned pacing for request sends."""
from __future__ import annotations
from unittest.mock import Mock
import pytest
from smartthings_local.errors import SessionClosedError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.coap import (
METHOD_GET,
METHOD_POST,
OBSERVE,
TYPE_ACK,
TYPE_CON,
build_coap,
parse_coap,
)
from smartthings_local.protocol.dtls_session import DtlsCoapSession
class _NullAuth:
def configure_context(self, _context):
return None
def _session():
session = DtlsCoapSession(
"device.example",
5684,
auth=_NullAuth(),
rate_limit_rps=1_000_000,
)
session.conn = object()
return session
def test_first_get_post_and_subscribe_are_paced_before_send():
session = _session()
order = []
requests = []
def pace():
order.append("pace")
def send(datagram):
order.append("send")
request = parse_coap(datagram)
requests.append(request)
_mtype, _code, mid, token, options, _payload = request
if any(number == OBSERVE for number, _value in options):
assert session._observe_tokens[token] == "/mode/vs/0"
return
session._dispatch_coap(
build_coap(TYPE_ACK, 0x45, mid, token, [], b"ok")
)
session.pace = pace
session._send_dgram = send
assert session.get(["device", "0"]) == (0x45, b"ok")
assert session.post(["mode", "vs", "0"], b"payload") == (0x45, b"ok")
observe_token = session.subscribe(["mode", "vs", "0"])
assert session._observe_tokens[observe_token] == "/mode/vs/0"
assert order == ["pace", "send", "pace", "send", "pace", "send"]
assert [request[1] for request in requests] == [
METHOD_GET,
METHOD_POST,
METHOD_GET,
]
def test_every_subscribe_in_registration_burst_honors_rate_limit(monkeypatch):
session = _session()
now = [100.0]
waits = []
sends = []
class StopEvent:
def wait(self, delay):
waits.append(delay)
now[0] += delay
def send(datagram):
sends.append(datagram)
session._last_send_ts = now[0]
monkeypatch.setattr(dtls_session.time, "monotonic", lambda: now[0])
session._stop = StopEvent()
session._min_req_interval = 0.2
session._last_send_ts = 0.0
session._send_dgram = send
for index in range(11):
session.subscribe(["resource", "vs", str(index)])
assert len(sends) == 11
assert waits == [pytest.approx(0.2)] * 10
def test_existing_caller_pacing_before_subscribe_does_not_wait_twice(
monkeypatch,
):
session = _session()
now = [100.05]
waits = []
class StopEvent:
def wait(self, delay):
waits.append(delay)
now[0] += delay
monkeypatch.setattr(dtls_session.time, "monotonic", lambda: now[0])
session._stop = StopEvent()
session._min_req_interval = 0.2
session._last_send_ts = 100.0
session._send_dgram = Mock()
session.pace()
session.subscribe(["mode", "vs", "0"])
assert waits == [pytest.approx(0.15)]
session._send_dgram.assert_called_once()
def test_subscribe_rechecks_liveness_after_pacing_before_registering():
session = _session()
session._send_dgram = Mock()
def close_during_pacing():
session.conn = None
session.pace = close_during_pacing
with pytest.raises(SessionClosedError):
session.subscribe(["mode", "vs", "0"])
assert session._observe_tokens == {}
session._send_dgram.assert_not_called()
def test_ack_ping_and_observe_deregister_are_not_paced():
session = _session()
session.pace = Mock(side_effect=AssertionError("control send was paced"))
class Connection:
def __init__(self):
self.sent = []
def send(self, datagram):
self.sent.append(datagram)
def bio_read(self, _size):
return b""
connection = Connection()
session.conn = connection
session.ping()
session._send_observe_dereg(b"\x40", ["mode", "vs", "0"])
session._dispatch_coap(
build_coap(TYPE_CON, 0x45, 0x1234, b"unknown", [], b"state")
)
session.pace.assert_not_called()
assert len(connection.sent) == 3
assert parse_coap(connection.sent[-1])[:4] == (TYPE_ACK, 0, 0x1234, b"")
+348
View File
@@ -0,0 +1,348 @@
"""Deterministic tests for bounded DTLS handshake timing."""
from __future__ import annotations
import socket
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionTimeoutError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.dtls_session import DtlsCoapSession
from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint
class _Clock:
def __init__(self):
self.now = 100.0
def monotonic(self):
return self.now
def advance(self, seconds):
self.now += seconds
class _Auth:
def __init__(self, clock=None, configure_delay=0.0):
self.clock = clock
self.configure_delay = configure_delay
def configure_context(self, _context):
if self.clock is not None:
self.clock.advance(self.configure_delay)
class _Connection:
def __init__(self, outcomes=None, outputs=None, timer=None):
self.outcomes = list(outcomes or ())
self.outputs = list(outputs or ())
self.timer = timer
self.bio_writes = []
self.timeout_calls = 0
def set_connect_state(self):
return None
def set_ciphertext_mtu(self, _mtu):
return None
def do_handshake(self):
outcome = self.outcomes.pop(0) if self.outcomes else "want-read"
if outcome == "want-read":
raise SSL.WantReadError()
if isinstance(outcome, Exception):
raise outcome
def bio_read(self, _size):
if self.outputs:
return self.outputs.pop(0)
raise SSL.WantReadError()
def bio_write(self, data):
self.bio_writes.append(data)
def DTLSv1_get_timeout(self):
return self.timer
def DTLSv1_handle_timeout(self):
self.timeout_calls += 1
class _Socket:
def __init__(self, clock, inbound=()):
self.clock = clock
self.inbound = list(inbound)
self.timeouts = []
self.sent = []
self.closed = False
def settimeout(self, timeout):
self.timeouts.append(timeout)
def send(self, data):
self.sent.append(data)
return len(data)
def recv(self, _size):
if self.inbound:
result = self.inbound.pop(0)
if isinstance(result, Exception):
self.clock.advance(self.timeouts[-1])
raise result
return result
self.clock.advance(self.timeouts[-1])
raise TimeoutError()
def close(self):
self.closed = True
def _session(auth=None):
return DtlsCoapSession(
"device.example",
5684,
auth=auth or _Auth(),
)
def _install_handshake(
monkeypatch,
clock,
*,
outcomes=(),
outputs=(),
inbound=(),
timer=None,
):
connection = _Connection(outcomes, outputs, timer)
sock = _Socket(clock, inbound)
endpoint = ResolvedUdpEndpoint(
socket.AF_INET,
("192.0.2.10", 5684),
)
open_calls = []
def open_socket(*args, **kwargs):
open_calls.append((args, kwargs))
sock.settimeout(kwargs["timeout"])
return sock, endpoint
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(
dtls_session.SSL,
"Connection",
lambda *_args: connection,
)
monkeypatch.setattr(
dtls_session,
"open_connected_udp_socket",
open_socket,
)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
monkeypatch.setattr(dtls_session.time, "sleep", clock.advance)
monkeypatch.setattr(
dtls_session.time,
"time",
lambda: pytest.fail("wall clock must not control handshake deadlines"),
)
return connection, sock, endpoint, open_calls
@pytest.mark.parametrize("timeout", (True, "1", object()))
def test_connect_timeout_type_is_explicit(timeout):
with pytest.raises(TypeError, match="number or None"):
_session().connect(timeout=timeout)
@pytest.mark.parametrize(
"timeout",
(
0,
-1,
float("nan"),
float("inf"),
float("-inf"),
10**1000,
),
)
def test_connect_timeout_must_be_positive_and_finite(timeout):
with pytest.raises(ValueError, match="positive finite"):
_session().connect(timeout=timeout)
def test_connect_timeout_caps_every_blocking_poll(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, open_calls = _install_handshake(
monkeypatch,
clock,
)
session = _session()
with pytest.raises(SessionTimeoutError):
session.connect(timeout=4.75)
assert clock.now == pytest.approx(104.75)
assert sock.closed
assert open_calls == [
(
("device.example", 5684),
{
"family": socket.AF_UNSPEC,
"local_port": None,
"timeout": 0.5,
},
)
]
assert max(sock.timeouts) <= 0.5
assert sock.timeouts[-1] == pytest.approx(0.25)
def test_short_timeout_is_not_rounded_up_to_poll_interval(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, open_calls = _install_handshake(
monkeypatch,
clock,
)
with pytest.raises(SessionTimeoutError):
_session().connect(timeout=0.125)
assert clock.now == pytest.approx(100.125)
assert open_calls[0][1]["timeout"] == pytest.approx(0.125)
assert sock.timeouts == pytest.approx([0.125, 0.125])
def test_default_timeout_uses_session_constant(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
)
session = _session()
session.HANDSHAKE_TIMEOUT_S = 0.2
with pytest.raises(SessionTimeoutError):
session.connect()
assert clock.now == pytest.approx(100.2)
assert sock.closed
def test_context_setup_consumes_the_same_deadline(monkeypatch):
clock = _Clock()
socket_opened = False
def open_socket(*_args, **_kwargs):
nonlocal socket_opened
socket_opened = True
raise AssertionError("expired setup must not open a socket")
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(dtls_session.SSL, "Connection", lambda *_args: _Connection())
monkeypatch.setattr(dtls_session, "open_connected_udp_socket", open_socket)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
session = _session(_Auth(clock, configure_delay=0.2))
with pytest.raises(SessionTimeoutError):
session.connect(timeout=0.1)
assert not socket_opened
def test_socket_setup_consumes_the_same_deadline(monkeypatch):
clock = _Clock()
connection = _Connection()
connection.do_handshake = lambda: pytest.fail(
"expired socket setup must not start a handshake"
)
sock = _Socket(clock)
endpoint = ResolvedUdpEndpoint(
socket.AF_INET,
("192.0.2.10", 5684),
)
def open_socket(*_args, **kwargs):
sock.settimeout(kwargs["timeout"])
clock.advance(0.2)
return sock, endpoint
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(
dtls_session.SSL,
"Connection",
lambda *_args: connection,
)
monkeypatch.setattr(dtls_session, "open_connected_udp_socket", open_socket)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
with pytest.raises(SessionTimeoutError):
_session().connect(timeout=0.1)
assert sock.closed
def test_successful_handshake_preserves_connected_session_state(monkeypatch):
clock = _Clock()
connection, sock, endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
outcomes=("want-read", "success"),
inbound=(b"synthetic server flight",),
)
session = _session()
session.connect(timeout=1.0)
assert connection.bio_writes == [b"synthetic server flight"]
assert session.conn is connection
assert session.sock is sock
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert not sock.closed
def test_connect_services_openssl_retransmit_timer(monkeypatch):
clock = _Clock()
outbound = b"\x16\xfe\xfd" + b"\x00" * 8 + b"\x00\x01x"
connection, sock, _endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
outcomes=("want-read", "want-read", "success"),
outputs=(outbound, outbound),
inbound=(TimeoutError(), b"synthetic server flight"),
timer=0.0,
)
_session().connect(timeout=2.0)
assert connection.timeout_calls == 1
assert sock.sent == [outbound, outbound]
assert connection.bio_writes == [b"synthetic server flight"]
@pytest.mark.parametrize("success_delay", (0.1, 0.2))
def test_handshake_success_at_or_after_deadline_is_retained(
monkeypatch,
success_delay,
):
clock = _Clock()
connection, sock, endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
)
def late_success():
clock.advance(success_delay)
connection.do_handshake = late_success
session = _session()
session.connect(timeout=0.1)
assert session.conn is connection
assert session.sock is sock
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert not sock.closed
+316
View File
@@ -0,0 +1,316 @@
"""Deterministic tests for connection-attempt cancellation."""
from __future__ import annotations
import select
import socket
import threading
import time
import traceback
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionClosedError, SessionError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.dtls_session import (
ConnectCancellation,
DtlsCoapSession,
)
from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint
class _Auth:
def __init__(self, on_configure=None):
self.on_configure = on_configure
def configure_context(self, _context):
if self.on_configure is not None:
self.on_configure()
class _Connection:
def __init__(self, *, started=None, on_success=None, succeed=False):
self.started = started
self.on_success = on_success
self.succeed = succeed
self.bio_writes = []
self.handshake_calls = 0
def set_connect_state(self):
return None
def set_ciphertext_mtu(self, _mtu):
return None
def do_handshake(self):
self.handshake_calls += 1
if self.started is not None:
self.started.set()
if self.succeed:
if self.on_success is not None:
self.on_success()
return
raise SSL.WantReadError()
def bio_read(self, _size):
raise SSL.WantReadError()
def bio_write(self, datagram):
self.bio_writes.append(datagram)
self.succeed = True
def DTLSv1_get_timeout(self):
return None
def shutdown(self):
return None
def _session(auth=None):
return DtlsCoapSession(
"device.example",
5684,
auth=auth or _Auth(),
)
def _install_connection(monkeypatch, connection, data_socket, *, on_open=None):
endpoint = ResolvedUdpEndpoint(
socket.AF_INET,
("192.0.2.10", 5684),
)
def open_socket(*_args, **_kwargs):
if on_open is not None:
on_open()
return data_socket, endpoint
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(
dtls_session.SSL,
"Connection",
lambda *_args: connection,
)
monkeypatch.setattr(
dtls_session,
"open_connected_udp_socket",
open_socket,
)
return endpoint
def _run_connect(session, cancel):
outcome = {}
def worker():
try:
session.connect(timeout=2.0, cancel=cancel)
except Exception as error: # noqa: BLE001 - captured for assertion
outcome["error"] = error
else:
outcome["connected"] = True
thread = threading.Thread(target=worker)
thread.start()
return thread, outcome
@pytest.mark.parametrize("cancel", (True, threading.Event(), object(), "signal"))
def test_connect_cancel_type_is_explicit(cancel):
with pytest.raises(TypeError, match="ConnectCancellation or None"):
_session().connect(cancel=cancel)
def test_pre_cancelled_connect_stops_before_context_setup(monkeypatch):
cancel = ConnectCancellation()
cancel.set()
monkeypatch.setattr(
dtls_session.SSL,
"Context",
lambda *_args: pytest.fail("cancelled connect configured TLS"),
)
with pytest.raises(SessionClosedError):
_session().connect(cancel=cancel)
def test_cancel_during_context_setup_stops_before_socket_setup(monkeypatch):
cancel = ConnectCancellation()
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(
dtls_session,
"open_connected_udp_socket",
lambda *_args, **_kwargs: pytest.fail(
"cancelled connect opened a socket"
),
)
with pytest.raises(SessionClosedError):
_session(_Auth(cancel.set)).connect(cancel=cancel)
def test_cancel_during_socket_setup_closes_before_handshake(monkeypatch):
cancel = ConnectCancellation()
connection = _Connection(succeed=True)
data_socket, peer = socket.socketpair()
_install_connection(
monkeypatch,
connection,
data_socket,
on_open=cancel.set,
)
try:
with pytest.raises(SessionClosedError):
_session().connect(cancel=cancel)
assert data_socket.fileno() == -1
assert connection.handshake_calls == 0
finally:
peer.close()
def test_socket_signal_wakes_every_subscribed_waiter():
cancel = ConnectCancellation()
first = cancel._subscribe()
second = cancel._subscribe()
try:
cancel.set()
readable, _, _ = select.select(
(first[0], second[0]),
(),
(),
0,
)
assert set(readable) == {first[0], second[0]}
assert cancel._unsubscribe(*first)
assert cancel._unsubscribe(*second)
finally:
for reader, writer in (first, second):
reader.close()
writer.close()
def test_cancel_wakes_blocked_connect_without_poll_latency(monkeypatch):
cancel = ConnectCancellation()
started = threading.Event()
connection = _Connection(started=started)
data_socket, peer = socket.socketpair()
_install_connection(monkeypatch, connection, data_socket)
session = _session()
thread, outcome = _run_connect(session, cancel)
try:
assert started.wait(1.0)
before = time.monotonic()
cancel.set()
thread.join(1.0)
elapsed = time.monotonic() - before
assert not thread.is_alive()
assert elapsed < 0.25
assert isinstance(outcome.get("error"), SessionClosedError)
assert data_socket.fileno() == -1
assert session.sock is None
assert session.conn is None
assert not cancel._writers
finally:
cancel.set()
thread.join(1.0)
peer.close()
def test_reported_handshake_success_wins_cancel_during_unsubscribe(
monkeypatch,
):
class CancelDuringUnsubscribe(ConnectCancellation):
def _unsubscribe(self, reader, writer):
self.set()
return super()._unsubscribe(reader, writer)
cancel = CancelDuringUnsubscribe()
connection = _Connection(succeed=True)
data_socket, peer = socket.socketpair()
endpoint = _install_connection(monkeypatch, connection, data_socket)
session = _session()
try:
session.connect(cancel=cancel)
assert cancel.is_set()
assert session.sock is data_socket
assert session.conn is connection
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert not cancel._writers
finally:
session.close()
peer.close()
def test_cancel_during_backend_failure_does_not_read_unset_completion(
monkeypatch,
):
cancel = ConnectCancellation()
connection = _Connection()
def fail_after_cancel():
cancel.set()
raise SSL.Error("synthetic backend failure")
connection.do_handshake = fail_after_cancel
data_socket, peer = socket.socketpair()
_install_connection(monkeypatch, connection, data_socket)
session = _session()
try:
with pytest.raises(SessionClosedError):
session.connect(cancel=cancel)
assert data_socket.fileno() == -1
assert session.sock is None
assert session.conn is None
assert not cancel._writers
finally:
peer.close()
def test_successful_connect_does_not_set_cancel_or_close_session(monkeypatch):
cancel = ConnectCancellation()
connection = _Connection()
data_socket, peer = socket.socketpair()
endpoint = _install_connection(monkeypatch, connection, data_socket)
peer.send(b"synthetic server flight")
session = _session()
try:
session.connect(cancel=cancel)
assert not cancel.is_set()
assert session.sock is data_socket
assert session.conn is connection
assert session.endpoint is endpoint
assert connection.bio_writes == [b"synthetic server flight"]
assert not cancel._writers
finally:
session.close()
peer.close()
def test_cancellation_socket_failure_is_redacted_and_closes_udp(monkeypatch):
class FailingCancellation(ConnectCancellation):
def _subscribe(self):
raise OSError("credential-value at device.example")
connection = _Connection()
data_socket, peer = socket.socketpair()
_install_connection(monkeypatch, connection, data_socket)
try:
with pytest.raises(SessionError) as exc:
_session().connect(cancel=FailingCancellation())
formatted = "".join(traceback.format_exception(exc.value))
assert data_socket.fileno() == -1
assert exc.value.__context__ is None
assert "credential-value" not in formatted
assert "device.example" not in formatted
finally:
peer.close()
+240
View File
@@ -3,7 +3,12 @@
from __future__ import annotations
import threading
from types import SimpleNamespace
import pytest
from mqtt_demo import bridge as bridge_module
from mqtt_demo.bridge import PushBridge
from smartthings_local.ocf.keepalive import KeepaliveTask
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
from smartthings_local.ocf.poll_scheduler import PollScheduler, PollTier
@@ -73,3 +78,238 @@ def test_poll_scheduler_worker_stops_without_leaking_thread():
[PollTier("idle", interval_s=3600.0, paths=())],
)
_assert_worker_stops(scheduler.run_forever, "test-poll-scheduler")
class _SessionWorker:
def __init__(self, *args, **kwargs):
self.stop = None
self.started = threading.Event()
self.exited = threading.Event()
self.on_reachable = kwargs.get("on_reachable")
self.on_unreachable = kwargs.get("on_unreachable")
self.last_success_ts = 0.0
def run_forever(self, stop):
self.stop = stop
self.started.set()
stop.wait()
self.exited.set()
class _JoinedSession:
def __init__(self, workers, start_index, error=None):
self.workers = workers
self.start_index = start_index
self.error = error
def join(self):
assert all(
worker.started.wait(_THREAD_DEADLINE_S)
for worker in self.workers[self.start_index:]
)
if self.error is not None:
raise self.error
class _BlockingJoinedSession(_JoinedSession):
def __init__(self, workers):
super().__init__(workers, 0)
self.joined = threading.Event()
self.release = threading.Event()
def join(self):
super().join()
self.joined.set()
assert self.release.wait(_THREAD_DEADLINE_S)
def _bridge():
bridge = object.__new__(PushBridge)
bridge.descriptor = SimpleNamespace(
observe_paths=(),
poll_tiers=[],
is_active=lambda _state: False,
)
bridge.shared = SimpleNamespace(PING_INTERVAL_S=3600.0)
bridge.app = SimpleNamespace(klass="test")
bridge.log = SimpleNamespace(info=lambda *args: None, warning=lambda *args: None)
bridge.cache = SimpleNamespace(links={})
bridge.stop = threading.Event()
bridge._session_stop_lock = threading.Lock()
bridge._session_stop = None
bridge.scheduler = None
bridge.keepalive = None
bridge.observe_refresh = None
bridge._seed_from_device0 = lambda _session: None
bridge._retag_logger_with_serial = lambda: None
bridge.maybe_publish_state = lambda **kwargs: None
bridge.set_availability = lambda _online: None
return bridge
@pytest.mark.parametrize(
("session_count", "join_error"),
[
pytest.param(2, None, id="reconnect"),
pytest.param(1, RuntimeError("reader failed"), id="reader-error"),
],
)
def test_bridge_retires_session_workers_before_returning(
monkeypatch, session_count, join_error
):
workers = []
def worker_factory(*args, **kwargs):
worker = _SessionWorker(*args, **kwargs)
workers.append(worker)
return worker
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
bridge = _bridge()
session_stops = []
for session_index in range(session_count):
start_index = len(workers)
session = _JoinedSession(
workers,
start_index,
error=join_error if session_index == session_count - 1 else None,
)
if session.error is None:
bridge._run_session_inner(session)
else:
with pytest.raises(RuntimeError, match="reader failed"):
bridge._run_session_inner(session)
session_workers = workers[start_index:]
assert len(session_workers) == 3
assert len({id(worker.stop) for worker in session_workers}) == 1
session_stops.append(session_workers[0].stop)
assert len(workers) == session_count * 3
assert len({id(stop) for stop in session_stops}) == session_count
assert all(stop is not bridge.stop for stop in session_stops)
assert all(stop.is_set() for stop in session_stops)
assert not bridge.stop.is_set()
assert all(worker.exited.is_set() for worker in workers)
keepalive_workers = tuple(
workers[index] for index in range(1, len(workers), 3)
)
assert all(worker.on_reachable is None for worker in keepalive_workers)
assert all(worker.on_unreachable is None for worker in keepalive_workers)
assert bridge.scheduler is None
assert bridge.keepalive is None
assert bridge.observe_refresh is None
assert bridge._session_stop is None
def test_bridge_retires_started_worker_when_later_thread_fails_to_start(
monkeypatch,
):
workers = []
def worker_factory(*args, **kwargs):
worker = _SessionWorker(*args, **kwargs)
workers.append(worker)
return worker
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
bridge = _bridge()
original_start = threading.Thread.start
start_count = 0
def fail_second_start(thread):
nonlocal start_count
start_count += 1
if start_count == 2:
raise RuntimeError("synthetic thread start failure")
original_start(thread)
monkeypatch.setattr(threading.Thread, "start", fail_second_start)
with pytest.raises(RuntimeError, match="synthetic thread start failure"):
bridge._run_session_inner(SimpleNamespace(join=lambda: None))
assert len(workers) == 3
assert workers[0].started.wait(_THREAD_DEADLINE_S)
assert workers[0].exited.wait(_THREAD_DEADLINE_S)
assert workers[0].stop is not bridge.stop
assert workers[0].stop.is_set()
assert workers[1].stop is None
assert workers[2].stop is None
assert bridge.scheduler is None
assert bridge.keepalive is None
assert bridge.observe_refresh is None
assert bridge._session_stop is None
def test_request_stop_wakes_current_session_workers(monkeypatch):
workers = []
def worker_factory(*args, **kwargs):
worker = _SessionWorker(*args, **kwargs)
workers.append(worker)
return worker
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
bridge = _bridge()
session = _BlockingJoinedSession(workers)
session_thread = threading.Thread(
target=bridge._run_session_inner,
args=(session,),
daemon=True,
)
session_thread.start()
try:
assert session.joined.wait(_THREAD_DEADLINE_S)
session_stop = bridge._session_stop
assert session_stop is not None
assert not session_stop.is_set()
bridge.request_stop()
assert bridge.stop.is_set()
assert session_stop.is_set()
assert all(
worker.exited.wait(_THREAD_DEADLINE_S) for worker in workers
)
assert session_thread.is_alive()
finally:
session.release.set()
session_thread.join(_THREAD_DEADLINE_S)
assert not session_thread.is_alive()
assert bridge._session_stop is None
def test_session_workers_observe_stop_requested_before_handoff(monkeypatch):
workers = []
def worker_factory(*args, **kwargs):
worker = _SessionWorker(*args, **kwargs)
workers.append(worker)
return worker
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
bridge = _bridge()
bridge.request_stop()
bridge._run_session_inner(_JoinedSession(workers, 0))
assert len(workers) == 3
assert all(worker.stop.is_set() for worker in workers)
assert all(worker.exited.wait(_THREAD_DEADLINE_S) for worker in workers)
assert bridge._session_stop is None
+2
View File
@@ -59,6 +59,7 @@ def check_wheel(path: Path) -> None:
f"{dist_info}/METADATA",
f"{dist_info}/WHEEL",
f"{dist_info}/licenses/LICENSE",
f"{dist_info}/licenses/NOTICE",
f"{dist_info}/RECORD",
}
if metadata != expected_metadata:
@@ -83,6 +84,7 @@ def check_sdist(path: Path) -> None:
relative = {name[len(root) + 1 :] for name in names if name.startswith(f"{root}/")}
required = _tracked_files() | {
"LICENSE",
"NOTICE",
"PKG-INFO",
"README.md",
"pyproject.toml",