Compare commits
27
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4999e8beac | ||
|
|
b3045f5ddb | ||
|
|
0d9d13b8fb | ||
|
|
0722c552c9 | ||
|
|
a45ee4c004 | ||
|
|
512df7ff36 | ||
|
|
6b9a508fd9 | ||
|
|
ab0ec4961c | ||
|
|
6db846563a | ||
|
|
8c374e17b4 | ||
|
|
cd86424ca0 | ||
|
|
6da7e95731 | ||
|
|
a3bec9470d | ||
|
|
79493fbd48 | ||
|
|
627fcb19da | ||
|
|
31be87061a | ||
|
|
bc4465b274 | ||
|
|
e63acb759f | ||
|
|
231e88a8c6 | ||
|
|
dec84c9ad8 | ||
|
|
e9aee1c235 | ||
|
|
9ef5598813 | ||
|
|
917b0e47c5 | ||
|
|
83a5973434 | ||
|
|
3f0e437880 | ||
|
|
a44930f9df | ||
|
|
b0d51abcc8 |
@@ -48,6 +48,35 @@ sess.subscribe(["operational", "state", "vs", "0"], # OBSERVE
|
|||||||
sess.close()
|
sess.close()
|
||||||
```
|
```
|
||||||
|
|
||||||
|
`connect()` uses a 12-second monotonic DTLS handshake deadline by default. A
|
||||||
|
caller that needs a shorter bounded attempt can pass a positive finite value
|
||||||
|
without changing later reader timeouts. OpenSSL's DTLS timer schedules flight
|
||||||
|
retransmissions within that same deadline:
|
||||||
|
|
||||||
|
```python
|
||||||
|
sess.connect(timeout=4.0)
|
||||||
|
```
|
||||||
|
|
||||||
|
The deadline stops further setup, retries, and network waits. If OpenSSL
|
||||||
|
reports that the handshake completed at the deadline boundary, the completed
|
||||||
|
session is retained rather than torn down as a timeout.
|
||||||
|
|
||||||
|
Connection attempts can also use a one-way cancellation signal. The signal is
|
||||||
|
backed by a socketpair, so setting it wakes the network wait immediately while
|
||||||
|
OpenSSL retains control of DTLS retransmission timing:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from smartthings_local.protocol.dtls_session import ConnectCancellation
|
||||||
|
|
||||||
|
cancel_connect = ConnectCancellation()
|
||||||
|
# Another thread may call cancel_connect.set().
|
||||||
|
sess.connect(timeout=8.0, cancel=cancel_connect)
|
||||||
|
```
|
||||||
|
|
||||||
|
Setting the signal stops subscribed connection attempts and closes their
|
||||||
|
temporary UDP sockets. It does not alter an already established session or add
|
||||||
|
new session lifecycle methods. Interrupted attempts raise `SessionClosedError`.
|
||||||
|
|
||||||
If the cert/key are minted at runtime and never written to disk (e.g. inside
|
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:
|
an HA config flow), create the provider from memory instead:
|
||||||
|
|
||||||
@@ -56,6 +85,76 @@ auth = CertificateAuth.from_memory(cert_pem, key_pem)
|
|||||||
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
|
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` /
|
For compatibility, the existing `cert_path` / `key_path` and `cert_pem` /
|
||||||
`key_pem` session arguments remain supported without a deprecation warning.
|
`key_pem` session arguments remain supported without a deprecation warning.
|
||||||
They are routed through `CertificateAuth` internally. Do not combine `auth`
|
They are routed through `CertificateAuth` internally. Do not combine `auth`
|
||||||
@@ -76,6 +175,32 @@ the key must be exactly 16 or 32 bytes. `PskAuth` selects only
|
|||||||
or persist credentials. Ownership transfer and credential discovery are
|
or persist credentials. Ownership transfer and credential discovery are
|
||||||
outside this package.
|
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
|
### Classified errors
|
||||||
|
|
||||||
Runtime transport failures use the public types in
|
Runtime transport failures use the public types in
|
||||||
@@ -137,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
|
candidate after a family, bind, or connect failure. Resolution and socket
|
||||||
setup failures raise the redacted `EndpointError` documented above.
|
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.
|
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
|
### What the demo bridge gives you
|
||||||
@@ -155,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.
|
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.
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -170,7 +323,7 @@ nmap -Pn -sU -p 5683,5684,49152-49160 "$APPLIANCE_IP"
|
|||||||
|
|
||||||
Read the result:
|
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.
|
- **`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.**
|
- **Only `8888/tcp` open (token-based HTTPS)** → older firmware (~2018–2022). **Not supported here.**
|
||||||
|
|
||||||
@@ -509,6 +662,8 @@ smartthings_local/ The installable library — `pip install sm
|
|||||||
coap.py CoAP wire protocol: message encode/decode, token handling
|
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_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_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_root_ca.pem Samsung OCF root CA, bundled for handshake verification
|
||||||
ocf/ OCF resource + state layer (reusable)
|
ocf/ OCF resource + state layer (reusable)
|
||||||
__init__.py
|
__init__.py
|
||||||
@@ -535,7 +690,7 @@ mqtt_demo/ MQTT bridge demo (consumes smartthings_loca
|
|||||||
.env.example Template — copy to .env, fill in
|
.env.example Template — copy to .env, fill in
|
||||||
setup_cert.py One-shot cert minting script (live-fetches AC14K_M + UUID)
|
setup_cert.py One-shot cert minting script (live-fetches AC14K_M + UUID)
|
||||||
pyproject.toml Packaging — PyPI dist `smartthings-local`, hatch-vcs versioning
|
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
|
.github/workflows/publish.yml Build + PyPI Trusted Publishing on `v*` tags
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|||||||
@@ -166,7 +166,7 @@ def main():
|
|||||||
stopping.set()
|
stopping.set()
|
||||||
logger.info("shutting down…")
|
logger.info("shutting down…")
|
||||||
for b in bridges:
|
for b in bridges:
|
||||||
b.stop.set()
|
b.request_stop()
|
||||||
try: b.set_availability(False)
|
try: b.set_availability(False)
|
||||||
except Exception: pass
|
except Exception: pass
|
||||||
|
|
||||||
|
|||||||
+45
-6
@@ -87,6 +87,7 @@ OCF_STANDARD_SECURE_PORT = 5684
|
|||||||
# fixed-source-port reconnect invariant is untouched (see session_once).
|
# fixed-source-port reconnect invariant is untouched (see session_once).
|
||||||
_GATE_RETRIES = 1
|
_GATE_RETRIES = 1
|
||||||
_GATE_TIMEOUT_S = 4.0
|
_GATE_TIMEOUT_S = 4.0
|
||||||
|
_WORKER_JOIN_TIMEOUT_S = 2.0
|
||||||
|
|
||||||
|
|
||||||
class PushBridge:
|
class PushBridge:
|
||||||
@@ -123,6 +124,8 @@ class PushBridge:
|
|||||||
self.last_cycle_pub = None
|
self.last_cycle_pub = None
|
||||||
self.last_avail_pub: str | None = None
|
self.last_avail_pub: str | None = None
|
||||||
self.stop = threading.Event()
|
self.stop = threading.Event()
|
||||||
|
self._session_stop_lock = threading.Lock()
|
||||||
|
self._session_stop: threading.Event | None = None
|
||||||
self.started_ts = time.time()
|
self.started_ts = time.time()
|
||||||
self.session_started_ts = None
|
self.session_started_ts = None
|
||||||
self.last_change_ts = None
|
self.last_change_ts = None
|
||||||
@@ -175,6 +178,13 @@ class PushBridge:
|
|||||||
app.topic_prefix, shared.HA_DISCOVERY_PREFIX, app.device_name,
|
app.topic_prefix, shared.HA_DISCOVERY_PREFIX, app.device_name,
|
||||||
model=descriptor.name.title()))
|
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 ---------------------------------------------
|
# ---- cache plumbing ---------------------------------------------
|
||||||
|
|
||||||
def _on_cache_change(self, changed: bool, source: str) -> None:
|
def _on_cache_change(self, changed: bool, source: str) -> None:
|
||||||
@@ -423,25 +433,54 @@ class PushBridge:
|
|||||||
self.keepalive = keepalive
|
self.keepalive = keepalive
|
||||||
self.observe_refresh = observe_refresh
|
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(
|
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')
|
daemon=True, name=f'{self.app.klass}-poll')
|
||||||
ka_t = threading.Thread(
|
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')
|
daemon=True, name=f'{self.app.klass}-ping')
|
||||||
ref_t = threading.Thread(
|
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')
|
daemon=True, name=f'{self.app.klass}-obsref')
|
||||||
sched_t.start()
|
workers = (sched_t, ka_t, ref_t)
|
||||||
ka_t.start()
|
started_workers = []
|
||||||
ref_t.start()
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
for worker in workers:
|
||||||
|
worker.start()
|
||||||
|
started_workers.append(worker)
|
||||||
sess.join()
|
sess.join()
|
||||||
finally:
|
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.scheduler = None
|
||||||
self.keepalive = None
|
self.keepalive = None
|
||||||
self.observe_refresh = 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):
|
def _seed_from_device0(self, sess):
|
||||||
code, pl = sess.get(self.descriptor.seed_path, timeout=15.0)
|
code, pl = sess.get(self.descriptor.seed_path, timeout=15.0)
|
||||||
|
|||||||
@@ -166,6 +166,14 @@ def flatten(links):
|
|||||||
if temps_items:
|
if temps_items:
|
||||||
cur_c = _int(temps_items[0].get('x.com.samsung.da.current'))
|
cur_c = _int(temps_items[0].get('x.com.samsung.da.current'))
|
||||||
des_c = _int(temps_items[0].get('x.com.samsung.da.desired'))
|
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
|
# Door
|
||||||
doors_items = g('/doors/vs/0', 'x.com.samsung.da.items') or []
|
doors_items = g('/doors/vs/0', 'x.com.samsung.da.items') or []
|
||||||
|
|||||||
@@ -2,16 +2,33 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
import re
|
import re
|
||||||
|
import warnings
|
||||||
|
from enum import Enum
|
||||||
from os import PathLike
|
from os import PathLike
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Protocol, runtime_checkable
|
from typing import Protocol, runtime_checkable
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
from cryptography.x509.oid import ExtensionOID
|
||||||
from OpenSSL import SSL, _util, crypto
|
from OpenSSL import SSL, _util, crypto
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_OCF_ROOT_CA = str(Path(__file__).with_name("ocf_root_ca.pem"))
|
_OCF_ROOT_CA = str(Path(__file__).with_name("ocf_root_ca.pem"))
|
||||||
_DTLS_CIPHERS = b"ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0"
|
_DTLS_CIPHERS = b"ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0"
|
||||||
_DTLS_PSK_CIPHERS = b"ECDHE-PSK-AES128-CBC-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 = (
|
_PSK_CLIENT_CALLBACK_CDEF = (
|
||||||
"unsigned int (*)(SSL *, char *, char *, unsigned int, "
|
"unsigned int (*)(SSL *, char *, char *, unsigned int, "
|
||||||
"unsigned char *, unsigned int)"
|
"unsigned char *, unsigned int)"
|
||||||
@@ -53,6 +70,299 @@ class AuthenticationProvider(Protocol):
|
|||||||
"""Configure a context while this provider remains session-owned."""
|
"""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:
|
class CertificateAuth:
|
||||||
"""Certificate authentication loaded from files or in-memory PEM data.
|
"""Certificate authentication loaded from files or in-memory PEM data.
|
||||||
|
|
||||||
@@ -65,6 +375,7 @@ class CertificateAuth:
|
|||||||
"_certificate_pem",
|
"_certificate_pem",
|
||||||
"_private_key_path",
|
"_private_key_path",
|
||||||
"_private_key_pem",
|
"_private_key_pem",
|
||||||
|
"_server_profile",
|
||||||
)
|
)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -74,6 +385,7 @@ class CertificateAuth:
|
|||||||
private_key_path: str | PathLike[str] | None = None,
|
private_key_path: str | PathLike[str] | None = None,
|
||||||
certificate_pem: str | None = None,
|
certificate_pem: str | None = None,
|
||||||
private_key_pem: str | None = None,
|
private_key_pem: str | None = None,
|
||||||
|
server_profile: SamsungServerProfile | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
file_supplied = (
|
file_supplied = (
|
||||||
certificate_path is not None or private_key_path is not None
|
certificate_path is not None or private_key_path is not None
|
||||||
@@ -101,6 +413,11 @@ class CertificateAuth:
|
|||||||
"must pass either certificate_path/private_key_path or "
|
"must pass either certificate_path/private_key_path or "
|
||||||
"certificate_pem/private_key_pem"
|
"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__(
|
object.__setattr__(
|
||||||
self,
|
self,
|
||||||
"_certificate_path",
|
"_certificate_path",
|
||||||
@@ -113,6 +430,7 @@ class CertificateAuth:
|
|||||||
)
|
)
|
||||||
object.__setattr__(self, "_certificate_pem", certificate_pem)
|
object.__setattr__(self, "_certificate_pem", certificate_pem)
|
||||||
object.__setattr__(self, "_private_key_pem", private_key_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:
|
def __setattr__(self, _name: str, _value: object) -> None:
|
||||||
raise AttributeError("CertificateAuth is immutable")
|
raise AttributeError("CertificateAuth is immutable")
|
||||||
@@ -125,11 +443,14 @@ class CertificateAuth:
|
|||||||
cls,
|
cls,
|
||||||
certificate_path: str | PathLike[str],
|
certificate_path: str | PathLike[str],
|
||||||
private_key_path: str | PathLike[str],
|
private_key_path: str | PathLike[str],
|
||||||
|
*,
|
||||||
|
server_profile: SamsungServerProfile | None = None,
|
||||||
) -> CertificateAuth:
|
) -> CertificateAuth:
|
||||||
"""Create a provider backed by certificate-chain and key files."""
|
"""Create a provider backed by certificate-chain and key files."""
|
||||||
return cls(
|
return cls(
|
||||||
certificate_path=certificate_path,
|
certificate_path=certificate_path,
|
||||||
private_key_path=private_key_path,
|
private_key_path=private_key_path,
|
||||||
|
server_profile=server_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -137,11 +458,14 @@ class CertificateAuth:
|
|||||||
cls,
|
cls,
|
||||||
certificate_pem: str,
|
certificate_pem: str,
|
||||||
private_key_pem: str,
|
private_key_pem: str,
|
||||||
|
*,
|
||||||
|
server_profile: SamsungServerProfile | None = None,
|
||||||
) -> CertificateAuth:
|
) -> CertificateAuth:
|
||||||
"""Create a provider backed by an in-memory PEM chain and key."""
|
"""Create a provider backed by an in-memory PEM chain and key."""
|
||||||
return cls(
|
return cls(
|
||||||
certificate_pem=certificate_pem,
|
certificate_pem=certificate_pem,
|
||||||
private_key_pem=private_key_pem,
|
private_key_pem=private_key_pem,
|
||||||
|
server_profile=server_profile,
|
||||||
)
|
)
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
def __repr__(self) -> str:
|
||||||
@@ -150,13 +474,7 @@ class CertificateAuth:
|
|||||||
|
|
||||||
def configure_context(self, context: SSL.Context) -> None:
|
def configure_context(self, context: SSL.Context) -> None:
|
||||||
"""Apply the existing certificate authentication profile to a context."""
|
"""Apply the existing certificate authentication profile to a context."""
|
||||||
context.load_verify_locations(_OCF_ROOT_CA)
|
_configure_certificate_server(context, self._server_profile)
|
||||||
context.set_verify(SSL.VERIFY_PEER, _verify_peer)
|
|
||||||
# @SECLEVEL=0 permits SHA-1 in Samsung's server cert chain (AC14K_M
|
|
||||||
# intermediate is SHA-1 signed). This is the only channel that reaches
|
|
||||||
# the OpenSSL instance cryptography bundles; ctypes and cffi bindings
|
|
||||||
# do not expose SSL_CTX_set_security_level on this build.
|
|
||||||
context.set_cipher_list(_DTLS_CIPHERS)
|
|
||||||
if self._certificate_pem is not None:
|
if self._certificate_pem is not None:
|
||||||
_load_pem_chain(
|
_load_pem_chain(
|
||||||
context,
|
context,
|
||||||
@@ -240,7 +558,14 @@ class PskAuth:
|
|||||||
"the installed OpenSSL binding does not support DTLS PSK"
|
"the installed OpenSSL binding does not support DTLS PSK"
|
||||||
)
|
)
|
||||||
context.set_cipher_list(_DTLS_PSK_CIPHERS)
|
context.set_cipher_list(_DTLS_PSK_CIPHERS)
|
||||||
setter(context._context, self._callback) # noqa: SLF001
|
setter(context._context, self._callback)
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["AuthenticationProvider", "CertificateAuth", "PskAuth"]
|
__all__ = [
|
||||||
|
"AuthenticationProvider",
|
||||||
|
"CertificateAuth",
|
||||||
|
"PskAuth",
|
||||||
|
"SamsungServerProfile",
|
||||||
|
"SamsungServerRole",
|
||||||
|
"ServerCertificateAuth",
|
||||||
|
]
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from ..errors import MalformedMessageError
|
|||||||
URI_PATH = 11
|
URI_PATH = 11
|
||||||
URI_QUERY = 15
|
URI_QUERY = 15
|
||||||
OBSERVE = 6
|
OBSERVE = 6
|
||||||
|
ETAG = 4
|
||||||
CONTENT_FORMAT = 12
|
CONTENT_FORMAT = 12
|
||||||
ACCEPT = 17
|
ACCEPT = 17
|
||||||
BLOCK2 = 23
|
BLOCK2 = 23
|
||||||
@@ -121,6 +122,14 @@ def block_value(num, more, szx):
|
|||||||
return struct.pack('>I', v)[1:]
|
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):
|
def fmt_code(c):
|
||||||
"""0x45 → '2.05', 0x84 → '4.04'. Used in log lines."""
|
"""0x45 → '2.05', 0x84 → '4.04'. Used in log lines."""
|
||||||
return f"{c >> 5}.{c & 0x1F:02d}"
|
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
|
||||||
@@ -39,8 +39,9 @@ from dataclasses import dataclass
|
|||||||
from OpenSSL import SSL
|
from OpenSSL import SSL
|
||||||
|
|
||||||
from ..errors import ProbeError
|
from ..errors import ProbeError
|
||||||
|
from .auth import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
|
||||||
from .coap import split_dtls
|
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
|
from .endpoint import open_connected_udp_socket
|
||||||
|
|
||||||
# DTLS record content types (RFC 6347 §4.1)
|
# DTLS record content types (RFC 6347 §4.1)
|
||||||
@@ -603,73 +604,40 @@ def diagnose_dtls_handshake(
|
|||||||
started = time.monotonic()
|
started = time.monotonic()
|
||||||
deadline = started + timeout
|
deadline = started + timeout
|
||||||
seen = set()
|
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:
|
try:
|
||||||
while time.monotonic() < deadline:
|
completed = _drive_dtls_handshake(
|
||||||
try:
|
conn,
|
||||||
conn.do_handshake()
|
sock,
|
||||||
result.outcome = COMPLETED
|
deadline=deadline,
|
||||||
if result.rtt_s is None:
|
retries=retries,
|
||||||
result.rtt_s = time.monotonic() - started
|
on_datagram=record_datagram,
|
||||||
break
|
)
|
||||||
except SSL.WantReadError:
|
if completed:
|
||||||
pass
|
result.outcome = COMPLETED
|
||||||
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
|
|
||||||
|
|
||||||
if result.rtt_s is None:
|
if result.rtt_s is None:
|
||||||
result.rtt_s = time.monotonic() - started
|
result.rtt_s = time.monotonic() - started
|
||||||
result.datagrams.append(d)
|
except SSL.Error:
|
||||||
for ct, detail in classify_datagram(d):
|
# A fatal Alert lands here; record_datagram() has already classified
|
||||||
if ct == _CT_HANDSHAKE:
|
# the alert record before it is fed back into OpenSSL.
|
||||||
if detail not in seen:
|
result.error = ProbeError()
|
||||||
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 OSError:
|
except OSError:
|
||||||
result.error = ProbeError()
|
result.error = ProbeError()
|
||||||
finally:
|
finally:
|
||||||
|
|||||||
@@ -12,13 +12,21 @@ Wire-level details that matter (from local-tools/oven-findings.md §17):
|
|||||||
or interleaved one-shot / OBSERVE traffic mis-attributes.
|
or interleaved one-shot / OBSERVE traffic mis-attributes.
|
||||||
* Multi-block GET requires the SAME CoAP token across every block
|
* Multi-block GET requires the SAME CoAP token across every block
|
||||||
of the response ("token-stable Block2"). Fresh-token-per-block
|
of the response ("token-stable Block2"). Fresh-token-per-block
|
||||||
is silently dropped by the server.
|
is silently dropped by the server, and a transfer that opens at
|
||||||
|
NUM>0 under a token the server has not seen gets no reply at all.
|
||||||
|
|
||||||
Reader thread owns the UDP socket. Callers issue get()/post() and block
|
Reader thread owns the UDP socket. Callers issue get()/post() and block
|
||||||
on a per-token Event the reader signals. OBSERVE notifications are
|
on a per-token Event the reader signals. OBSERVE notifications are
|
||||||
delivered via the on_notification callback.
|
delivered via the on_notification callback.
|
||||||
|
|
||||||
|
A notification carries only the first block of a large representation
|
||||||
|
(RFC 7959 §2.6). Because a continuation cannot borrow the observation's
|
||||||
|
token (§3.4) and this server will not continue a transfer it did not
|
||||||
|
start, such a notification is withheld and the resource is re-read from
|
||||||
|
block 0 on a fresh one-shot token by a worker thread. See #39.
|
||||||
"""
|
"""
|
||||||
import errno
|
import errno
|
||||||
|
import math
|
||||||
import os
|
import os
|
||||||
import socket
|
import socket
|
||||||
import threading
|
import threading
|
||||||
@@ -34,11 +42,12 @@ from ..errors import (
|
|||||||
SessionTimeoutError,
|
SessionTimeoutError,
|
||||||
)
|
)
|
||||||
from .coap import (
|
from .coap import (
|
||||||
URI_PATH, URI_QUERY, OBSERVE, CONTENT_FORMAT, ACCEPT, BLOCK2, SIZE2,
|
URI_PATH, URI_QUERY, OBSERVE, ETAG, CONTENT_FORMAT, ACCEPT, BLOCK2, SIZE2,
|
||||||
TYPE_CON, TYPE_NON, TYPE_ACK, TYPE_RST,
|
TYPE_CON, TYPE_NON, TYPE_ACK, TYPE_RST,
|
||||||
METHOD_GET, METHOD_POST, CF_CBOR,
|
METHOD_GET, METHOD_POST, CF_CBOR,
|
||||||
OBSERVE_REGISTER, OBSERVE_DEREGISTER, BLOCK_SZX,
|
OBSERVE_REGISTER, OBSERVE_DEREGISTER, BLOCK_SZX,
|
||||||
encode_options, parse_coap, build_coap, block_value, fmt_code,
|
encode_options, parse_coap, build_coap, block_value, block_fields,
|
||||||
|
fmt_code,
|
||||||
split_dtls as _split_dtls,
|
split_dtls as _split_dtls,
|
||||||
)
|
)
|
||||||
from .auth import (
|
from .auth import (
|
||||||
@@ -48,6 +57,11 @@ from .auth import (
|
|||||||
_OCF_ROOT_CA,
|
_OCF_ROOT_CA,
|
||||||
_load_pem_chain,
|
_load_pem_chain,
|
||||||
)
|
)
|
||||||
|
from .dtls_handshake import (
|
||||||
|
_HANDSHAKE_POLL_S,
|
||||||
|
_HandshakeCancelled,
|
||||||
|
_drive_dtls_handshake,
|
||||||
|
)
|
||||||
from .endpoint import open_connected_udp_socket
|
from .endpoint import open_connected_udp_socket
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
@@ -66,6 +80,11 @@ DEBUG_BRIDGE = os.environ.get('DEBUG_BRIDGE') == '1'
|
|||||||
_BLOCK_MAX_ATTEMPTS = 3
|
_BLOCK_MAX_ATTEMPTS = 3
|
||||||
_BLOCK_ACK_TIMEOUT = 4.0
|
_BLOCK_ACK_TIMEOUT = 4.0
|
||||||
|
|
||||||
|
# How often a block wait re-checks that the reader is still alive. Short
|
||||||
|
# enough that a mid-transfer reader death fails fast instead of burning
|
||||||
|
# the whole per-block timeout, long enough to stay off the CPU.
|
||||||
|
_BLOCK_LIVENESS_POLL_S = 0.25
|
||||||
|
|
||||||
# Inter-request pacing: minimum seconds between CoAP CON sends on one session.
|
# Inter-request pacing: minimum seconds between CoAP CON sends on one session.
|
||||||
# Samsung's RT-OCF stacks drop requests when hit faster than their firmware
|
# Samsung's RT-OCF stacks drop requests when hit faster than their firmware
|
||||||
# ceiling (dryer ~14 req/s, oven ~8 req/s, dishwasher unknown). 5 req/s
|
# ceiling (dryer ~14 req/s, oven ~8 req/s, dishwasher unknown). 5 req/s
|
||||||
@@ -73,6 +92,22 @@ _BLOCK_ACK_TIMEOUT = 4.0
|
|||||||
# once the ceiling is measured empirically.
|
# once the ceiling is measured empirically.
|
||||||
_DEFAULT_RATE_LIMIT_RPS = 5.0
|
_DEFAULT_RATE_LIMIT_RPS = 5.0
|
||||||
|
|
||||||
|
# Maximum hrefs held for OBSERVE refetch at once. A notification storm
|
||||||
|
# on more resources than this is already past what the 5/s ceiling can
|
||||||
|
# drain, so the excess is dropped rather than queued indefinitely.
|
||||||
|
_MAX_PENDING_REFETCH = 16
|
||||||
|
|
||||||
|
# Timeout for one notification refetch. Generous relative to a poll:
|
||||||
|
# the resource is known large (that is why it blocked) and the worker
|
||||||
|
# is serialized, so a slow one delays only later refetches.
|
||||||
|
_REFETCH_TIMEOUT_S = 15.0
|
||||||
|
|
||||||
|
|
||||||
|
class _EtagChanged(Exception):
|
||||||
|
"""Internal: the server's ETag changed partway through a Block2
|
||||||
|
transfer, so the blocks in hand are from two different versions."""
|
||||||
|
|
||||||
|
|
||||||
# ICMP errors a connected UDP socket surfaces on the next recv. On these
|
# ICMP errors a connected UDP socket surfaces on the next recv. On these
|
||||||
# appliances they show up while the device is rebooting, while it holds an
|
# appliances they show up while the device is rebooting, while it holds an
|
||||||
# orphaned association, or across a router blip, and the next datagram
|
# orphaned association, or across a router blip, and the next datagram
|
||||||
@@ -88,6 +123,75 @@ _ADVISORY_ERRNOS = frozenset(
|
|||||||
) if value is not None
|
) if value is not None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _validate_handshake_timeout(timeout, default):
|
||||||
|
"""Return one finite, positive DTLS handshake timeout."""
|
||||||
|
value = default if timeout is None else timeout
|
||||||
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||||
|
raise TypeError('timeout must be a number or None')
|
||||||
|
try:
|
||||||
|
value = float(value)
|
||||||
|
except OverflowError:
|
||||||
|
raise ValueError(
|
||||||
|
'timeout must be a positive finite number or None') from None
|
||||||
|
if not math.isfinite(value) or value <= 0:
|
||||||
|
raise ValueError('timeout must be a positive finite number or None')
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
class ConnectCancellation:
|
||||||
|
"""One-way, socket-backed cancellation signal for ``connect()``.
|
||||||
|
|
||||||
|
Each active connection attempt receives its own wake socket. ``set()``
|
||||||
|
makes every subscribed socket readable immediately, without a polling
|
||||||
|
thread or a session-level abort API.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_is_set", "_lock", "_writers")
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._is_set = False
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
self._writers: set[socket.socket] = set()
|
||||||
|
|
||||||
|
def set(self) -> None:
|
||||||
|
"""Cancel current and future connection attempts using this signal."""
|
||||||
|
with self._lock:
|
||||||
|
if self._is_set:
|
||||||
|
return
|
||||||
|
self._is_set = True
|
||||||
|
for writer in self._writers:
|
||||||
|
try:
|
||||||
|
writer.send(b"\0")
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def is_set(self) -> bool:
|
||||||
|
"""Return whether cancellation has been requested."""
|
||||||
|
with self._lock:
|
||||||
|
return self._is_set
|
||||||
|
|
||||||
|
def _subscribe(self) -> tuple[socket.socket, socket.socket]:
|
||||||
|
reader, writer = socket.socketpair()
|
||||||
|
reader.setblocking(False)
|
||||||
|
with self._lock:
|
||||||
|
self._writers.add(writer)
|
||||||
|
if self._is_set:
|
||||||
|
writer.send(b"\0")
|
||||||
|
return reader, writer
|
||||||
|
|
||||||
|
def _unsubscribe(
|
||||||
|
self,
|
||||||
|
reader: socket.socket,
|
||||||
|
writer: socket.socket,
|
||||||
|
) -> bool:
|
||||||
|
with self._lock:
|
||||||
|
self._writers.discard(writer)
|
||||||
|
interrupted = self._is_set
|
||||||
|
reader.close()
|
||||||
|
writer.close()
|
||||||
|
return interrupted
|
||||||
|
|
||||||
|
|
||||||
class DtlsCoapSession:
|
class DtlsCoapSession:
|
||||||
"""Single sustained DTLS-CoAP session.
|
"""Single sustained DTLS-CoAP session.
|
||||||
|
|
||||||
@@ -166,6 +270,11 @@ class DtlsCoapSession:
|
|||||||
self.endpoint = None
|
self.endpoint = None
|
||||||
|
|
||||||
self._send_lock = threading.Lock()
|
self._send_lock = threading.Lock()
|
||||||
|
# Guards the MID/token counters and _pending. The refetch worker
|
||||||
|
# makes the session its own second concurrent get() caller, so
|
||||||
|
# two threads can mint tokens at once; without this they can
|
||||||
|
# collide and one transfer silently absorbs the other's blocks.
|
||||||
|
self._state_lock = threading.Lock()
|
||||||
# Randomize MID and token counter starting points so reconnects
|
# Randomize MID and token counter starting points so reconnects
|
||||||
# don't reuse identifiers from previous sessions — Samsung's
|
# don't reuse identifiers from previous sessions — Samsung's
|
||||||
# RT-OCF appears to remember observer state across DTLS
|
# RT-OCF appears to remember observer state across DTLS
|
||||||
@@ -182,6 +291,14 @@ class DtlsCoapSession:
|
|||||||
# token (bytes) → href (str)
|
# token (bytes) → href (str)
|
||||||
self._observe_tokens = {}
|
self._observe_tokens = {}
|
||||||
|
|
||||||
|
# OBSERVE refetch queue: href → sequence number of the newest
|
||||||
|
# notification that asked for it. Drained by a worker thread
|
||||||
|
# because _dispatch_coap cannot block (see _queue_refetch).
|
||||||
|
self._refetch_cond = threading.Condition()
|
||||||
|
self._refetch_pending = {}
|
||||||
|
self._refetch_seq = 0
|
||||||
|
self._refetch_thread = None
|
||||||
|
|
||||||
self._stop = threading.Event()
|
self._stop = threading.Event()
|
||||||
self._reader_thread = None
|
self._reader_thread = None
|
||||||
# Set while the reader owns the socket. Cleared when it exits for
|
# Set while the reader owns the socket. Cleared when it exits for
|
||||||
@@ -199,69 +316,106 @@ class DtlsCoapSession:
|
|||||||
|
|
||||||
# ---- lifecycle ---------------------------------------------------
|
# ---- lifecycle ---------------------------------------------------
|
||||||
|
|
||||||
def connect(self):
|
def connect(
|
||||||
"""DTLS handshake. Blocks up to HANDSHAKE_TIMEOUT_S. Raises
|
self,
|
||||||
ConnectionError / TimeoutError on failure."""
|
*,
|
||||||
|
timeout: float | None = None,
|
||||||
|
cancel: ConnectCancellation | None = None,
|
||||||
|
):
|
||||||
|
"""Perform a cancellable DTLS handshake using a monotonic deadline.
|
||||||
|
|
||||||
|
``timeout`` overrides ``HANDSHAKE_TIMEOUT_S`` for this call. OpenSSL
|
||||||
|
owns DTLS retransmission timing while every receive is capped by the
|
||||||
|
remaining budget, so wall-clock adjustments cannot change the bound.
|
||||||
|
Once OpenSSL reports completion, that completed session is retained
|
||||||
|
even if the call returns just after the deadline. A
|
||||||
|
``ConnectCancellation`` wakes the network wait immediately and does not
|
||||||
|
alter an already established session.
|
||||||
|
"""
|
||||||
|
handshake_timeout = _validate_handshake_timeout(
|
||||||
|
timeout, self.HANDSHAKE_TIMEOUT_S)
|
||||||
|
if cancel is not None and not isinstance(cancel, ConnectCancellation):
|
||||||
|
raise TypeError("cancel must be a ConnectCancellation or None")
|
||||||
|
if cancel is not None and cancel.is_set():
|
||||||
|
raise SessionClosedError()
|
||||||
|
deadline = time.monotonic() + handshake_timeout
|
||||||
ctx = SSL.Context(SSL.DTLS_METHOD)
|
ctx = SSL.Context(SSL.DTLS_METHOD)
|
||||||
self.auth.configure_context(ctx)
|
self.auth.configure_context(ctx)
|
||||||
|
if cancel is not None and cancel.is_set():
|
||||||
|
raise SessionClosedError()
|
||||||
|
|
||||||
conn = SSL.Connection(ctx, None)
|
conn = SSL.Connection(ctx, None)
|
||||||
conn.set_connect_state()
|
conn.set_connect_state()
|
||||||
conn.set_ciphertext_mtu(self.mtu)
|
conn.set_ciphertext_mtu(self.mtu)
|
||||||
|
if cancel is not None and cancel.is_set():
|
||||||
|
raise SessionClosedError()
|
||||||
|
|
||||||
|
remaining = deadline - time.monotonic()
|
||||||
|
if remaining <= 0:
|
||||||
|
raise SessionTimeoutError()
|
||||||
sock, endpoint = open_connected_udp_socket(
|
sock, endpoint = open_connected_udp_socket(
|
||||||
self.host,
|
self.host,
|
||||||
self.port,
|
self.port,
|
||||||
family=self.family,
|
family=self.family,
|
||||||
local_port=self.local_port,
|
local_port=self.local_port,
|
||||||
timeout=2.0,
|
timeout=min(_HANDSHAKE_POLL_S, remaining),
|
||||||
)
|
)
|
||||||
dest = endpoint.sockaddr
|
dest = endpoint.sockaddr
|
||||||
|
if cancel is not None and cancel.is_set():
|
||||||
|
sock.close()
|
||||||
|
raise SessionClosedError()
|
||||||
|
|
||||||
|
wake_subscription = None
|
||||||
|
subscription_failed = False
|
||||||
|
if cancel is not None:
|
||||||
|
try:
|
||||||
|
wake_subscription = cancel._subscribe()
|
||||||
|
except OSError:
|
||||||
|
subscription_failed = True
|
||||||
|
if subscription_failed:
|
||||||
|
sock.close()
|
||||||
|
raise SessionError() from OSError(
|
||||||
|
"connection cancellation setup failed"
|
||||||
|
)
|
||||||
|
|
||||||
t0 = time.time()
|
|
||||||
backend_failed = False
|
backend_failed = False
|
||||||
while time.time() - t0 < self.HANDSHAKE_TIMEOUT_S:
|
io_failed = False
|
||||||
|
cancelled = False
|
||||||
|
interrupted = False
|
||||||
|
completed = False
|
||||||
|
try:
|
||||||
try:
|
try:
|
||||||
conn.do_handshake()
|
completed = _drive_dtls_handshake(
|
||||||
break
|
conn,
|
||||||
except SSL.WantReadError:
|
sock,
|
||||||
pass
|
deadline=deadline,
|
||||||
|
wake_socket=(
|
||||||
|
wake_subscription[0]
|
||||||
|
if wake_subscription is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
except _HandshakeCancelled:
|
||||||
|
cancelled = True
|
||||||
except SSL.Error:
|
except SSL.Error:
|
||||||
sock.close()
|
|
||||||
backend_failed = True
|
backend_failed = True
|
||||||
break
|
|
||||||
send_failed = False
|
|
||||||
try:
|
|
||||||
o = conn.bio_read(65535)
|
|
||||||
if o:
|
|
||||||
for r in _split_dtls(o):
|
|
||||||
if sock.send(r) != len(r):
|
|
||||||
raise OSError('incomplete UDP send')
|
|
||||||
except SSL.WantReadError:
|
|
||||||
pass
|
|
||||||
except OSError:
|
except OSError:
|
||||||
sock.close()
|
io_failed = True
|
||||||
send_failed = True
|
finally:
|
||||||
if send_failed:
|
if wake_subscription is not None:
|
||||||
raise EndpointError() from OSError('UDP send failed')
|
interrupted = cancel._unsubscribe(*wake_subscription)
|
||||||
receive_failed = False
|
if cancelled or (interrupted and not completed):
|
||||||
try:
|
sock.close()
|
||||||
d = sock.recv(65535)
|
raise SessionClosedError()
|
||||||
if d:
|
if backend_failed:
|
||||||
conn.bio_write(d)
|
sock.close()
|
||||||
except socket.timeout:
|
raise SessionError() from ConnectionError('DTLS backend failed')
|
||||||
pass
|
if io_failed:
|
||||||
except OSError:
|
sock.close()
|
||||||
sock.close()
|
raise EndpointError() from OSError('UDP handshake I/O failed')
|
||||||
receive_failed = True
|
if not completed:
|
||||||
if receive_failed:
|
|
||||||
raise EndpointError() from OSError('UDP receive failed')
|
|
||||||
time.sleep(0.05)
|
|
||||||
else:
|
|
||||||
sock.close()
|
sock.close()
|
||||||
raise SessionTimeoutError()
|
raise SessionTimeoutError()
|
||||||
if backend_failed:
|
|
||||||
raise SessionError() from ConnectionError('DTLS backend failed')
|
|
||||||
|
|
||||||
self.sock = sock
|
self.sock = sock
|
||||||
self.conn = conn
|
self.conn = conn
|
||||||
@@ -298,6 +452,8 @@ class DtlsCoapSession:
|
|||||||
"""Block until the reader thread exits (i.e. socket dies)."""
|
"""Block until the reader thread exits (i.e. socket dies)."""
|
||||||
if self._reader_thread is not None:
|
if self._reader_thread is not None:
|
||||||
self._reader_thread.join()
|
self._reader_thread.join()
|
||||||
|
if self._refetch_thread is not None:
|
||||||
|
self._refetch_thread.join()
|
||||||
|
|
||||||
def _send_observe_dereg(self, tok, path_segs):
|
def _send_observe_dereg(self, tok, path_segs):
|
||||||
"""Send a single OBSERVE deregister GET (Observe option = 1)
|
"""Send a single OBSERVE deregister GET (Observe option = 1)
|
||||||
@@ -329,6 +485,11 @@ class DtlsCoapSession:
|
|||||||
time.sleep(0.1)
|
time.sleep(0.1)
|
||||||
|
|
||||||
self._stop.set()
|
self._stop.set()
|
||||||
|
# Wake the refetch worker so it sees _stop instead of sitting on
|
||||||
|
# its condition for up to a second after the socket is gone.
|
||||||
|
with self._refetch_cond:
|
||||||
|
self._refetch_pending.clear()
|
||||||
|
self._refetch_cond.notify_all()
|
||||||
if self.conn is not None:
|
if self.conn is not None:
|
||||||
try:
|
try:
|
||||||
self.conn.shutdown()
|
self.conn.shutdown()
|
||||||
@@ -339,10 +500,12 @@ class DtlsCoapSession:
|
|||||||
self.sock.close()
|
self.sock.close()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
for tok, (ev, container) in list(self._pending.items()):
|
with self._state_lock:
|
||||||
|
pending = list(self._pending.items())
|
||||||
|
self._pending.clear()
|
||||||
|
for tok, (ev, container) in pending:
|
||||||
container.setdefault('err', SessionClosedError())
|
container.setdefault('err', SessionClosedError())
|
||||||
ev.set()
|
ev.set()
|
||||||
self._pending.clear()
|
|
||||||
self._observe_tokens.clear()
|
self._observe_tokens.clear()
|
||||||
self.sock = None
|
self.sock = None
|
||||||
self.conn = None
|
self.conn = None
|
||||||
@@ -352,27 +515,30 @@ class DtlsCoapSession:
|
|||||||
# ---- send / receive plumbing -------------------------------------
|
# ---- send / receive plumbing -------------------------------------
|
||||||
|
|
||||||
def _next_mid(self):
|
def _next_mid(self):
|
||||||
self._mid = (self._mid + 1) & 0xFFFF
|
with self._state_lock:
|
||||||
return self._mid
|
self._mid = (self._mid + 1) & 0xFFFF
|
||||||
|
return self._mid
|
||||||
|
|
||||||
def _next_tok(self):
|
def _next_tok(self):
|
||||||
self._tok_counter = (self._tok_counter + 1) & 0xFFFFFFFF
|
with self._state_lock:
|
||||||
# 4-byte tokens — fits within tkl=8 cap with headroom and
|
self._tok_counter = (self._tok_counter + 1) & 0xFFFFFFFF
|
||||||
# avoids collisions across long-running OBSERVE subscriptions.
|
# 4-byte tokens — fits within tkl=8 cap with headroom and
|
||||||
return self._tok_counter.to_bytes(4, 'big')
|
# avoids collisions across long-running OBSERVE subscriptions.
|
||||||
|
return self._tok_counter.to_bytes(4, 'big')
|
||||||
|
|
||||||
def _next_observe_tok(self):
|
def _next_observe_tok(self):
|
||||||
# Single-byte tokens for OBSERVE registrations. Samsung
|
with self._state_lock:
|
||||||
# RT-OCF accepts these but silently drops TKL=4 OBSERVE
|
# Single-byte tokens for OBSERVE registrations. Samsung
|
||||||
# registrations. Counter is randomly seeded per session so
|
# RT-OCF accepts these but silently drops TKL=4 OBSERVE
|
||||||
# reconnects don't collide with stale observer state Samsung
|
# registrations. Counter is randomly seeded per session so
|
||||||
# may still be holding from the previous run.
|
# reconnects don't collide with stale observer state Samsung
|
||||||
self._observe_tok_counter = (self._observe_tok_counter + 1) & 0xFF
|
# may still be holding from the previous run.
|
||||||
# Avoid 0x00 — some CoAP stacks treat an all-zero token as
|
self._observe_tok_counter = (self._observe_tok_counter + 1) & 0xFF
|
||||||
# equivalent to "no token" / empty (TKL=0).
|
# Avoid 0x00 — some CoAP stacks treat an all-zero token as
|
||||||
if self._observe_tok_counter == 0:
|
# equivalent to "no token" / empty (TKL=0).
|
||||||
self._observe_tok_counter = 1
|
if self._observe_tok_counter == 0:
|
||||||
return bytes([self._observe_tok_counter])
|
self._observe_tok_counter = 1
|
||||||
|
return bytes([self._observe_tok_counter])
|
||||||
|
|
||||||
def _send_dgram(self, datagram):
|
def _send_dgram(self, datagram):
|
||||||
"""Send a CoAP datagram. Holds the send lock for the
|
"""Send a CoAP datagram. Holds the send lock for the
|
||||||
@@ -474,9 +640,15 @@ class DtlsCoapSession:
|
|||||||
# Reader no longer owns the socket — callers must fail fast.
|
# Reader no longer owns the socket — callers must fail fast.
|
||||||
self._reader_running.clear()
|
self._reader_running.clear()
|
||||||
# Make sure pending waiters don't hang if the reader dies.
|
# Make sure pending waiters don't hang if the reader dies.
|
||||||
for tok, (ev, container) in list(self._pending.items()):
|
with self._state_lock:
|
||||||
|
pending = list(self._pending.items())
|
||||||
|
for tok, (ev, container) in pending:
|
||||||
container.setdefault('err', SessionClosedError())
|
container.setdefault('err', SessionClosedError())
|
||||||
ev.set()
|
ev.set()
|
||||||
|
# Nothing will answer a refetch now either.
|
||||||
|
with self._refetch_cond:
|
||||||
|
self._refetch_pending.clear()
|
||||||
|
self._refetch_cond.notify_all()
|
||||||
|
|
||||||
def _dispatch_coap(self, datagram):
|
def _dispatch_coap(self, datagram):
|
||||||
try:
|
try:
|
||||||
@@ -506,7 +678,8 @@ class DtlsCoapSession:
|
|||||||
return
|
return
|
||||||
|
|
||||||
# Pending one-shot? Resolve and return.
|
# Pending one-shot? Resolve and return.
|
||||||
rec = self._pending.get(tok)
|
with self._state_lock:
|
||||||
|
rec = self._pending.get(tok)
|
||||||
if rec is not None:
|
if rec is not None:
|
||||||
ev, container = rec
|
ev, container = rec
|
||||||
container['code'] = code
|
container['code'] = code
|
||||||
@@ -524,6 +697,16 @@ class DtlsCoapSession:
|
|||||||
logger.warning("observe %s: non-2.05 %s",
|
logger.warning("observe %s: non-2.05 %s",
|
||||||
href, fmt_code(code))
|
href, fmt_code(code))
|
||||||
return
|
return
|
||||||
|
# RFC 7959 §2.6: a notification carries only the first block
|
||||||
|
# of the representation. Handing the callback a partial CBOR
|
||||||
|
# buffer is what #39 was about, so anything with M=1 (or a
|
||||||
|
# block past the first) goes to the refetch worker instead.
|
||||||
|
b2 = [v for n, v in ropts if n == BLOCK2]
|
||||||
|
if b2:
|
||||||
|
num, more, _ = block_fields(b2[0])
|
||||||
|
if more or num:
|
||||||
|
self._queue_refetch(href)
|
||||||
|
return
|
||||||
cb = self.on_notification
|
cb = self.on_notification
|
||||||
if cb is not None:
|
if cb is not None:
|
||||||
try:
|
try:
|
||||||
@@ -535,6 +718,110 @@ class DtlsCoapSession:
|
|||||||
|
|
||||||
# Stale token (post-reconnect or unknown) — drop quietly.
|
# Stale token (post-reconnect or unknown) — drop quietly.
|
||||||
|
|
||||||
|
# ---- OBSERVE refetch ---------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _log_refetch(msg, *args):
|
||||||
|
"""Refetch outcomes are debug-level in normal operation, which is
|
||||||
|
below the bridge's INFO default, so a healthy session stays quiet.
|
||||||
|
DEBUG_BRIDGE=1 promotes them to INFO for hardware validation:
|
||||||
|
that shows which token the re-read used and whether it completed,
|
||||||
|
without also turning on every per-block retransmit line."""
|
||||||
|
(logger.info if DEBUG_BRIDGE else logger.debug)(msg, *args)
|
||||||
|
|
||||||
|
def _queue_refetch(self, href):
|
||||||
|
"""Queue a blockwise notification for re-reading.
|
||||||
|
|
||||||
|
Called from the reader thread, so it must not block: _dispatch_coap
|
||||||
|
runs there and _blockwise_get waits on an Event only that same
|
||||||
|
thread can set, which would deadlock the session outright. Latest
|
||||||
|
wins per href — a burst of notifications on one resource collapses
|
||||||
|
into a single re-read of its final state."""
|
||||||
|
with self._refetch_cond:
|
||||||
|
if (href not in self._refetch_pending
|
||||||
|
and len(self._refetch_pending) >= _MAX_PENDING_REFETCH):
|
||||||
|
self._log_refetch(
|
||||||
|
"refetch %s dropped: queue full (%d pending)",
|
||||||
|
href, len(self._refetch_pending))
|
||||||
|
return
|
||||||
|
self._refetch_seq += 1
|
||||||
|
self._refetch_pending[href] = self._refetch_seq
|
||||||
|
self._refetch_cond.notify()
|
||||||
|
self._start_refetch_worker()
|
||||||
|
|
||||||
|
def _start_refetch_worker(self):
|
||||||
|
"""Start the refetch worker on first use. Sessions that never see
|
||||||
|
a blockwise notification never grow the thread."""
|
||||||
|
if self._refetch_thread is not None or self._stop.is_set():
|
||||||
|
return
|
||||||
|
with self._state_lock:
|
||||||
|
if self._refetch_thread is not None or self._stop.is_set():
|
||||||
|
return
|
||||||
|
self._refetch_thread = threading.Thread(
|
||||||
|
target=self._refetch_loop, daemon=True,
|
||||||
|
name=f'stl-refetch-{self.host}')
|
||||||
|
self._refetch_thread.start()
|
||||||
|
|
||||||
|
def _refetch_loop(self):
|
||||||
|
"""Re-read blockwise-notified resources, one at a time.
|
||||||
|
|
||||||
|
Serialized on purpose. Each re-read is a multi-block transfer and
|
||||||
|
_blockwise_get paces between blocks, so running one at a time is
|
||||||
|
what keeps a notification storm under the firmware's request
|
||||||
|
ceiling."""
|
||||||
|
while self._refetch_alive():
|
||||||
|
with self._refetch_cond:
|
||||||
|
while not self._refetch_pending and self._refetch_alive():
|
||||||
|
self._refetch_cond.wait(1.0)
|
||||||
|
if not self._refetch_pending:
|
||||||
|
return
|
||||||
|
href, seq = next(iter(self._refetch_pending.items()))
|
||||||
|
del self._refetch_pending[href]
|
||||||
|
self._refetch_one(href, seq)
|
||||||
|
|
||||||
|
def _refetch_alive(self):
|
||||||
|
"""False once the session is closing or the reader has died. A
|
||||||
|
refetch needs the reader to resolve its token, so outliving it
|
||||||
|
would leave join() waiting on a thread with nothing to do."""
|
||||||
|
if self._stop.is_set():
|
||||||
|
return False
|
||||||
|
return self._reader_thread is None or self._reader_running.is_set()
|
||||||
|
|
||||||
|
def _refetch_one(self, href, seq):
|
||||||
|
"""Re-read one href from block 0 and deliver it if it is still
|
||||||
|
the freshest thing we know about that resource."""
|
||||||
|
self.pace()
|
||||||
|
segs = [s for s in href.split('/') if s]
|
||||||
|
try:
|
||||||
|
code, payload, blocks, tok = self._blockwise_get(
|
||||||
|
segs, (), _REFETCH_TIMEOUT_S)
|
||||||
|
except Exception as e:
|
||||||
|
# Device silent, session gone, ETag never settled, block cap
|
||||||
|
# hit. Whatever the reason, dropping the notification is the
|
||||||
|
# contract: the poll tiers still carry freshness, and handing
|
||||||
|
# over the first block is the bug this replaced.
|
||||||
|
self._log_refetch("refetch %s failed: %s", href, e)
|
||||||
|
return
|
||||||
|
if code != 0x45:
|
||||||
|
self._log_refetch("refetch %s returned %s", href, fmt_code(code))
|
||||||
|
return
|
||||||
|
with self._refetch_cond:
|
||||||
|
# A newer notification landed while we were reading. That one
|
||||||
|
# has its own refetch queued, so this result is already stale.
|
||||||
|
if self._refetch_pending.get(href, 0) > seq:
|
||||||
|
self._log_refetch(
|
||||||
|
"refetch %s tok=%s blocks=%d bytes=%d superseded",
|
||||||
|
href, tok.hex(), blocks, len(payload))
|
||||||
|
return
|
||||||
|
self._log_refetch("refetch %s tok=%s blocks=%d bytes=%d ok",
|
||||||
|
href, tok.hex(), blocks, len(payload))
|
||||||
|
cb = self.on_notification
|
||||||
|
if cb is not None:
|
||||||
|
try:
|
||||||
|
cb(href, payload)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("notification callback %s: %s", href, e)
|
||||||
|
|
||||||
# ---- request primitives ------------------------------------------
|
# ---- request primitives ------------------------------------------
|
||||||
|
|
||||||
def get(self, path_segs, query=(), timeout=10.0):
|
def get(self, path_segs, query=(), timeout=10.0):
|
||||||
@@ -545,76 +832,183 @@ class DtlsCoapSession:
|
|||||||
token, and dropping a fresh token on block 1+ silently drops
|
token, and dropping a fresh token on block 1+ silently drops
|
||||||
the request."""
|
the request."""
|
||||||
self._check_live()
|
self._check_live()
|
||||||
|
code, blob, _blocks, _tok = self._blockwise_get(
|
||||||
|
path_segs, query, timeout)
|
||||||
|
return code, blob
|
||||||
|
|
||||||
|
def _blockwise_get(self, path_segs, query=(), timeout=10.0):
|
||||||
|
"""Shared token-stable Block2 reassembly (RFC 7959 §2.4).
|
||||||
|
|
||||||
|
Returns (code, payload, block_count, token). The last two are
|
||||||
|
diagnostics for the refetch log; get() drops them.
|
||||||
|
|
||||||
|
Mints one fresh 4-byte token and holds it across every block of
|
||||||
|
the transfer. Also the notification-refetch primitive: RFC 7959
|
||||||
|
§3.4 forbids continuing a blockwise notification on the
|
||||||
|
observation's token, and Samsung's RT-OCF drops a transfer that
|
||||||
|
opens at NUM>0 under a token it has not seen, so a truncated
|
||||||
|
notification is recovered by re-reading from block 0 through
|
||||||
|
this same path rather than by a §2.6 continuation.
|
||||||
|
|
||||||
|
Restarts once if the server's ETag changes mid-transfer, then
|
||||||
|
gives up: RFC 7959 §2.4 requires the client to compare ETags
|
||||||
|
when the server supplies them. None of the tested appliances
|
||||||
|
emit option 4, so on those this is inert."""
|
||||||
|
try:
|
||||||
|
return self._blockwise_get_once(path_segs, query, timeout)
|
||||||
|
except _EtagChanged:
|
||||||
|
logger.debug("GET %s /%s: ETag changed mid-transfer, restarting",
|
||||||
|
self.host, '/'.join(path_segs))
|
||||||
|
try:
|
||||||
|
return self._blockwise_get_once(path_segs, query, timeout)
|
||||||
|
except _EtagChanged:
|
||||||
|
logger.debug(
|
||||||
|
"GET %s /%s: representation kept changing mid-transfer",
|
||||||
|
self.host, '/'.join(path_segs))
|
||||||
|
raise BlockwiseError() from None
|
||||||
|
|
||||||
|
def _blockwise_get_once(self, path_segs, query, timeout):
|
||||||
|
"""One attempt at a full Block2 transfer. Raises _EtagChanged if
|
||||||
|
the server's representation changed while we were reassembling."""
|
||||||
tok = self._next_tok()
|
tok = self._next_tok()
|
||||||
blob = b''
|
blob = b''
|
||||||
num = 0
|
num = 0
|
||||||
|
blocks = 0
|
||||||
last_code = None
|
last_code = None
|
||||||
last_opts = []
|
etag = None
|
||||||
deadline = time.time() + timeout
|
deadline = time.time() + timeout
|
||||||
szx = BLOCK_SZX # server may negotiate down; track per-transfer
|
szx = BLOCK_SZX # server may negotiate down; track per-transfer
|
||||||
while True:
|
while True:
|
||||||
if num > 0:
|
self.pace()
|
||||||
self.pace()
|
self._check_live()
|
||||||
container = {}
|
container = self._exchange_block(
|
||||||
for attempt in range(_BLOCK_MAX_ATTEMPTS):
|
tok, path_segs, query, num, szx, deadline)
|
||||||
ev = threading.Event()
|
|
||||||
container = {}
|
|
||||||
self._pending[tok] = (ev, container)
|
|
||||||
try:
|
|
||||||
mid = self._next_mid()
|
|
||||||
opts = [(URI_PATH, s.encode()) for s in path_segs]
|
|
||||||
for q in query:
|
|
||||||
opts.append((URI_QUERY, q.encode()))
|
|
||||||
opts.append((ACCEPT, CF_CBOR))
|
|
||||||
if num > 0:
|
|
||||||
opts.append((BLOCK2, block_value(num, 0, szx)))
|
|
||||||
self._send_dgram(
|
|
||||||
build_coap(TYPE_CON, METHOD_GET, mid, tok, opts))
|
|
||||||
per_wait = min(_BLOCK_ACK_TIMEOUT,
|
|
||||||
max(0.1, deadline - time.time()))
|
|
||||||
if ev.wait(per_wait):
|
|
||||||
break # got a response
|
|
||||||
remaining = deadline - time.time()
|
|
||||||
if remaining <= 0 or attempt == _BLOCK_MAX_ATTEMPTS - 1:
|
|
||||||
logger.debug(
|
|
||||||
"GET %s /%s block %d: timed out after %d attempt(s)",
|
|
||||||
self.host, '/'.join(path_segs), num, attempt + 1,
|
|
||||||
)
|
|
||||||
raise SessionTimeoutError()
|
|
||||||
logger.debug(
|
|
||||||
"GET %s /%s block %d: attempt %d/%d timeout, retrying",
|
|
||||||
self.host, '/'.join(path_segs), num,
|
|
||||||
attempt + 1, _BLOCK_MAX_ATTEMPTS,
|
|
||||||
)
|
|
||||||
finally:
|
|
||||||
self._pending.pop(tok, None)
|
|
||||||
if 'err' in container:
|
if 'err' in container:
|
||||||
raise container['err']
|
raise container['err']
|
||||||
|
blocks += 1
|
||||||
|
|
||||||
code = container['code']
|
code = container['code']
|
||||||
payload = container['payload']
|
payload = container['payload']
|
||||||
ropts = container['options']
|
ropts = container['options']
|
||||||
last_code = code
|
last_code = code
|
||||||
last_opts = ropts
|
|
||||||
blob += payload
|
|
||||||
# 4.xx / 5.xx responses don't carry Block2 continuation —
|
# 4.xx / 5.xx responses don't carry Block2 continuation —
|
||||||
# bail with whatever we got. Caller decides if 4.xx is fatal.
|
# bail with whatever we got. Caller decides if 4.xx is fatal.
|
||||||
if code >> 5 != 2:
|
if code >> 5 != 2:
|
||||||
return code, blob
|
return code, blob, blocks, tok
|
||||||
|
|
||||||
|
# RFC 7959 §2.4: compare ETags across blocks, or we splice
|
||||||
|
# two versions of the resource into one buffer.
|
||||||
|
block_etag = next((v for n, v in ropts if n == ETAG), None)
|
||||||
|
if num == 0:
|
||||||
|
etag = block_etag
|
||||||
|
elif etag is not None and block_etag != etag:
|
||||||
|
raise _EtagChanged()
|
||||||
|
|
||||||
|
blob += payload
|
||||||
b2 = [v for n, v in ropts if n == BLOCK2]
|
b2 = [v for n, v in ropts if n == BLOCK2]
|
||||||
more = 0
|
if not b2:
|
||||||
if b2:
|
break
|
||||||
bv = int.from_bytes(b2[0], 'big')
|
_, more, server_szx = block_fields(b2[0])
|
||||||
more = (bv >> 3) & 1
|
|
||||||
server_szx = bv & 0x07
|
|
||||||
if server_szx != szx:
|
|
||||||
szx = server_szx
|
|
||||||
if not more:
|
if not more:
|
||||||
break
|
break
|
||||||
num += 1
|
if server_szx != szx:
|
||||||
|
# Server negotiated the block size down. Block numbers
|
||||||
|
# are indices into the new size, so the next one has to
|
||||||
|
# come off the byte offset we have actually accumulated,
|
||||||
|
# not off num + 1.
|
||||||
|
szx = server_szx
|
||||||
|
num = len(blob) >> (szx + 4)
|
||||||
|
else:
|
||||||
|
num += 1
|
||||||
if num > self.MAX_BLOCKS:
|
if num > self.MAX_BLOCKS:
|
||||||
raise BlockwiseError()
|
raise BlockwiseError()
|
||||||
return last_code, blob
|
return last_code, blob, blocks, tok
|
||||||
|
|
||||||
|
def _exchange_block(self, tok, path_segs, query, num, szx, deadline):
|
||||||
|
"""Send one block request under `tok` and return its response
|
||||||
|
container, retransmitting up to _BLOCK_MAX_ATTEMPTS times.
|
||||||
|
|
||||||
|
A response whose Block2 NUM is not the one we asked for is a
|
||||||
|
retransmit of an earlier block, not the next one. Concatenating
|
||||||
|
it would corrupt the buffer, so keep waiting on the same
|
||||||
|
attempt budget instead."""
|
||||||
|
for attempt in range(_BLOCK_MAX_ATTEMPTS):
|
||||||
|
ev = threading.Event()
|
||||||
|
container = {}
|
||||||
|
with self._state_lock:
|
||||||
|
self._pending[tok] = (ev, container)
|
||||||
|
try:
|
||||||
|
mid = self._next_mid()
|
||||||
|
opts = [(URI_PATH, s.encode()) for s in path_segs]
|
||||||
|
for q in query:
|
||||||
|
opts.append((URI_QUERY, q.encode()))
|
||||||
|
opts.append((ACCEPT, CF_CBOR))
|
||||||
|
if num > 0:
|
||||||
|
opts.append((BLOCK2, block_value(num, 0, szx)))
|
||||||
|
self._send_dgram(
|
||||||
|
build_coap(TYPE_CON, METHOD_GET, mid, tok, opts))
|
||||||
|
while True:
|
||||||
|
per_wait = min(_BLOCK_ACK_TIMEOUT,
|
||||||
|
max(0.1, deadline - time.time()))
|
||||||
|
if not self._wait_for_block(ev, per_wait):
|
||||||
|
break # attempt timed out
|
||||||
|
if 'err' in container or self._block_num_matches(
|
||||||
|
container, num):
|
||||||
|
return container
|
||||||
|
logger.debug(
|
||||||
|
"GET %s /%s block %d: stale block, still waiting",
|
||||||
|
self.host, '/'.join(path_segs), num)
|
||||||
|
ev.clear()
|
||||||
|
container.clear()
|
||||||
|
if deadline - time.time() <= 0:
|
||||||
|
break
|
||||||
|
remaining = deadline - time.time()
|
||||||
|
if remaining <= 0 or attempt == _BLOCK_MAX_ATTEMPTS - 1:
|
||||||
|
logger.debug(
|
||||||
|
"GET %s /%s block %d: timed out after %d attempt(s)",
|
||||||
|
self.host, '/'.join(path_segs), num, attempt + 1,
|
||||||
|
)
|
||||||
|
raise SessionTimeoutError()
|
||||||
|
logger.debug(
|
||||||
|
"GET %s /%s block %d: attempt %d/%d timeout, retrying",
|
||||||
|
self.host, '/'.join(path_segs), num,
|
||||||
|
attempt + 1, _BLOCK_MAX_ATTEMPTS,
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
with self._state_lock:
|
||||||
|
self._pending.pop(tok, None)
|
||||||
|
raise SessionTimeoutError()
|
||||||
|
|
||||||
|
def _wait_for_block(self, ev, per_wait):
|
||||||
|
"""Wait for one block response, giving up early if the reader
|
||||||
|
dies underneath us.
|
||||||
|
|
||||||
|
Only the reader thread can resolve a token, so once it is gone
|
||||||
|
the wait can never succeed. Polling in slices turns what would
|
||||||
|
be a full per-block timeout into an immediate SessionClosedError,
|
||||||
|
which is the same fail-fast contract get() gets from _check_live()
|
||||||
|
at entry — it just has to hold for every block, not only the
|
||||||
|
first."""
|
||||||
|
deadline = time.time() + per_wait
|
||||||
|
while True:
|
||||||
|
slice_s = min(_BLOCK_LIVENESS_POLL_S, deadline - time.time())
|
||||||
|
if slice_s <= 0:
|
||||||
|
return False
|
||||||
|
if ev.wait(slice_s):
|
||||||
|
return True
|
||||||
|
self._check_live()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _block_num_matches(container, num):
|
||||||
|
"""True if this response carries the block we asked for. A
|
||||||
|
response with no Block2 option is the whole representation, so
|
||||||
|
it only answers block 0."""
|
||||||
|
if container.get('code', 0) >> 5 != 2:
|
||||||
|
return True # error responses end the transfer either way
|
||||||
|
b2 = [v for n, v in container.get('options', ()) if n == BLOCK2]
|
||||||
|
if not b2:
|
||||||
|
return num == 0
|
||||||
|
return block_fields(b2[0])[0] == num
|
||||||
|
|
||||||
def post(self, path_segs, body_cbor, timeout=8.0):
|
def post(self, path_segs, body_cbor, timeout=8.0):
|
||||||
"""Single-frame POST with a CBOR-encoded body. Returns
|
"""Single-frame POST with a CBOR-encoded body. Returns
|
||||||
@@ -629,8 +1023,11 @@ class DtlsCoapSession:
|
|||||||
body_cbor)
|
body_cbor)
|
||||||
ev = threading.Event()
|
ev = threading.Event()
|
||||||
container = {}
|
container = {}
|
||||||
self._pending[tok] = (ev, container)
|
with self._state_lock:
|
||||||
|
self._pending[tok] = (ev, container)
|
||||||
try:
|
try:
|
||||||
|
self.pace()
|
||||||
|
self._check_live()
|
||||||
self._send_dgram(datagram)
|
self._send_dgram(datagram)
|
||||||
if not ev.wait(timeout):
|
if not ev.wait(timeout):
|
||||||
raise SessionTimeoutError()
|
raise SessionTimeoutError()
|
||||||
@@ -638,7 +1035,8 @@ class DtlsCoapSession:
|
|||||||
raise container['err']
|
raise container['err']
|
||||||
return container['code'], container['payload']
|
return container['code'], container['payload']
|
||||||
finally:
|
finally:
|
||||||
self._pending.pop(tok, None)
|
with self._state_lock:
|
||||||
|
self._pending.pop(tok, None)
|
||||||
|
|
||||||
def ping(self):
|
def ping(self):
|
||||||
"""RFC 7252 §4.4 CoAP Ping — empty CON, no token, no payload.
|
"""RFC 7252 §4.4 CoAP Ping — empty CON, no token, no payload.
|
||||||
@@ -693,6 +1091,8 @@ class DtlsCoapSession:
|
|||||||
Returns the token used (in case the caller wants to deregister
|
Returns the token used (in case the caller wants to deregister
|
||||||
later)."""
|
later)."""
|
||||||
self._check_live()
|
self._check_live()
|
||||||
|
self.pace()
|
||||||
|
self._check_live()
|
||||||
tok = self._next_observe_tok()
|
tok = self._next_observe_tok()
|
||||||
href = '/' + '/'.join(path_segs)
|
href = '/' + '/'.join(path_segs)
|
||||||
# Register the token BEFORE sending — otherwise the device
|
# Register the token BEFORE sending — otherwise the device
|
||||||
|
|||||||
@@ -0,0 +1,322 @@
|
|||||||
|
"""Bounded discovery of a known host's plaintext OCF response port.
|
||||||
|
|
||||||
|
Some OCF devices receive multicast discovery on UDP 5683 but reply from an
|
||||||
|
ephemeral port. This module records only token-correlated response ports from
|
||||||
|
the caller's expected IPv4 address. The results are candidates: callers still
|
||||||
|
need directory parsing, a DTLS probe, and authenticated identity validation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import ipaddress
|
||||||
|
import math
|
||||||
|
import secrets
|
||||||
|
import selectors
|
||||||
|
import socket
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
from ..errors import MalformedMessageError
|
||||||
|
from .coap import (
|
||||||
|
ACCEPT,
|
||||||
|
CF_CBOR,
|
||||||
|
METHOD_GET,
|
||||||
|
TYPE_ACK,
|
||||||
|
TYPE_CON,
|
||||||
|
TYPE_NON,
|
||||||
|
URI_PATH,
|
||||||
|
URI_QUERY,
|
||||||
|
build_coap,
|
||||||
|
parse_coap,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"OcfResponderPortDiscoveryResult",
|
||||||
|
"discover_ocf_responder_ports",
|
||||||
|
]
|
||||||
|
|
||||||
|
_OCF_MULTICAST_GROUP = socket.inet_ntoa(bytes((224, 0, 1, 187)))
|
||||||
|
_OCF_DISCOVERY_PORT = 5683
|
||||||
|
_OCF_CBOR = (10_000).to_bytes(2, "big")
|
||||||
|
_OCF_CONTENT_FORMAT_VERSION = 2049
|
||||||
|
_OCF_VERSION_1_0 = (2048).to_bytes(2, "big")
|
||||||
|
_CONTENT = 0x45
|
||||||
|
_MAX_DATAGRAM_BYTES = 8192
|
||||||
|
_MAX_DATAGRAMS_PER_ROUND = 64
|
||||||
|
_MAX_PORTS = 8
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True, repr=False)
|
||||||
|
class OcfResponderPortDiscoveryResult:
|
||||||
|
"""Redacted result of one known-host multicast discovery operation."""
|
||||||
|
|
||||||
|
ports: tuple[int, ...]
|
||||||
|
attempts: int
|
||||||
|
responses: int
|
||||||
|
error_code: str | None = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def found(self) -> bool:
|
||||||
|
"""Return whether at least one response port was discovered."""
|
||||||
|
return bool(self.ports)
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
return (
|
||||||
|
"OcfResponderPortDiscoveryResult("
|
||||||
|
f"found={self.found!r}, port_count={len(self.ports)}, "
|
||||||
|
f"attempts={self.attempts}, responses={self.responses}, "
|
||||||
|
f"error_code={self.error_code!r})"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_address(value: object, name: str) -> tuple[str, bytes]:
|
||||||
|
if not isinstance(value, str):
|
||||||
|
raise TypeError(f"{name} must be an IPv4 address string")
|
||||||
|
try:
|
||||||
|
address = ipaddress.IPv4Address(value)
|
||||||
|
except ipaddress.AddressValueError as exc:
|
||||||
|
raise ValueError(f"{name} must be a valid IPv4 address") from exc
|
||||||
|
if address.is_multicast or address.is_unspecified or address.is_reserved:
|
||||||
|
raise ValueError(f"{name} must be a unicast IPv4 address")
|
||||||
|
return str(address), address.packed
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_options(
|
||||||
|
*,
|
||||||
|
discovery_port: object,
|
||||||
|
timeout: object,
|
||||||
|
rounds: object,
|
||||||
|
) -> tuple[int, float, int]:
|
||||||
|
if isinstance(discovery_port, bool) or not isinstance(discovery_port, int):
|
||||||
|
raise TypeError("discovery_port must be an integer")
|
||||||
|
if not 1 <= discovery_port <= 65535:
|
||||||
|
raise ValueError("discovery_port must be between 1 and 65535")
|
||||||
|
if isinstance(timeout, bool) or not isinstance(timeout, (int, float)):
|
||||||
|
raise TypeError("timeout must be a number")
|
||||||
|
timeout_value = float(timeout)
|
||||||
|
if not math.isfinite(timeout_value) or not 0 < timeout_value <= 30:
|
||||||
|
raise ValueError("timeout must be greater than zero and at most 30")
|
||||||
|
if isinstance(rounds, bool) or not isinstance(rounds, int):
|
||||||
|
raise TypeError("rounds must be an integer")
|
||||||
|
if not 1 <= rounds <= 4:
|
||||||
|
raise ValueError("rounds must be between one and four")
|
||||||
|
return discovery_port, timeout_value, rounds
|
||||||
|
|
||||||
|
|
||||||
|
def _request(
|
||||||
|
token: bytes,
|
||||||
|
message_id: int,
|
||||||
|
*,
|
||||||
|
versioned: bool,
|
||||||
|
filtered: bool,
|
||||||
|
) -> bytes:
|
||||||
|
options = [
|
||||||
|
(URI_PATH, b"oic"),
|
||||||
|
(URI_PATH, b"res"),
|
||||||
|
(ACCEPT, _OCF_CBOR if versioned else CF_CBOR),
|
||||||
|
]
|
||||||
|
if filtered:
|
||||||
|
options.append((URI_QUERY, b"rt=oic.r.doxm"))
|
||||||
|
if versioned:
|
||||||
|
options.append((_OCF_CONTENT_FORMAT_VERSION, _OCF_VERSION_1_0))
|
||||||
|
return build_coap(TYPE_NON, METHOD_GET, message_id, token, options)
|
||||||
|
|
||||||
|
|
||||||
|
def _result(
|
||||||
|
ports: tuple[int, ...],
|
||||||
|
attempts: int,
|
||||||
|
responses: int,
|
||||||
|
error_code: str | None = None,
|
||||||
|
) -> OcfResponderPortDiscoveryResult:
|
||||||
|
return OcfResponderPortDiscoveryResult(
|
||||||
|
ports=ports,
|
||||||
|
attempts=attempts,
|
||||||
|
responses=responses,
|
||||||
|
error_code=error_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def discover_ocf_responder_ports(
|
||||||
|
target_address: str,
|
||||||
|
*,
|
||||||
|
interface_address: str,
|
||||||
|
discovery_port: int = _OCF_DISCOVERY_PORT,
|
||||||
|
timeout: float = 3.0,
|
||||||
|
rounds: int = 2,
|
||||||
|
) -> OcfResponderPortDiscoveryResult:
|
||||||
|
"""Find plaintext OCF response ports for one known IPv4 host.
|
||||||
|
|
||||||
|
Each round sends modern OCF and legacy IoTivity NON requests to the
|
||||||
|
link-local multicast group. Only a 2.05 response with a request token and
|
||||||
|
the exact target source address contributes a candidate. One monotonic
|
||||||
|
deadline bounds all rounds, and every socket is closed before return.
|
||||||
|
"""
|
||||||
|
|
||||||
|
_target_address, target_key = _validate_address(target_address, "target_address")
|
||||||
|
interface_address, interface_key = _validate_address(
|
||||||
|
interface_address, "interface_address"
|
||||||
|
)
|
||||||
|
discovery_port, timeout, rounds = _validate_options(
|
||||||
|
discovery_port=discovery_port,
|
||||||
|
timeout=timeout,
|
||||||
|
rounds=rounds,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
selector = selectors.DefaultSelector()
|
||||||
|
except (OSError, ValueError):
|
||||||
|
return _result((), 0, 0, "interface_unavailable")
|
||||||
|
active = None
|
||||||
|
try:
|
||||||
|
active = socket.socket(socket.AF_INET, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
|
||||||
|
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_IF, interface_key)
|
||||||
|
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_TTL, 1)
|
||||||
|
active.setsockopt(socket.IPPROTO_IP, socket.IP_MULTICAST_LOOP, 0)
|
||||||
|
active.bind((interface_address, 0))
|
||||||
|
active.setblocking(False)
|
||||||
|
selector.register(active, selectors.EVENT_READ)
|
||||||
|
except (OSError, ValueError):
|
||||||
|
if active is not None:
|
||||||
|
try:
|
||||||
|
active.close()
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
selector.close()
|
||||||
|
return _result((), 0, 0, "interface_unavailable")
|
||||||
|
|
||||||
|
started = time.monotonic()
|
||||||
|
deadline = started + timeout
|
||||||
|
accepted_tokens: set[bytes] = set()
|
||||||
|
observations: set[tuple[bytes, int]] = set()
|
||||||
|
ports: list[int] = []
|
||||||
|
seen_ports: set[int] = set()
|
||||||
|
attempts = 0
|
||||||
|
responses = 0
|
||||||
|
too_many_ports = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
for round_number in range(rounds):
|
||||||
|
if time.monotonic() >= deadline:
|
||||||
|
break
|
||||||
|
# Preserve the unfiltered modern and legacy requests used by the
|
||||||
|
# installed appliance generations. Older media firmware can omit
|
||||||
|
# usable endpoint policy from its large unfiltered directory but
|
||||||
|
# answer the smaller legacy DOXM-filtered lookup, so send that as
|
||||||
|
# a third bounded fallback rather than narrowing every request.
|
||||||
|
for versioned, filtered in (
|
||||||
|
(True, False),
|
||||||
|
(False, False),
|
||||||
|
(False, True),
|
||||||
|
):
|
||||||
|
token = secrets.token_bytes(8)
|
||||||
|
while token in accepted_tokens:
|
||||||
|
token = secrets.token_bytes(8)
|
||||||
|
accepted_tokens.add(token)
|
||||||
|
request = _request(
|
||||||
|
token,
|
||||||
|
secrets.randbits(16),
|
||||||
|
versioned=versioned,
|
||||||
|
filtered=filtered,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
sent = active.sendto(
|
||||||
|
request,
|
||||||
|
(_OCF_MULTICAST_GROUP, discovery_port),
|
||||||
|
)
|
||||||
|
except OSError:
|
||||||
|
continue
|
||||||
|
attempts += 1
|
||||||
|
if sent != len(request):
|
||||||
|
continue
|
||||||
|
|
||||||
|
round_deadline = started + timeout * (round_number + 1) / rounds
|
||||||
|
datagrams = 0
|
||||||
|
while datagrams < _MAX_DATAGRAMS_PER_ROUND:
|
||||||
|
remaining = min(deadline, round_deadline) - time.monotonic()
|
||||||
|
if remaining <= 0:
|
||||||
|
break
|
||||||
|
try:
|
||||||
|
events = selector.select(remaining)
|
||||||
|
except (OSError, ValueError):
|
||||||
|
break
|
||||||
|
if not events:
|
||||||
|
break
|
||||||
|
try:
|
||||||
|
datagram, source = active.recvfrom(_MAX_DATAGRAM_BYTES + 1)
|
||||||
|
except (BlockingIOError, OSError):
|
||||||
|
continue
|
||||||
|
datagrams += 1
|
||||||
|
if len(datagram) > _MAX_DATAGRAM_BYTES:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
len(datagram) < 4
|
||||||
|
or datagram[0] >> 6 != 1
|
||||||
|
or datagram[0] & 0x0F > 8
|
||||||
|
or 4 + (datagram[0] & 0x0F) > len(datagram)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
if not isinstance(source, tuple) or len(source) != 2:
|
||||||
|
continue
|
||||||
|
source_host, source_port = source
|
||||||
|
if not isinstance(source_host, str):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
source_key = socket.inet_pton(socket.AF_INET, source_host)
|
||||||
|
except OSError:
|
||||||
|
continue
|
||||||
|
if source_key != target_key:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
isinstance(source_port, bool)
|
||||||
|
or not isinstance(source_port, int)
|
||||||
|
or not 1 <= source_port <= 65535
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
message_type, code, mid, token, _options, payload = parse_coap(
|
||||||
|
datagram
|
||||||
|
)
|
||||||
|
except (IndexError, ValueError, MalformedMessageError):
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
token not in accepted_tokens
|
||||||
|
or code != _CONTENT
|
||||||
|
or message_type not in (TYPE_NON, TYPE_CON)
|
||||||
|
or not payload
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
if message_type == TYPE_CON:
|
||||||
|
try:
|
||||||
|
active.sendto(build_coap(TYPE_ACK, 0, mid, b"", []), source)
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
observation = (token, source_port)
|
||||||
|
if observation in observations:
|
||||||
|
continue
|
||||||
|
observations.add(observation)
|
||||||
|
responses += 1
|
||||||
|
if source_port in seen_ports:
|
||||||
|
continue
|
||||||
|
if len(ports) >= _MAX_PORTS:
|
||||||
|
too_many_ports = True
|
||||||
|
continue
|
||||||
|
seen_ports.add(source_port)
|
||||||
|
ports.append(source_port)
|
||||||
|
|
||||||
|
if too_many_ports:
|
||||||
|
return _result((), attempts, responses, "ambiguous_response")
|
||||||
|
if ports:
|
||||||
|
return _result(tuple(ports), attempts, responses)
|
||||||
|
if attempts == 0:
|
||||||
|
return _result((), attempts, responses, "interface_unavailable")
|
||||||
|
return _result((), attempts, responses, "no_response")
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
selector.unregister(active)
|
||||||
|
except (KeyError, OSError, ValueError):
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
active.close()
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
selector.close()
|
||||||
@@ -0,0 +1,128 @@
|
|||||||
|
"""Pure IoTivity manufacturer-certificate OwnerPSK derivation."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Mapping
|
||||||
|
import hashlib
|
||||||
|
import hmac
|
||||||
|
from types import MappingProxyType
|
||||||
|
from typing import Final
|
||||||
|
|
||||||
|
|
||||||
|
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL: Final = b"x.org.iotivity.conmfgcert"
|
||||||
|
STANDARD_MFG_CERTIFICATE_OXM_LABEL: Final = b"oic.sec.doxm.mfgcert"
|
||||||
|
|
||||||
|
# OpenSSL cipher names mapped to the key-block lengths used by IoTivity's
|
||||||
|
# CAGenerateOwnerPSK implementation.
|
||||||
|
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS: Final[Mapping[str, int]] = MappingProxyType(
|
||||||
|
{
|
||||||
|
"ECDHE-ECDSA-AES128-SHA256": 96,
|
||||||
|
"ECDHE-ECDSA-AES128-CCM": 40,
|
||||||
|
"ECDHE-ECDSA-AES128-CCM8": 40,
|
||||||
|
"ECDHE-ECDSA-AES128-GCM-SHA256": 120,
|
||||||
|
"AES256-SHA256": 128,
|
||||||
|
"ECDHE-ECDSA-AES256-SHA384": 160,
|
||||||
|
"ECDHE-ECDSA-AES256-GCM-SHA384": 184,
|
||||||
|
"AES128-GCM-SHA256": 120,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
_TLS_MASTER_SECRET_BYTES: Final = 48
|
||||||
|
_TLS_RANDOM_BYTES: Final = 32
|
||||||
|
_OCF_UUID_BYTES: Final = 16
|
||||||
|
_OWNER_PSK_BYTES: Final = 16
|
||||||
|
|
||||||
|
|
||||||
|
def _require_bytes(name: str, value: bytes, length: int) -> bytes:
|
||||||
|
if not isinstance(value, bytes):
|
||||||
|
raise TypeError(f"{name} must be bytes")
|
||||||
|
if len(value) != length:
|
||||||
|
raise ValueError(f"{name} must be exactly {length} bytes")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _require_uuid(name: str, value: bytes) -> bytes:
|
||||||
|
value = _require_bytes(name, value, _OCF_UUID_BYTES)
|
||||||
|
if not any(value):
|
||||||
|
raise ValueError(f"{name} must not be the nil UUID")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _tls12_p_hash_sha256(
|
||||||
|
key: bytes,
|
||||||
|
label: bytes,
|
||||||
|
random1: bytes,
|
||||||
|
random2: bytes,
|
||||||
|
length: int,
|
||||||
|
) -> bytes:
|
||||||
|
seed = label + random1 + random2
|
||||||
|
a_value = hmac.new(key, seed, hashlib.sha256).digest()
|
||||||
|
output = bytearray()
|
||||||
|
while len(output) < length:
|
||||||
|
output.extend(hmac.new(key, a_value + seed, hashlib.sha256).digest())
|
||||||
|
a_value = hmac.new(key, a_value, hashlib.sha256).digest()
|
||||||
|
return bytes(output[:length])
|
||||||
|
|
||||||
|
|
||||||
|
def derive_mfg_certificate_owner_psk(
|
||||||
|
*,
|
||||||
|
master_secret: bytes,
|
||||||
|
client_random: bytes,
|
||||||
|
server_random: bytes,
|
||||||
|
owner_uuid: bytes,
|
||||||
|
device_uuid: bytes,
|
||||||
|
cipher_name: str,
|
||||||
|
oxm_label: bytes,
|
||||||
|
) -> bytes:
|
||||||
|
"""Derive a 128-bit OwnerPSK from caller-supplied DTLS state.
|
||||||
|
|
||||||
|
This implements IoTivity's two-stage TLS 1.2 SHA-256 P_hash operation.
|
||||||
|
It performs no session access, network I/O, ownership writes, or storage.
|
||||||
|
The caller must supply state from an authenticated manufacturer-certificate
|
||||||
|
session and explicitly select the OXM label used by that transaction.
|
||||||
|
IoTivity's other 96-byte ECDH_ANON, ECDHE_PSK, and ECDHE_RSA mappings are
|
||||||
|
intentionally outside this helper's manufacturer-certificate allowlist.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not isinstance(cipher_name, str):
|
||||||
|
raise TypeError("cipher_name must be a string")
|
||||||
|
key_block_bytes = MFG_CERTIFICATE_KEY_BLOCK_LENGTHS.get(cipher_name)
|
||||||
|
if key_block_bytes is None:
|
||||||
|
raise ValueError("unexpected manufacturer-certificate DTLS cipher")
|
||||||
|
|
||||||
|
master_secret = _require_bytes(
|
||||||
|
"master_secret", master_secret, _TLS_MASTER_SECRET_BYTES
|
||||||
|
)
|
||||||
|
client_random = _require_bytes(
|
||||||
|
"client_random", client_random, _TLS_RANDOM_BYTES
|
||||||
|
)
|
||||||
|
server_random = _require_bytes(
|
||||||
|
"server_random", server_random, _TLS_RANDOM_BYTES
|
||||||
|
)
|
||||||
|
owner_uuid = _require_uuid("owner_uuid", owner_uuid)
|
||||||
|
device_uuid = _require_uuid("device_uuid", device_uuid)
|
||||||
|
if not isinstance(oxm_label, bytes):
|
||||||
|
raise TypeError("oxm_label must be bytes")
|
||||||
|
if oxm_label not in {
|
||||||
|
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
|
||||||
|
STANDARD_MFG_CERTIFICATE_OXM_LABEL,
|
||||||
|
}:
|
||||||
|
raise ValueError("unexpected manufacturer-certificate OXM label")
|
||||||
|
|
||||||
|
key_block = _tls12_p_hash_sha256(
|
||||||
|
master_secret,
|
||||||
|
b"key expansion",
|
||||||
|
server_random,
|
||||||
|
client_random,
|
||||||
|
key_block_bytes,
|
||||||
|
)
|
||||||
|
# IoTivity's OTM callers pass the owner UUID first and the target device
|
||||||
|
# UUID second. The lower adapter's historical rsrc/prov parameter names
|
||||||
|
# describe those arguments inconsistently, so preserve the caller order.
|
||||||
|
return _tls12_p_hash_sha256(
|
||||||
|
key_block,
|
||||||
|
oxm_label,
|
||||||
|
owner_uuid,
|
||||||
|
device_uuid,
|
||||||
|
_OWNER_PSK_BYTES,
|
||||||
|
)
|
||||||
File diff suppressed because it is too large
Load Diff
+19
-1
@@ -2,7 +2,8 @@ import pytest
|
|||||||
|
|
||||||
from smartthings_local.errors import MalformedMessageError
|
from smartthings_local.errors import MalformedMessageError
|
||||||
from smartthings_local.protocol.coap import (
|
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,
|
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
|
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():
|
def test_fmt_code_formats_class_dot_detail():
|
||||||
assert fmt_code(0x45) == '2.05'
|
assert fmt_code(0x45) == '2.05'
|
||||||
assert fmt_code(0x84) == '4.04'
|
assert fmt_code(0x84) == '4.04'
|
||||||
|
|||||||
@@ -376,7 +376,7 @@ def test_session_uses_connected_socket_send_and_recv(monkeypatch):
|
|||||||
assert open_calls == [(('device.example', 5684), {
|
assert open_calls == [(('device.example', 5684), {
|
||||||
'family': socket.AF_INET6,
|
'family': socket.AF_INET6,
|
||||||
'local_port': None,
|
'local_port': None,
|
||||||
'timeout': 2.0,
|
'timeout': 0.5,
|
||||||
})]
|
})]
|
||||||
|
|
||||||
session.close()
|
session.close()
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ def test_smartthings_local_imports_without_mqtt_demo_present(tmp_path):
|
|||||||
|
|
||||||
import_lines = [
|
import_lines = [
|
||||||
"import smartthings_local.protocol.coap",
|
"import smartthings_local.protocol.coap",
|
||||||
|
"import smartthings_local.protocol.ocf_multicast",
|
||||||
"import smartthings_local.protocol.dtls_session",
|
"import smartthings_local.protocol.dtls_session",
|
||||||
"import smartthings_local.ocf.state_cache",
|
"import smartthings_local.ocf.state_cache",
|
||||||
"import smartthings_local.ocf.poll_scheduler",
|
"import smartthings_local.ocf.poll_scheduler",
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -0,0 +1,392 @@
|
|||||||
|
"""Known-host OCF multicast responder-port discovery tests."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from smartthings_local.protocol.coap import (
|
||||||
|
ACCEPT,
|
||||||
|
TYPE_CON,
|
||||||
|
TYPE_NON,
|
||||||
|
URI_PATH,
|
||||||
|
URI_QUERY,
|
||||||
|
build_coap,
|
||||||
|
parse_coap,
|
||||||
|
)
|
||||||
|
from smartthings_local.protocol.ocf_multicast import (
|
||||||
|
_OCF_MULTICAST_GROUP,
|
||||||
|
OcfResponderPortDiscoveryResult,
|
||||||
|
discover_ocf_responder_ports,
|
||||||
|
)
|
||||||
|
|
||||||
|
_TARGET = "192.0.2.20"
|
||||||
|
_OTHER = "192.0.2.21"
|
||||||
|
_INTERFACE = "192.0.2.10"
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeSocket:
|
||||||
|
def __init__(self, responder=None):
|
||||||
|
self.responder = responder
|
||||||
|
self.incoming = []
|
||||||
|
self.sent = []
|
||||||
|
self.options = []
|
||||||
|
self.bound = None
|
||||||
|
self.blocking = None
|
||||||
|
self.closed = False
|
||||||
|
|
||||||
|
def setsockopt(self, level, option, value):
|
||||||
|
self.options.append((level, option, value))
|
||||||
|
|
||||||
|
def bind(self, address):
|
||||||
|
self.bound = address
|
||||||
|
|
||||||
|
def setblocking(self, enabled):
|
||||||
|
self.blocking = enabled
|
||||||
|
|
||||||
|
def sendto(self, datagram, destination):
|
||||||
|
self.sent.append((datagram, destination))
|
||||||
|
if self.responder is not None:
|
||||||
|
self.incoming.extend(self.responder(datagram, destination))
|
||||||
|
return len(datagram)
|
||||||
|
|
||||||
|
def recvfrom(self, _size):
|
||||||
|
return self.incoming.pop(0)
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
self.closed = True
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeSelector:
|
||||||
|
def __init__(self, active):
|
||||||
|
self.active = active
|
||||||
|
self.closed = False
|
||||||
|
|
||||||
|
def register(self, *_args):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def unregister(self, *_args):
|
||||||
|
return None
|
||||||
|
|
||||||
|
def select(self, _timeout):
|
||||||
|
if not self.active.incoming:
|
||||||
|
return []
|
||||||
|
return [(SimpleNamespace(fileobj=self.active), selectors_event_read())]
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
self.closed = True
|
||||||
|
|
||||||
|
|
||||||
|
def selectors_event_read():
|
||||||
|
return 1
|
||||||
|
|
||||||
|
|
||||||
|
def _response(
|
||||||
|
datagram, _destination=None, *, host=_TARGET, port=43123, message_type=TYPE_NON
|
||||||
|
):
|
||||||
|
_mtype, _code, mid, token, _options, _payload = parse_coap(datagram)
|
||||||
|
response_mid = mid if message_type != TYPE_CON else (mid + 1) & 0xFFFF
|
||||||
|
return [
|
||||||
|
(
|
||||||
|
build_coap(message_type, 0x45, response_mid, token, [], b"directory"),
|
||||||
|
(host, port),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def patch_socket(monkeypatch):
|
||||||
|
created = []
|
||||||
|
|
||||||
|
def install(responder=None):
|
||||||
|
active = _FakeSocket(responder)
|
||||||
|
selector = _FakeSelector(active)
|
||||||
|
created.append((active, selector))
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"smartthings_local.protocol.ocf_multicast.socket.socket",
|
||||||
|
lambda *_args: active,
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"smartthings_local.protocol.ocf_multicast.selectors.DefaultSelector",
|
||||||
|
lambda: selector,
|
||||||
|
)
|
||||||
|
return active, selector
|
||||||
|
|
||||||
|
return install
|
||||||
|
|
||||||
|
|
||||||
|
def test_sends_proven_directory_requests_and_filtered_fallback_on_interface(
|
||||||
|
patch_socket,
|
||||||
|
):
|
||||||
|
active, selector = patch_socket(_response)
|
||||||
|
|
||||||
|
result = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
rounds=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.ports == (43123,)
|
||||||
|
assert result.attempts == 3
|
||||||
|
assert result.responses == 3
|
||||||
|
assert active.bound == (_INTERFACE, 0)
|
||||||
|
assert active.blocking is False
|
||||||
|
assert active.closed
|
||||||
|
assert selector.closed
|
||||||
|
assert all(
|
||||||
|
destination == (_OCF_MULTICAST_GROUP, 5683) for _, destination in active.sent
|
||||||
|
)
|
||||||
|
|
||||||
|
requests = [parse_coap(datagram) for datagram, _ in active.sent]
|
||||||
|
assert all(
|
||||||
|
[value for number, value in item[4] if number == URI_PATH] == [b"oic", b"res"]
|
||||||
|
for item in requests
|
||||||
|
)
|
||||||
|
queries = [
|
||||||
|
[value for number, value in item[4] if number == URI_QUERY] for item in requests
|
||||||
|
]
|
||||||
|
assert queries == [[], [], [b"rt=oic.r.doxm"]]
|
||||||
|
accepts = [
|
||||||
|
[value for number, value in item[4] if number == ACCEPT] for item in requests
|
||||||
|
]
|
||||||
|
assert accepts == [[b"\x27\x10"], [b"\x3c"], [b"\x3c"]]
|
||||||
|
assert [value for number, value in requests[0][4] if number == 2049] == [
|
||||||
|
b"\x08\x00"
|
||||||
|
]
|
||||||
|
assert [value for number, value in requests[1][4] if number == 2049] == []
|
||||||
|
assert [value for number, value in requests[2][4] if number == 2049] == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("accepted_accept", "accepted_query"),
|
||||||
|
(
|
||||||
|
(b"\x27\x10", ()),
|
||||||
|
(b"\x3c", ()),
|
||||||
|
(b"\x3c", (b"rt=oic.r.doxm",)),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
def test_each_directory_request_profile_can_find_the_responder(
|
||||||
|
patch_socket, accepted_accept, accepted_query
|
||||||
|
):
|
||||||
|
def responder(datagram, destination):
|
||||||
|
_mtype, _code, _mid, _token, options, _payload = parse_coap(datagram)
|
||||||
|
accept = next(value for number, value in options if number == ACCEPT)
|
||||||
|
query = tuple(value for number, value in options if number == URI_QUERY)
|
||||||
|
if (accept, query) != (accepted_accept, accepted_query):
|
||||||
|
return []
|
||||||
|
return _response(datagram, destination)
|
||||||
|
|
||||||
|
patch_socket(responder)
|
||||||
|
result = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
rounds=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.ports == (43123,)
|
||||||
|
assert result.responses == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_accepts_only_token_correlated_content_from_the_target(patch_socket):
|
||||||
|
def responder(datagram, _destination):
|
||||||
|
responses = _response(datagram)
|
||||||
|
_mtype, _code, mid, token, _options, _payload = parse_coap(datagram)
|
||||||
|
responses.extend(
|
||||||
|
[
|
||||||
|
(build_coap(TYPE_NON, 0x45, mid, b"wrong", [], b"x"), (_TARGET, 49999)),
|
||||||
|
(build_coap(TYPE_NON, 0x44, mid, token, [], b"x"), (_TARGET, 49998)),
|
||||||
|
(build_coap(TYPE_NON, 0x45, mid, token, [], b"x"), (_OTHER, 49997)),
|
||||||
|
(build_coap(TYPE_NON, 0x45, mid, token, [], b""), (_TARGET, 49996)),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
return responses
|
||||||
|
|
||||||
|
patch_socket(responder)
|
||||||
|
result = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
rounds=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.ports == (43123,)
|
||||||
|
assert result.responses == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_ignores_malformed_and_oversized_datagrams(patch_socket):
|
||||||
|
def responder(datagram, _destination):
|
||||||
|
valid = _response(datagram)
|
||||||
|
invalid_version = bytes([valid[0][0][0] & 0x3F]) + valid[0][0][1:]
|
||||||
|
invalid_token_length = bytes([0x49]) + valid[0][0][1:]
|
||||||
|
return [
|
||||||
|
(b"\x40", (_TARGET, 49999)),
|
||||||
|
(invalid_version, (_TARGET, 49998)),
|
||||||
|
(invalid_token_length, (_TARGET, 49997)),
|
||||||
|
(b"x" * 8193, (_TARGET, 49996)),
|
||||||
|
*valid,
|
||||||
|
]
|
||||||
|
|
||||||
|
patch_socket(responder)
|
||||||
|
result = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
rounds=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.ports == (43123,)
|
||||||
|
assert result.responses == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_acknowledges_confirmable_responses(patch_socket):
|
||||||
|
active, _selector = patch_socket(
|
||||||
|
lambda datagram, _destination: _response(datagram, message_type=TYPE_CON)
|
||||||
|
)
|
||||||
|
|
||||||
|
result = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
rounds=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.found
|
||||||
|
acknowledgements = [
|
||||||
|
(parse_coap(datagram), destination)
|
||||||
|
for datagram, destination in active.sent
|
||||||
|
if parse_coap(datagram)[0] == 2
|
||||||
|
]
|
||||||
|
assert len(acknowledgements) == 3
|
||||||
|
assert all(item[0][1] == 0 and item[0][3] == b"" for item in acknowledgements)
|
||||||
|
assert all(
|
||||||
|
destination == (_TARGET, 43123) for _item, destination in acknowledgements
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_collects_a_small_bounded_candidate_set(patch_socket):
|
||||||
|
response_number = 0
|
||||||
|
|
||||||
|
def responder(datagram, _destination):
|
||||||
|
nonlocal response_number
|
||||||
|
port = 40000 + response_number % 8
|
||||||
|
response_number += 1
|
||||||
|
return _response(datagram, port=port)
|
||||||
|
|
||||||
|
patch_socket(responder)
|
||||||
|
result = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
rounds=4,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.ports == tuple(range(40000, 40008))
|
||||||
|
assert result.error_code is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_fails_closed_when_too_many_distinct_ports_answer(patch_socket):
|
||||||
|
next_port = 40000
|
||||||
|
|
||||||
|
def responder(datagram, _destination):
|
||||||
|
nonlocal next_port
|
||||||
|
responses = [
|
||||||
|
*_response(datagram, port=next_port),
|
||||||
|
*_response(datagram, port=next_port + 1),
|
||||||
|
]
|
||||||
|
next_port += 2
|
||||||
|
return responses
|
||||||
|
|
||||||
|
patch_socket(responder)
|
||||||
|
result = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
rounds=4,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.ports == ()
|
||||||
|
assert result.responses == 24
|
||||||
|
assert result.error_code == "ambiguous_response"
|
||||||
|
|
||||||
|
|
||||||
|
def test_no_response_and_interface_failures_are_fixed_results(
|
||||||
|
patch_socket, monkeypatch
|
||||||
|
):
|
||||||
|
active, _selector = patch_socket()
|
||||||
|
no_response = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
rounds=1,
|
||||||
|
)
|
||||||
|
assert no_response == OcfResponderPortDiscoveryResult(
|
||||||
|
ports=(),
|
||||||
|
attempts=3,
|
||||||
|
responses=0,
|
||||||
|
error_code="no_response",
|
||||||
|
)
|
||||||
|
assert active.closed
|
||||||
|
|
||||||
|
class BrokenSocket:
|
||||||
|
def setsockopt(self, *_args):
|
||||||
|
raise OSError("synthetic")
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
return None
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"smartthings_local.protocol.ocf_multicast.socket.socket",
|
||||||
|
lambda *_args: BrokenSocket(),
|
||||||
|
)
|
||||||
|
unavailable = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
)
|
||||||
|
assert unavailable.error_code == "interface_unavailable"
|
||||||
|
assert unavailable.attempts == 0
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"smartthings_local.protocol.ocf_multicast.selectors.DefaultSelector",
|
||||||
|
lambda: (_ for _ in ()).throw(OSError("synthetic")),
|
||||||
|
)
|
||||||
|
unavailable = discover_ocf_responder_ports(
|
||||||
|
_TARGET,
|
||||||
|
interface_address=_INTERFACE,
|
||||||
|
)
|
||||||
|
assert unavailable.error_code == "interface_unavailable"
|
||||||
|
assert unavailable.attempts == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("kwargs", "exception"),
|
||||||
|
[
|
||||||
|
({"target_address": 123}, TypeError),
|
||||||
|
({"target_address": "not-an-address"}, ValueError),
|
||||||
|
({"target_address": _OCF_MULTICAST_GROUP}, ValueError),
|
||||||
|
({"interface_address": "0.0.0.0"}, ValueError),
|
||||||
|
({"discovery_port": True}, TypeError),
|
||||||
|
({"discovery_port": 0}, ValueError),
|
||||||
|
({"timeout": float("nan")}, ValueError),
|
||||||
|
({"rounds": 0}, ValueError),
|
||||||
|
({"rounds": 5}, ValueError),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_rejects_invalid_options_without_opening_a_socket(
|
||||||
|
monkeypatch, kwargs, exception
|
||||||
|
):
|
||||||
|
values = {
|
||||||
|
"target_address": _TARGET,
|
||||||
|
"interface_address": _INTERFACE,
|
||||||
|
**kwargs,
|
||||||
|
}
|
||||||
|
socket_factory = pytest.fail
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"smartthings_local.protocol.ocf_multicast.socket.socket",
|
||||||
|
socket_factory,
|
||||||
|
)
|
||||||
|
with pytest.raises(exception):
|
||||||
|
discover_ocf_responder_ports(**values)
|
||||||
|
|
||||||
|
|
||||||
|
def test_result_repr_omits_ports_and_addresses():
|
||||||
|
result = OcfResponderPortDiscoveryResult(ports=(43123,), attempts=2, responses=1)
|
||||||
|
|
||||||
|
rendered = repr(result)
|
||||||
|
assert "43123" not in rendered
|
||||||
|
assert _TARGET not in rendered
|
||||||
|
assert "port_count=1" in rendered
|
||||||
@@ -0,0 +1,86 @@
|
|||||||
|
"""Oven descriptor flatten() contracts for the HA Number entity's range.
|
||||||
|
|
||||||
|
The oven reports ``x.com.samsung.da.desired = 0`` whenever no cycle is
|
||||||
|
set. That is "no setpoint", not a 0 °C target, and publishing it as one
|
||||||
|
makes Home Assistant reject every state message against the Number
|
||||||
|
entity's declared 30-270 range.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from mqtt_demo.samples import oven
|
||||||
|
|
||||||
|
|
||||||
|
def _links(desired, current=180):
|
||||||
|
"""A /temperatures/vs/0 link tree carrying one desired/current pair."""
|
||||||
|
return {
|
||||||
|
'/temperatures/vs/0': {
|
||||||
|
'x.com.samsung.da.items': [{
|
||||||
|
'x.com.samsung.da.current': str(current),
|
||||||
|
'x.com.samsung.da.desired': str(desired),
|
||||||
|
}],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize('desired', [
|
||||||
|
oven.SETPOINT_MIN_C,
|
||||||
|
oven.SETPOINT_MIN_C + oven.SETPOINT_STEP_C,
|
||||||
|
180,
|
||||||
|
oven.SETPOINT_MAX_C,
|
||||||
|
])
|
||||||
|
def test_settable_setpoints_are_published_unchanged(desired):
|
||||||
|
assert oven.flatten(_links(desired))['target_temp_c'] == desired
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize('desired', [
|
||||||
|
0, # the idle oven; see module docstring
|
||||||
|
oven.SETPOINT_MIN_C - 1,
|
||||||
|
oven.SETPOINT_MAX_C + 1,
|
||||||
|
])
|
||||||
|
def test_unsettable_setpoints_are_published_as_absent(desired):
|
||||||
|
assert oven.flatten(_links(desired))['target_temp_c'] is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_out_of_range_setpoint_does_not_suppress_current_temperature():
|
||||||
|
"""The guard applies to the setpoint alone. A cooling oven still
|
||||||
|
reports its cavity temperature after the cycle ends."""
|
||||||
|
sensors = oven.flatten(_links(0, current=210))
|
||||||
|
|
||||||
|
assert sensors['target_temp_c'] is None
|
||||||
|
assert sensors['current_temp_c'] == 210
|
||||||
|
|
||||||
|
|
||||||
|
def test_missing_temperature_resource_leaves_both_absent():
|
||||||
|
sensors = oven.flatten({})
|
||||||
|
|
||||||
|
assert sensors['target_temp_c'] is None
|
||||||
|
assert sensors['current_temp_c'] is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_every_committed_write_is_a_value_flatten_will_publish():
|
||||||
|
"""The write path snaps to the step grid *before* bounds-checking, so
|
||||||
|
it accepts more than flatten() publishes: 29 commits as 30, and 271 as
|
||||||
|
270. That is fine for a slider, but it means the two range checks are
|
||||||
|
not symmetric. What has to hold is the weaker invariant: any setpoint
|
||||||
|
the oven is actually told to adopt is one flatten() will show back,
|
||||||
|
otherwise a write appears to succeed and then reads as unknown."""
|
||||||
|
handler = oven.command_handlers()[oven.CMD_SETPOINT]
|
||||||
|
|
||||||
|
for requested in range(-20, oven.SETPOINT_MAX_C + 40):
|
||||||
|
write = handler(str(requested), _links(180))
|
||||||
|
if write is None:
|
||||||
|
continue
|
||||||
|
_path, body = write
|
||||||
|
committed = int(body['x.com.samsung.da.items'][0][
|
||||||
|
'x.com.samsung.da.desired'])
|
||||||
|
assert oven.flatten(_links(committed))['target_temp_c'] == committed
|
||||||
|
|
||||||
|
|
||||||
|
def test_zero_is_rejected_on_the_write_path_too():
|
||||||
|
"""0 is the one value that neither snaps into range nor publishes."""
|
||||||
|
handler = oven.command_handlers()[oven.CMD_SETPOINT]
|
||||||
|
|
||||||
|
assert handler('0', _links(180)) is None
|
||||||
@@ -0,0 +1,144 @@
|
|||||||
|
"""IoTivity manufacturer-certificate OwnerPSK derivation contracts."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from smartthings_local.protocol.owner_psk import (
|
||||||
|
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
|
||||||
|
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS,
|
||||||
|
STANDARD_MFG_CERTIFICATE_OXM_LABEL,
|
||||||
|
derive_mfg_certificate_owner_psk,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_VALID_INPUTS = {
|
||||||
|
"master_secret": bytes(range(48)),
|
||||||
|
"client_random": bytes(range(32)),
|
||||||
|
"server_random": bytes(range(32, 64)),
|
||||||
|
"owner_uuid": bytes.fromhex("00112233445566778899aabbccddeeff"),
|
||||||
|
"device_uuid": bytes.fromhex("ffeeddccbbaa99887766554433221100"),
|
||||||
|
"cipher_name": "ECDHE-ECDSA-AES128-GCM-SHA256",
|
||||||
|
"oxm_label": CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# These synthetic expected values were generated from the fixed inputs above;
|
||||||
|
# they were not captured from IoTivity or a device. They lock deterministic
|
||||||
|
# output for the selected mappings, while the table test below states the
|
||||||
|
# key-block-length contract explicitly.
|
||||||
|
def test_fixed_synthetic_gcm_regression_vector():
|
||||||
|
assert derive_mfg_certificate_owner_psk(**_VALID_INPUTS).hex() == (
|
||||||
|
"ccd6c618a91290dee8c106544ed79a33"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_owner_then_device_uuid_order_matches_iotivity_callers():
|
||||||
|
reversed_context = derive_mfg_certificate_owner_psk(
|
||||||
|
**{
|
||||||
|
**_VALID_INPUTS,
|
||||||
|
"owner_uuid": _VALID_INPUTS["device_uuid"],
|
||||||
|
"device_uuid": _VALID_INPUTS["owner_uuid"],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert reversed_context.hex() == "8f0f2416c483546dc1806db769b21b68"
|
||||||
|
assert reversed_context != derive_mfg_certificate_owner_psk(**_VALID_INPUTS)
|
||||||
|
|
||||||
|
|
||||||
|
def test_fixed_synthetic_ccm8_regression_vector():
|
||||||
|
inputs = {
|
||||||
|
**_VALID_INPUTS,
|
||||||
|
"cipher_name": "ECDHE-ECDSA-AES128-CCM8",
|
||||||
|
}
|
||||||
|
assert derive_mfg_certificate_owner_psk(**inputs).hex() == (
|
||||||
|
"ddd3d945e266ee3dc27ff3a2c4321d32"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_standard_and_confirmed_labels_derive_distinct_keys():
|
||||||
|
confirmed = derive_mfg_certificate_owner_psk(**_VALID_INPUTS)
|
||||||
|
standard = derive_mfg_certificate_owner_psk(
|
||||||
|
**{**_VALID_INPUTS, "oxm_label": STANDARD_MFG_CERTIFICATE_OXM_LABEL}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert standard.hex() == "26ee1fe4c3e74509a2f5db5ab41b1e47"
|
||||||
|
assert standard != confirmed
|
||||||
|
|
||||||
|
|
||||||
|
def test_iotivity_cipher_key_block_lengths_are_immutable():
|
||||||
|
assert dict(MFG_CERTIFICATE_KEY_BLOCK_LENGTHS) == {
|
||||||
|
"ECDHE-ECDSA-AES128-SHA256": 96,
|
||||||
|
"ECDHE-ECDSA-AES128-CCM": 40,
|
||||||
|
"ECDHE-ECDSA-AES128-CCM8": 40,
|
||||||
|
"ECDHE-ECDSA-AES128-GCM-SHA256": 120,
|
||||||
|
"AES256-SHA256": 128,
|
||||||
|
"ECDHE-ECDSA-AES256-SHA384": 160,
|
||||||
|
"ECDHE-ECDSA-AES256-GCM-SHA384": 184,
|
||||||
|
"AES128-GCM-SHA256": 120,
|
||||||
|
}
|
||||||
|
with pytest.raises(TypeError):
|
||||||
|
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS["new-cipher"] = 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("field", "length"),
|
||||||
|
[
|
||||||
|
("master_secret", 48),
|
||||||
|
("client_random", 32),
|
||||||
|
("server_random", 32),
|
||||||
|
("owner_uuid", 16),
|
||||||
|
("device_uuid", 16),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_binary_inputs_require_exact_bytes_and_lengths(field, length):
|
||||||
|
with pytest.raises(TypeError, match=f"{field} must be bytes"):
|
||||||
|
derive_mfg_certificate_owner_psk(
|
||||||
|
**{**_VALID_INPUTS, field: bytearray(length)}
|
||||||
|
)
|
||||||
|
for invalid_length in (length - 1, length + 1):
|
||||||
|
with pytest.raises(ValueError, match=f"exactly {length} bytes"):
|
||||||
|
derive_mfg_certificate_owner_psk(
|
||||||
|
**{**_VALID_INPUTS, field: b"x" * invalid_length}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("field", ["owner_uuid", "device_uuid"])
|
||||||
|
def test_nil_uuid_is_rejected(field):
|
||||||
|
with pytest.raises(ValueError, match="must not be the nil UUID"):
|
||||||
|
derive_mfg_certificate_owner_psk(
|
||||||
|
**{**_VALID_INPUTS, field: bytes(16)}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_cipher_and_label_must_be_explicit_supported_values():
|
||||||
|
with pytest.raises(TypeError, match="cipher_name must be a string"):
|
||||||
|
derive_mfg_certificate_owner_psk(
|
||||||
|
**{**_VALID_INPUTS, "cipher_name": b"cipher"}
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match="unexpected.*cipher"):
|
||||||
|
derive_mfg_certificate_owner_psk(
|
||||||
|
**{**_VALID_INPUTS, "cipher_name": "ECDHE-RSA-AES128-GCM-SHA256"}
|
||||||
|
)
|
||||||
|
with pytest.raises(TypeError, match="oxm_label must be bytes"):
|
||||||
|
derive_mfg_certificate_owner_psk(
|
||||||
|
**{**_VALID_INPUTS, "oxm_label": "oic.sec.doxm.mfgcert"}
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match="unexpected.*label"):
|
||||||
|
derive_mfg_certificate_owner_psk(
|
||||||
|
**{**_VALID_INPUTS, "oxm_label": b"unsupported"}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_failures_do_not_include_key_material():
|
||||||
|
key_material = b"private-master-secret"
|
||||||
|
with pytest.raises(ValueError) as raised:
|
||||||
|
derive_mfg_certificate_owner_psk(
|
||||||
|
**{
|
||||||
|
**_VALID_INPUTS,
|
||||||
|
"master_secret": key_material,
|
||||||
|
"cipher_name": "unsupported",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
assert key_material.hex() not in str(raised.value)
|
||||||
|
assert "private-master-secret" not in str(raised.value)
|
||||||
@@ -10,8 +10,19 @@ from smartthings_local.protocol.auth import (
|
|||||||
AuthenticationProvider,
|
AuthenticationProvider,
|
||||||
CertificateAuth,
|
CertificateAuth,
|
||||||
PskAuth,
|
PskAuth,
|
||||||
|
SamsungServerProfile,
|
||||||
|
SamsungServerRole,
|
||||||
|
ServerCertificateAuth,
|
||||||
)
|
)
|
||||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
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:
|
def _assert_compatible_signature(callable_object, expected: list[str]) -> None:
|
||||||
@@ -50,10 +61,77 @@ def test_dtls_session_constructor_keeps_file_memory_and_local_port_inputs():
|
|||||||
assert auth_parameter.default is None
|
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():
|
def test_certificate_auth_is_a_public_authentication_provider():
|
||||||
provider = CertificateAuth.from_files("/synthetic/cert.pem", "/synthetic/key")
|
provider = CertificateAuth.from_files("/synthetic/cert.pem", "/synthetic/key")
|
||||||
assert isinstance(provider, AuthenticationProvider)
|
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():
|
def test_psk_auth_is_a_public_authentication_provider():
|
||||||
provider = PskAuth(identity=b"i" * 16, key=b"k" * 16)
|
provider = PskAuth(identity=b"i" * 16, key=b"k" * 16)
|
||||||
@@ -67,6 +145,26 @@ def test_psk_auth_is_a_public_authentication_provider():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_owner_psk_derivation_keeps_every_security_input_explicit():
|
||||||
|
parameters = inspect.signature(
|
||||||
|
derive_mfg_certificate_owner_psk
|
||||||
|
).parameters
|
||||||
|
assert list(parameters) == [
|
||||||
|
"master_secret",
|
||||||
|
"client_random",
|
||||||
|
"server_random",
|
||||||
|
"owner_uuid",
|
||||||
|
"device_uuid",
|
||||||
|
"cipher_name",
|
||||||
|
"oxm_label",
|
||||||
|
]
|
||||||
|
assert all(
|
||||||
|
parameter.kind is inspect.Parameter.KEYWORD_ONLY
|
||||||
|
and parameter.default is inspect.Parameter.empty
|
||||||
|
for parameter in parameters.values()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_dtls_session_keeps_current_consumer_methods():
|
def test_dtls_session_keeps_current_consumer_methods():
|
||||||
expected = {
|
expected = {
|
||||||
"close",
|
"close",
|
||||||
@@ -81,6 +179,20 @@ def test_dtls_session_keeps_current_consumer_methods():
|
|||||||
"subscribe",
|
"subscribe",
|
||||||
}
|
}
|
||||||
assert expected <= set(dir(DtlsCoapSession))
|
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(
|
_assert_compatible_signature(
|
||||||
DtlsCoapSession.get,
|
DtlsCoapSession.get,
|
||||||
[
|
[
|
||||||
|
|||||||
@@ -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"")
|
||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -3,7 +3,12 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import threading
|
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.keepalive import KeepaliveTask
|
||||||
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
|
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
|
||||||
from smartthings_local.ocf.poll_scheduler import PollScheduler, PollTier
|
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=())],
|
[PollTier("idle", interval_s=3600.0, paths=())],
|
||||||
)
|
)
|
||||||
_assert_worker_stops(scheduler.run_forever, "test-poll-scheduler")
|
_assert_worker_stops(scheduler.run_forever, "test-poll-scheduler")
|
||||||
|
|
||||||
|
|
||||||
|
class _SessionWorker:
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
self.stop = None
|
||||||
|
self.started = threading.Event()
|
||||||
|
self.exited = threading.Event()
|
||||||
|
self.on_reachable = kwargs.get("on_reachable")
|
||||||
|
self.on_unreachable = kwargs.get("on_unreachable")
|
||||||
|
self.last_success_ts = 0.0
|
||||||
|
|
||||||
|
def run_forever(self, stop):
|
||||||
|
self.stop = stop
|
||||||
|
self.started.set()
|
||||||
|
stop.wait()
|
||||||
|
self.exited.set()
|
||||||
|
|
||||||
|
|
||||||
|
class _JoinedSession:
|
||||||
|
def __init__(self, workers, start_index, error=None):
|
||||||
|
self.workers = workers
|
||||||
|
self.start_index = start_index
|
||||||
|
self.error = error
|
||||||
|
|
||||||
|
def join(self):
|
||||||
|
assert all(
|
||||||
|
worker.started.wait(_THREAD_DEADLINE_S)
|
||||||
|
for worker in self.workers[self.start_index:]
|
||||||
|
)
|
||||||
|
if self.error is not None:
|
||||||
|
raise self.error
|
||||||
|
|
||||||
|
|
||||||
|
class _BlockingJoinedSession(_JoinedSession):
|
||||||
|
def __init__(self, workers):
|
||||||
|
super().__init__(workers, 0)
|
||||||
|
self.joined = threading.Event()
|
||||||
|
self.release = threading.Event()
|
||||||
|
|
||||||
|
def join(self):
|
||||||
|
super().join()
|
||||||
|
self.joined.set()
|
||||||
|
assert self.release.wait(_THREAD_DEADLINE_S)
|
||||||
|
|
||||||
|
|
||||||
|
def _bridge():
|
||||||
|
bridge = object.__new__(PushBridge)
|
||||||
|
bridge.descriptor = SimpleNamespace(
|
||||||
|
observe_paths=(),
|
||||||
|
poll_tiers=[],
|
||||||
|
is_active=lambda _state: False,
|
||||||
|
)
|
||||||
|
bridge.shared = SimpleNamespace(PING_INTERVAL_S=3600.0)
|
||||||
|
bridge.app = SimpleNamespace(klass="test")
|
||||||
|
bridge.log = SimpleNamespace(info=lambda *args: None, warning=lambda *args: None)
|
||||||
|
bridge.cache = SimpleNamespace(links={})
|
||||||
|
bridge.stop = threading.Event()
|
||||||
|
bridge._session_stop_lock = threading.Lock()
|
||||||
|
bridge._session_stop = None
|
||||||
|
bridge.scheduler = None
|
||||||
|
bridge.keepalive = None
|
||||||
|
bridge.observe_refresh = None
|
||||||
|
bridge._seed_from_device0 = lambda _session: None
|
||||||
|
bridge._retag_logger_with_serial = lambda: None
|
||||||
|
bridge.maybe_publish_state = lambda **kwargs: None
|
||||||
|
bridge.set_availability = lambda _online: None
|
||||||
|
return bridge
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("session_count", "join_error"),
|
||||||
|
[
|
||||||
|
pytest.param(2, None, id="reconnect"),
|
||||||
|
pytest.param(1, RuntimeError("reader failed"), id="reader-error"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_bridge_retires_session_workers_before_returning(
|
||||||
|
monkeypatch, session_count, join_error
|
||||||
|
):
|
||||||
|
workers = []
|
||||||
|
|
||||||
|
def worker_factory(*args, **kwargs):
|
||||||
|
worker = _SessionWorker(*args, **kwargs)
|
||||||
|
workers.append(worker)
|
||||||
|
return worker
|
||||||
|
|
||||||
|
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
|
||||||
|
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
|
||||||
|
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
|
||||||
|
|
||||||
|
bridge = _bridge()
|
||||||
|
|
||||||
|
session_stops = []
|
||||||
|
for session_index in range(session_count):
|
||||||
|
start_index = len(workers)
|
||||||
|
session = _JoinedSession(
|
||||||
|
workers,
|
||||||
|
start_index,
|
||||||
|
error=join_error if session_index == session_count - 1 else None,
|
||||||
|
)
|
||||||
|
if session.error is None:
|
||||||
|
bridge._run_session_inner(session)
|
||||||
|
else:
|
||||||
|
with pytest.raises(RuntimeError, match="reader failed"):
|
||||||
|
bridge._run_session_inner(session)
|
||||||
|
|
||||||
|
session_workers = workers[start_index:]
|
||||||
|
assert len(session_workers) == 3
|
||||||
|
assert len({id(worker.stop) for worker in session_workers}) == 1
|
||||||
|
session_stops.append(session_workers[0].stop)
|
||||||
|
|
||||||
|
assert len(workers) == session_count * 3
|
||||||
|
assert len({id(stop) for stop in session_stops}) == session_count
|
||||||
|
assert all(stop is not bridge.stop for stop in session_stops)
|
||||||
|
assert all(stop.is_set() for stop in session_stops)
|
||||||
|
assert not bridge.stop.is_set()
|
||||||
|
assert all(worker.exited.is_set() for worker in workers)
|
||||||
|
keepalive_workers = tuple(
|
||||||
|
workers[index] for index in range(1, len(workers), 3)
|
||||||
|
)
|
||||||
|
assert all(worker.on_reachable is None for worker in keepalive_workers)
|
||||||
|
assert all(worker.on_unreachable is None for worker in keepalive_workers)
|
||||||
|
assert bridge.scheduler is None
|
||||||
|
assert bridge.keepalive is None
|
||||||
|
assert bridge.observe_refresh is None
|
||||||
|
assert bridge._session_stop is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_bridge_retires_started_worker_when_later_thread_fails_to_start(
|
||||||
|
monkeypatch,
|
||||||
|
):
|
||||||
|
workers = []
|
||||||
|
|
||||||
|
def worker_factory(*args, **kwargs):
|
||||||
|
worker = _SessionWorker(*args, **kwargs)
|
||||||
|
workers.append(worker)
|
||||||
|
return worker
|
||||||
|
|
||||||
|
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
|
||||||
|
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
|
||||||
|
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
|
||||||
|
|
||||||
|
bridge = _bridge()
|
||||||
|
|
||||||
|
original_start = threading.Thread.start
|
||||||
|
start_count = 0
|
||||||
|
|
||||||
|
def fail_second_start(thread):
|
||||||
|
nonlocal start_count
|
||||||
|
start_count += 1
|
||||||
|
if start_count == 2:
|
||||||
|
raise RuntimeError("synthetic thread start failure")
|
||||||
|
original_start(thread)
|
||||||
|
|
||||||
|
monkeypatch.setattr(threading.Thread, "start", fail_second_start)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="synthetic thread start failure"):
|
||||||
|
bridge._run_session_inner(SimpleNamespace(join=lambda: None))
|
||||||
|
|
||||||
|
assert len(workers) == 3
|
||||||
|
assert workers[0].started.wait(_THREAD_DEADLINE_S)
|
||||||
|
assert workers[0].exited.wait(_THREAD_DEADLINE_S)
|
||||||
|
assert workers[0].stop is not bridge.stop
|
||||||
|
assert workers[0].stop.is_set()
|
||||||
|
assert workers[1].stop is None
|
||||||
|
assert workers[2].stop is None
|
||||||
|
assert bridge.scheduler is None
|
||||||
|
assert bridge.keepalive is None
|
||||||
|
assert bridge.observe_refresh is None
|
||||||
|
assert bridge._session_stop is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_request_stop_wakes_current_session_workers(monkeypatch):
|
||||||
|
workers = []
|
||||||
|
|
||||||
|
def worker_factory(*args, **kwargs):
|
||||||
|
worker = _SessionWorker(*args, **kwargs)
|
||||||
|
workers.append(worker)
|
||||||
|
return worker
|
||||||
|
|
||||||
|
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
|
||||||
|
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
|
||||||
|
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
|
||||||
|
|
||||||
|
bridge = _bridge()
|
||||||
|
session = _BlockingJoinedSession(workers)
|
||||||
|
session_thread = threading.Thread(
|
||||||
|
target=bridge._run_session_inner,
|
||||||
|
args=(session,),
|
||||||
|
daemon=True,
|
||||||
|
)
|
||||||
|
session_thread.start()
|
||||||
|
|
||||||
|
try:
|
||||||
|
assert session.joined.wait(_THREAD_DEADLINE_S)
|
||||||
|
session_stop = bridge._session_stop
|
||||||
|
assert session_stop is not None
|
||||||
|
assert not session_stop.is_set()
|
||||||
|
|
||||||
|
bridge.request_stop()
|
||||||
|
|
||||||
|
assert bridge.stop.is_set()
|
||||||
|
assert session_stop.is_set()
|
||||||
|
assert all(
|
||||||
|
worker.exited.wait(_THREAD_DEADLINE_S) for worker in workers
|
||||||
|
)
|
||||||
|
assert session_thread.is_alive()
|
||||||
|
finally:
|
||||||
|
session.release.set()
|
||||||
|
session_thread.join(_THREAD_DEADLINE_S)
|
||||||
|
|
||||||
|
assert not session_thread.is_alive()
|
||||||
|
assert bridge._session_stop is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_session_workers_observe_stop_requested_before_handoff(monkeypatch):
|
||||||
|
workers = []
|
||||||
|
|
||||||
|
def worker_factory(*args, **kwargs):
|
||||||
|
worker = _SessionWorker(*args, **kwargs)
|
||||||
|
workers.append(worker)
|
||||||
|
return worker
|
||||||
|
|
||||||
|
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
|
||||||
|
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
|
||||||
|
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
|
||||||
|
|
||||||
|
bridge = _bridge()
|
||||||
|
bridge.request_stop()
|
||||||
|
bridge._run_session_inner(_JoinedSession(workers, 0))
|
||||||
|
|
||||||
|
assert len(workers) == 3
|
||||||
|
assert all(worker.stop.is_set() for worker in workers)
|
||||||
|
assert all(worker.exited.wait(_THREAD_DEADLINE_S) for worker in workers)
|
||||||
|
assert bridge._session_stop is None
|
||||||
|
|||||||
Reference in New Issue
Block a user