Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b3045f5ddb | ||
|
|
0d9d13b8fb | ||
|
|
0722c552c9 | ||
|
|
a45ee4c004 | ||
|
|
6b9a508fd9 | ||
|
|
ab0ec4961c | ||
|
|
6db846563a | ||
|
|
8c374e17b4 | ||
|
|
cd86424ca0 | ||
|
|
6da7e95731 | ||
|
|
a3bec9470d | ||
|
|
79493fbd48 | ||
|
|
627fcb19da | ||
|
|
31be87061a | ||
|
|
bc4465b274 | ||
|
|
e63acb759f | ||
|
|
231e88a8c6 | ||
|
|
dec84c9ad8 | ||
|
|
e9aee1c235 | ||
|
|
9ef5598813 | ||
|
|
917b0e47c5 | ||
|
|
83a5973434 | ||
|
|
3f0e437880 | ||
|
|
a44930f9df | ||
|
|
b0d51abcc8 | ||
|
|
7a74a955f3 | ||
|
|
d4aebebca4 | ||
|
|
7f0f1e5531 | ||
|
|
2a6fc627f2 | ||
|
|
8fb37ca2ed | ||
|
|
74d78543e1 | ||
|
|
da8f917541 | ||
|
|
a63bc0bb01 | ||
|
|
11c87b275b | ||
|
|
7ec999b3ff | ||
|
|
31a52d6ff6 | ||
|
|
2a82a764bb | ||
|
|
dd453ebdfb | ||
|
|
d677c72f89 | ||
|
|
c7e15a7dd3 | ||
|
|
280939d646 | ||
|
|
98e0020e2f | ||
|
|
597a88ff25 | ||
|
|
93e39de079 | ||
|
|
fc6240b72e | ||
|
|
6dc9dca339 | ||
|
|
2c93cb3097 | ||
|
|
23338995bf | ||
|
|
119c114daa | ||
|
|
e5bd9456d4 | ||
|
|
a494e73e89 | ||
|
|
e0622eb087 |
@@ -0,0 +1 @@
|
||||
buy_me_a_coffee: quiteyellow
|
||||
@@ -0,0 +1,130 @@
|
||||
name: Validate
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: validate-${{ github.workflow }}-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
tests:
|
||||
name: Python ${{ matrix.python-version }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python-version:
|
||||
- "3.11"
|
||||
- "3.12"
|
||||
- "3.13"
|
||||
- "3.14"
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
- uses: actions/setup-python@v7
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- run: python -m pip install --upgrade pip
|
||||
- run: python -m pip install -e ".[dev]"
|
||||
- run: python -m pytest -q
|
||||
|
||||
dependency-bounds:
|
||||
name: Dependencies (${{ matrix.mode }})
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- mode: floor
|
||||
python-version: "3.11"
|
||||
- mode: latest
|
||||
python-version: "3.14"
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
- uses: actions/setup-python@v7
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
- run: python -m pip install --upgrade pip
|
||||
- if: matrix.mode == 'floor'
|
||||
run: >-
|
||||
python -m pip install
|
||||
"cbor2==5.6.0"
|
||||
"pyOpenSSL==23.1.0"
|
||||
"pytest==8.0.0"
|
||||
- if: matrix.mode == 'floor'
|
||||
run: python -m pip install --no-deps -e .
|
||||
- if: matrix.mode == 'latest'
|
||||
run: python -m pip install -e ".[dev]"
|
||||
- run: python -m pytest -q
|
||||
|
||||
package:
|
||||
name: Package artifacts
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- uses: actions/setup-python@v7
|
||||
with:
|
||||
python-version: "3.14"
|
||||
- run: python -m pip install --upgrade pip
|
||||
- run: python -m pip install build hatchling hatch-vcs
|
||||
- run: python -m build
|
||||
- run: python tools/check_distribution.py dist
|
||||
- name: Install and import wheel
|
||||
run: |
|
||||
python -m venv "$RUNNER_TEMP/wheel-smoke"
|
||||
"$RUNNER_TEMP/wheel-smoke/bin/python" -m pip install \
|
||||
dist/*.whl
|
||||
cd "$RUNNER_TEMP"
|
||||
"$RUNNER_TEMP/wheel-smoke/bin/python" -I -c \
|
||||
"from smartthings_local.protocol.dtls_session import DtlsCoapSession"
|
||||
- name: Install and import sdist
|
||||
run: |
|
||||
python -m venv "$RUNNER_TEMP/sdist-smoke"
|
||||
"$RUNNER_TEMP/sdist-smoke/bin/python" -m pip install \
|
||||
dist/*.tar.gz
|
||||
cd "$RUNNER_TEMP"
|
||||
"$RUNNER_TEMP/sdist-smoke/bin/python" -I -c \
|
||||
"from smartthings_local.ocf.state_cache import StateCache"
|
||||
|
||||
share-safety:
|
||||
name: Share safety
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- uses: actions/checkout@v7
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- uses: actions/setup-python@v7
|
||||
with:
|
||||
python-version: "3.14"
|
||||
- name: Select comparison base
|
||||
id: comparison
|
||||
env:
|
||||
PR_BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
PUSH_BEFORE_SHA: ${{ github.event.before }}
|
||||
run: |
|
||||
if [ -n "$PR_BASE_SHA" ]; then
|
||||
echo "sha=$PR_BASE_SHA" >> "$GITHUB_OUTPUT"
|
||||
elif [ -n "$PUSH_BEFORE_SHA" ] && \
|
||||
[ "$PUSH_BEFORE_SHA" != "0000000000000000000000000000000000000000" ]; then
|
||||
echo "sha=$PUSH_BEFORE_SHA" >> "$GITHUB_OUTPUT"
|
||||
else
|
||||
echo "sha=$(git rev-parse HEAD^)" >> "$GITHUB_OUTPUT"
|
||||
fi
|
||||
- run: >-
|
||||
python tools/check_share_safety.py
|
||||
--changed-since "${{ steps.comparison.outputs.sha }}"
|
||||
@@ -15,10 +15,10 @@ jobs:
|
||||
name: Build sdist + wheel
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v7
|
||||
with:
|
||||
fetch-depth: 0 # hatch-vcs needs full history + tags to derive the version
|
||||
- uses: actions/setup-python@v5
|
||||
- uses: actions/setup-python@v7
|
||||
with:
|
||||
python-version: "3.12"
|
||||
- run: python -m pip install --upgrade build
|
||||
@@ -30,7 +30,7 @@ jobs:
|
||||
if ! ls dist/ | grep -q "smartthings_local-${version}"; then
|
||||
echo "Built artifacts do not match tag version ${version}"; exit 1
|
||||
fi
|
||||
- uses: actions/upload-artifact@v4
|
||||
- uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: dist
|
||||
path: dist/
|
||||
@@ -43,7 +43,7 @@ jobs:
|
||||
permissions:
|
||||
id-token: write # required for Trusted Publishing (OIDC)
|
||||
steps:
|
||||
- uses: actions/download-artifact@v4
|
||||
- uses: actions/download-artifact@v8
|
||||
with:
|
||||
name: dist
|
||||
path: dist/
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
SmartThings-Local
|
||||
Copyright (c) 2026 Jack Nagy
|
||||
|
||||
This software is licensed under the MIT License. See the LICENSE file for
|
||||
the full terms.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
Trademarks & disclaimer
|
||||
------------------------------------------------------------------------
|
||||
|
||||
This is an independent, unofficial project. It is NOT affiliated with,
|
||||
authorised, endorsed, or sponsored by Samsung Electronics Co., Ltd. or any
|
||||
of its subsidiaries.
|
||||
|
||||
"Samsung", "SmartThings", and any related names, marks, and logos are
|
||||
trademarks of Samsung Electronics Co., Ltd. They are used in this project
|
||||
only nominatively -- to identify the hardware and protocols this software
|
||||
interoperates with -- and no claim is made to any right in them. Use of
|
||||
these marks does not imply any affiliation with or endorsement by their
|
||||
owner.
|
||||
|
||||
The software is provided for interoperability with hardware you own,
|
||||
without warranty of any kind.
|
||||
@@ -1,6 +1,6 @@
|
||||
# SmartThings-Local
|
||||
|
||||
**`smartthings-local` is a Python library for local, cloud-free control of newer-generation Samsung connected appliances over cert-authenticated CoAP-DTLS.** It gives you the DTLS-CoAP transport, a tiered polling + OBSERVE state layer, and one-command identity-cert minting. That covers everything needed to read state from and write commands to a Samsung dryer, oven, fridge, etc. on your LAN, with no SmartThings cloud round-trip.
|
||||
**`smartthings-local` is a Python library for local, cloud-free control of Samsung connected appliances over authenticated CoAP-DTLS.** It gives you the DTLS-CoAP transport, a tiered polling + OBSERVE state layer, and identity-cert tooling for AC14K_M-compatible firmware. Newer OCF-PKI appliances require a different authentication profile; see [the laundry compatibility findings](https://github.com/QuiteYellow/SmartThings-Local/blob/main/docs/ocf-pki-laundry.md). Supported profiles can read state and write commands on the LAN with no SmartThings cloud round-trip.
|
||||
|
||||
The repo also ships a self-contained **reference bridge demo** (`mqtt_demo/`) that turns the library into auto-discovered Home Assistant entities over MQTT. One process supervises multiple appliances, each on its own DTLS session.
|
||||
|
||||
@@ -21,16 +21,22 @@ The repo also ships a self-contained **reference bridge demo** (`mqtt_demo/`) th
|
||||
pip install smartthings-local
|
||||
```
|
||||
|
||||
Mint a client cert once (see [Part 2](#part-2--auth-get-the-identity-cert)), then drive a session directly:
|
||||
For compatible firmware, mint a client cert once (see
|
||||
[Part 2](#part-2--auth-for-ac14k_m-compatible-firmware)), then drive a
|
||||
session directly:
|
||||
|
||||
```python
|
||||
import cbor2
|
||||
from smartthings_local.protocol.auth import CertificateAuth
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
||||
|
||||
auth = CertificateAuth.from_files(
|
||||
"certs/client_fullchain.pem",
|
||||
"certs/client.key",
|
||||
)
|
||||
sess = DtlsCoapSession(
|
||||
"192.168.1.100", 49154,
|
||||
cert_path="certs/client_fullchain.pem",
|
||||
key_path="certs/client.key",
|
||||
"192.0.2.100", 49154,
|
||||
auth=auth,
|
||||
)
|
||||
sess.connect()
|
||||
sess.start_reader()
|
||||
@@ -42,12 +48,220 @@ sess.subscribe(["operational", "state", "vs", "0"], # OBSERVE
|
||||
sess.close()
|
||||
```
|
||||
|
||||
If the cert/key are minted at runtime and never written to disk (e.g. inside an HA config flow), pass them in memory instead of by path:
|
||||
`connect()` uses a 12-second monotonic DTLS handshake deadline by default. A
|
||||
caller that needs a shorter bounded attempt can pass a positive finite value
|
||||
without changing later reader timeouts. OpenSSL's DTLS timer schedules flight
|
||||
retransmissions within that same deadline:
|
||||
|
||||
```python
|
||||
sess = DtlsCoapSession("192.168.1.100", 49154, cert_pem=cert_pem, key_pem=key_pem)
|
||||
sess.connect(timeout=4.0)
|
||||
```
|
||||
|
||||
The deadline stops further setup, retries, and network waits. If OpenSSL
|
||||
reports that the handshake completed at the deadline boundary, the completed
|
||||
session is retained rather than torn down as a timeout.
|
||||
|
||||
Connection attempts can also use a one-way cancellation signal. The signal is
|
||||
backed by a socketpair, so setting it wakes the network wait immediately while
|
||||
OpenSSL retains control of DTLS retransmission timing:
|
||||
|
||||
```python
|
||||
from smartthings_local.protocol.dtls_session import ConnectCancellation
|
||||
|
||||
cancel_connect = ConnectCancellation()
|
||||
# Another thread may call cancel_connect.set().
|
||||
sess.connect(timeout=8.0, cancel=cancel_connect)
|
||||
```
|
||||
|
||||
Setting the signal stops subscribed connection attempts and closes their
|
||||
temporary UDP sockets. It does not alter an already established session or add
|
||||
new session lifecycle methods. Interrupted attempts raise `SessionClosedError`.
|
||||
|
||||
If the cert/key are minted at runtime and never written to disk (e.g. inside
|
||||
an HA config flow), create the provider from memory instead:
|
||||
|
||||
```python
|
||||
auth = CertificateAuth.from_memory(cert_pem, key_pem)
|
||||
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
|
||||
```
|
||||
|
||||
Some newer OCF-PKI devices require an exact Samsung DTLS offer and present a
|
||||
hardware certificate whose subject contains a certificate UUID. That UUID can
|
||||
be distinct from the runtime OCF device UUID reported by `/oic/d`, so callers
|
||||
must obtain and verify the certificate identity independently. When the caller
|
||||
already has an authorized client certificate and a previously verified
|
||||
hardware-certificate UUID, opt in to both requirements explicitly:
|
||||
|
||||
```python
|
||||
from smartthings_local.protocol.auth import (
|
||||
CertificateAuth,
|
||||
SamsungServerProfile,
|
||||
)
|
||||
|
||||
server_profile = SamsungServerProfile.bound_device(
|
||||
expected_certificate_uuid,
|
||||
additional_ca_pem=additional_samsung_ca_pem,
|
||||
)
|
||||
auth = CertificateAuth.from_memory(
|
||||
cert_pem,
|
||||
key_pem,
|
||||
server_profile=server_profile,
|
||||
)
|
||||
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
|
||||
```
|
||||
|
||||
The default profile is restricted to Samsung home-appliance leaves with
|
||||
`OU=OCF HA Device`. The profile limits the ClientHello to P-256,
|
||||
`ECDHE-ECDSA-AES128-GCM-SHA256`, and the observed SHA-256/SHA-1 RSA/ECDSA
|
||||
signature set, disables session tickets, preserves certificate-chain
|
||||
verification, and requires the exact subject role
|
||||
`C=KR, O=Samsung Electronics, OU=OCF HA Device` with a common name ending in
|
||||
the expected certificate UUID. `additional_ca_pem` is optional and accepts
|
||||
only a bounded PEM CA-certificate chain; it is applied only to this profiled
|
||||
context. Without a profile, `CertificateAuth` retains its existing verification
|
||||
behavior.
|
||||
|
||||
Samsung VD-family devices can present the same wire profile with the distinct
|
||||
`OU=OCF VD Device` role. Select that role explicitly; profiles never fall back
|
||||
between device classes:
|
||||
|
||||
```python
|
||||
from smartthings_local.protocol.auth import (
|
||||
SamsungServerProfile,
|
||||
SamsungServerRole,
|
||||
ServerCertificateAuth,
|
||||
)
|
||||
|
||||
server_profile = SamsungServerProfile.bound_device(
|
||||
expected_certificate_uuid,
|
||||
role=SamsungServerRole.VD_DEVICE,
|
||||
)
|
||||
auth = ServerCertificateAuth(server_profile=server_profile)
|
||||
sess = DtlsCoapSession("192.0.2.100", 5684, auth=auth)
|
||||
```
|
||||
|
||||
`ServerCertificateAuth` is for a server-authenticated channel that does not
|
||||
send a client certificate, such as the initial DTLS carrier used by
|
||||
manufacturer-certificate OTM. It still verifies the CA chain, exact selected
|
||||
subject role, and pinned certificate UUID. It does not learn an identity from
|
||||
the first endpoint it reaches and cannot be combined with client credentials.
|
||||
|
||||
This API deliberately does not discover, mint, authorize, provision, rotate,
|
||||
or persist credentials, and it performs no ownership transfer or OCF security
|
||||
resource writes. In particular, the server-only provider can authenticate the
|
||||
initial manufacturer-certificate channel, but it does not implement the OTM
|
||||
that follows. The already-owned new-PKI case in
|
||||
[issue #16](https://github.com/QuiteYellow/SmartThings-Local/issues/16) still
|
||||
requires an authorized client identity before ordinary protected resources
|
||||
can be used.
|
||||
|
||||
For compatibility, the existing `cert_path` / `key_path` and `cert_pem` /
|
||||
`key_pem` session arguments remain supported without a deprecation warning.
|
||||
They are routed through `CertificateAuth` internally. Do not combine `auth`
|
||||
with those legacy arguments.
|
||||
|
||||
An existing OCF PSK credential can be supplied through `PskAuth`:
|
||||
|
||||
```python
|
||||
from smartthings_local.protocol.auth import PskAuth
|
||||
|
||||
auth = PskAuth(identity=psk_identity, key=psk_key)
|
||||
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
|
||||
```
|
||||
|
||||
The identity must be the raw 16-byte OCF UUID and cannot contain a NUL byte;
|
||||
the key must be exactly 16 or 32 bytes. `PskAuth` selects only
|
||||
`ECDHE-PSK-AES128-CBC-SHA256` and does not acquire, derive, provision, rotate,
|
||||
or persist credentials. Ownership transfer and credential discovery are
|
||||
outside this package.
|
||||
|
||||
Code that has already completed an authenticated manufacturer-certificate
|
||||
session can derive IoTivity's 128-bit OwnerPSK from the resulting TLS state:
|
||||
|
||||
```python
|
||||
from smartthings_local.protocol.owner_psk import derive_mfg_certificate_owner_psk
|
||||
|
||||
owner_psk = derive_mfg_certificate_owner_psk(
|
||||
master_secret=master_secret,
|
||||
client_random=client_random,
|
||||
server_random=server_random,
|
||||
owner_uuid=owner_uuid,
|
||||
device_uuid=device_uuid,
|
||||
cipher_name=cipher_name,
|
||||
oxm_label=selected_oxm_label,
|
||||
)
|
||||
```
|
||||
|
||||
The caller must supply the exact authenticated TLS values, non-nil raw OCF
|
||||
UUIDs, negotiated cipher name, and label for the selected OXM. Use
|
||||
`STANDARD_MFG_CERTIFICATE_OXM_LABEL` for `oic.sec.doxm.mfgcert` and
|
||||
`CONFIRMED_MFG_CERTIFICATE_OXM_LABEL` for
|
||||
`x.org.iotivity.conmfgcert`; do not infer the label from the appliance model.
|
||||
The helper performs deterministic key derivation only: it does not access a
|
||||
session, discover credentials, choose an ownership method, write security
|
||||
resources, run OTM, or persist the result.
|
||||
|
||||
### Classified errors
|
||||
|
||||
Runtime transport failures use the public types in
|
||||
|
||||
```python
|
||||
from smartthings_local.errors import SessionClosedError, SmartThingsLocalError
|
||||
```
|
||||
|
||||
All classified errors inherit from `SmartThingsLocalError` and expose a stable
|
||||
`code`. Their messages are fixed and deliberately omit remote endpoints, local
|
||||
paths, credential metadata, raw packets, and backend exception text. Existing
|
||||
callers can keep catching the built-in types used by earlier releases:
|
||||
|
||||
| Error | Stable code | Compatible built-in |
|
||||
| --- | --- | --- |
|
||||
| `EndpointError` | `endpoint` | `OSError` |
|
||||
| `ProbeError` | `probe` | `ConnectionError` |
|
||||
| `SessionError` | `session` | `ConnectionError` |
|
||||
| `AuthenticationError` | `authentication` | `ConnectionError` |
|
||||
| `AuthorizationError` | `authorization` | `PermissionError` |
|
||||
| `SessionTimeoutError` | `timeout` | `TimeoutError` |
|
||||
| `SessionClosedError` | `session_closed` | `ConnectionError` |
|
||||
| `MalformedMessageError` | `malformed_message` | `ValueError` |
|
||||
| `BlockwiseError` | `blockwise` | `ConnectionError` |
|
||||
| `ObserveError` | `observe` | `ConnectionError` |
|
||||
|
||||
Constructor argument validation remains a normal `ValueError`. When a backend
|
||||
failure is chained for debugging, the cause is replaced with a fixed redacted
|
||||
marker; raw backend text is not copied into the public error or its formatted
|
||||
traceback.
|
||||
|
||||
### Resolved UDP endpoints
|
||||
|
||||
Sessions resolve a host to a first-class `ResolvedUdpEndpoint` and use a
|
||||
connected UDP socket for the DTLS transport. Connecting the datagram socket
|
||||
pins it to the exact resolved peer, so unrelated datagrams from another host
|
||||
using the same port are discarded by the operating system. IPv4, IPv6, and
|
||||
scoped IPv6 tuples are preserved without putting the address or scope in the
|
||||
endpoint's `repr`.
|
||||
|
||||
Address family and fixed source-port behavior are explicit and optional:
|
||||
|
||||
```python
|
||||
import socket
|
||||
|
||||
sess = DtlsCoapSession(
|
||||
"device.example",
|
||||
49154,
|
||||
cert_pem=cert_pem,
|
||||
key_pem=key_pem,
|
||||
family=socket.AF_INET6,
|
||||
local_port=56830,
|
||||
)
|
||||
sess.connect()
|
||||
assert sess.endpoint.family == socket.AF_INET6
|
||||
```
|
||||
|
||||
The resolver retains candidate order and the socket setup tries the next
|
||||
candidate after a family, bind, or connect failure. Resolution and socket
|
||||
setup failures raise the redacted `EndpointError` documented above.
|
||||
|
||||
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
|
||||
@@ -66,7 +280,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.
|
||||
|
||||
Authentication uses a client cert keyed to the UUID published in Samsung's own wildcard cloud TLS cert. Every Samsung Tizen/RT-OCF appliance's factory ACL grants that UUID `perm=31` (full CRUDN) on `href=*`, so a single cert chain works across the whole fleet. Setup is one Python script.
|
||||
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.
|
||||
|
||||
---
|
||||
|
||||
@@ -75,23 +289,24 @@ Authentication uses a client cert keyed to the UUID published in Samsung's own w
|
||||
Check before anything else; if it's older firmware, this project doesn't target it.
|
||||
|
||||
```sh
|
||||
# UDP scan for DTLS-CoAP ports
|
||||
nmap -Pn -sU -p 49152-49160 "$APPLIANCE_IP"
|
||||
# UDP scan for public/secure standard OCF plus the dynamic appliance band
|
||||
nmap -Pn -sU -p 5683,5684,49152-49160 "$APPLIANCE_IP"
|
||||
```
|
||||
|
||||
Read the result:
|
||||
|
||||
- **`49154/udp` (or similar 4915x) open|filtered with a DTLS handshake responding** → newer firmware (Tizen RT 3.x with DAWIT 3.0). This is what the bridge talks to.
|
||||
- **`5684/udp` or a 4915x port with a DTLS first-flight response** → an OCF DTLS listener. Standard-port OCF-PKI firmware needs the Samsung server-certificate profile (`SamsungServerProfile` / `ServerCertificateAuth`, see Quick start), and no working client credential for it exists yet.
|
||||
- **`5683/udp` responds to public OCF security/resource GETs** → use `/oic/res` to learn the device's advertised secure endpoint; do not assume that endpoint is fixed.
|
||||
- **Only `8888/tcp` open (token-based HTTPS)** → older firmware (~2018–2022). **Not supported here.**
|
||||
|
||||
nmap's `open|filtered` can't tell a real DTLS server from a silent UDP port. Confirm which of the candidate ports actually speaks DTLS with the ClientHello probe, which sends one ClientHello and reports back per port:
|
||||
|
||||
```sh
|
||||
# Stateless liveness check: one ClientHello round trip, leaves no state on the device
|
||||
.venv/bin/python -m smartthings_local.protocol.dtls_probe "$APPLIANCE_IP" 49153 49154 49155 49156 --stateless
|
||||
python -m smartthings_local.protocol.dtls_probe "$APPLIANCE_IP" 5684 49153 49154 49155 49156 --stateless
|
||||
```
|
||||
|
||||
`live` means a DTLS server answered its `HelloVerifyRequest` (that's your control port); `dead` means silent / not DTLS. Once you have the client cert (Part 2), drop `--stateless` to run the default *diagnostic* drive, which reports `completed` (cert accepted) or `rejected` with the server's fatal alert. An `unsupported_certificate` / `unknown_ca` alert is the signature of a newer OCF-PKI device that won't accept the AC14K_M cert. The same probe gates the bridge's own reconnect loop and auto-discovers the port when `OCF_PORT` is unset.
|
||||
`live` means a DTLS server answered its first flight; `dead` means silent or not DTLS. Once you have the client cert (Part 2), add the explicit `--diagnostic` flag to run the stateful diagnostic drive, which reports `completed` (cert accepted) or `rejected` with the server's fatal alert. Diagnostic mode can allocate appliance-side DTLS state and is never used by discovery or reconnect. An `unsupported_certificate` / `unknown_ca` alert means the endpoint is reachable but this certificate profile was rejected. It is not a reason to disable verification or keep retrying. The same bounded stateless API gates the bridge's reconnect loop and, when `OCF_PORT` is unset, probes both standard 5684 and ports 49152–49160.
|
||||
|
||||
### Tested combinations
|
||||
|
||||
@@ -104,6 +319,13 @@ nmap's `open|filtered` can't tell a real DTLS server from a silent UDP port. Con
|
||||
|
||||
Other appliances on the same firmware family (dishwashers, AC units) almost certainly speak the same protocol: the auth path and read primitives are common, and a washer on the shared `DA_WM_TP2_20_COMMON` controller is already confirmed above. You'd write one new descriptor for the `localthings` registry.
|
||||
|
||||
The Bespoke AI Laundry Combo `WD53DBA900HZ[A1]` on Tizen 7 software
|
||||
`20260416.215549` is a known OCF-PKI profile, but is not yet supported by the
|
||||
public authentication path. Its endpoint and manufacturer-OTM/OwnerPSK findings
|
||||
are documented [here](https://github.com/QuiteYellow/SmartThings-Local/blob/main/docs/ocf-pki-laundry.md), including the exact relationship
|
||||
to issues [#16](https://github.com/QuiteYellow/SmartThings-Local/issues/16) and
|
||||
[#20](https://github.com/QuiteYellow/SmartThings-Local/issues/20).
|
||||
|
||||
### Firmware families: a limitation
|
||||
|
||||
Descriptors are firmware-family-specific. Each descriptor hardcodes the resource layout of one firmware family: which hrefs it polls, which fields it reads, which write surfaces it exposes. There's no runtime feature detection. The three sample descriptors here (`mqtt_demo/samples/`) are frozen references.
|
||||
@@ -127,9 +349,12 @@ Which path is doing the work is visible in Home Assistant. The bridge publishes
|
||||
|
||||
---
|
||||
|
||||
## Part 2 — Auth: get the identity cert
|
||||
## Part 2 — Auth for AC14K_M-compatible firmware
|
||||
|
||||
The bridge authenticates with a **client cert** signed by `AC14K_M`, an intermediate CA that has been public for years and remains in current firmware trust stores. The cert's Subject DN carries a UUID that the on-device ACL grants full access to.
|
||||
For a compatible firmware family, the bridge authenticates with a **client
|
||||
cert** signed by `AC14K_M`, an intermediate CA that has been public for years.
|
||||
The cert's Subject DN carries a UUID that those appliances' on-device ACLs
|
||||
grant full access to.
|
||||
|
||||
You can read the UUID yourself out of the relevant server cert:
|
||||
|
||||
@@ -146,7 +371,7 @@ This README doesn't pin the literal UUID: the setup script extracts it live each
|
||||
|
||||
### Why this works
|
||||
|
||||
- Every Samsung Tizen/RT-OCF appliance has a **factory-baked ACE** in `/oic/sec/acl` granting this UUID `perm=31` on `href=*`.
|
||||
- Each currently supported Tizen/RT-OCF firmware family has a **factory-baked ACE** in `/oic/sec/acl` granting this UUID `perm=31` on `href=*`.
|
||||
- TizenRT iotivity derives peerId from `memmem(subject_dn, "uuid:")`, which is RDN-agnostic. A cert with the UUID in CN authenticates the same as one with it in OU.
|
||||
- We don't need the matching private key from the original keyholder. We mint our own key and have `AC14K_M` sign our leaf. Different key, same identity, same access.
|
||||
|
||||
@@ -171,9 +396,15 @@ Output in `./certs/`: `client_fullchain.pem` + `client.key`.
|
||||
|
||||
Neither the UUID nor the AC14K_M bundle is hardcoded in this repo; both are fetched live each run, so the script self-updates if upstream rotates. If either fetch fails, the script prints an inline workaround: supply the UUID via `UUID=<uuid>` env, or supply the AC14K_M bundle via `AC14K_M_CERT_BUNDLE=/path/to/cert.pem`. `BRAYSTORM_URL=<mirror>` points at a different bundle source.
|
||||
|
||||
### How durable is this?
|
||||
On Fedora/RHEL (and other hardened OpenSSL 3.x builds) the default crypto policy blocks SHA-1 signing, which step 5 needs. The script detects this, retries the signing step once with SHA-1 force-enabled for just that command, and only fails if the retry also fails. If it does, it prints the remedy: `sudo update-crypto-policies --set DEFAULT:SHA1` (undo afterward with `sudo update-crypto-policies --set DEFAULT`).
|
||||
|
||||
Rotating the published UUID would require Samsung to re-issue TLS certs across their IoT cloud, push new ACLs to every device in the field, and update the on-device daemon identity: a multi-quarter change with a long backwards-compat tail. `AC14K_M` has been public for years and is still in 2026 firmware trust stores. Local access via this path is roughly as durable as cloud control of these appliances.
|
||||
### How durable is this on the compatible firmware families?
|
||||
|
||||
Rotating the published UUID would require coordinated cloud certificate, ACL,
|
||||
and device identity changes across the compatible firmware families.
|
||||
`AC14K_M` has been public for years and remains accepted by the tested rows
|
||||
above, but it is already rejected by other 2026 appliance profiles. Do not
|
||||
extrapolate this certificate path to an untested model.
|
||||
|
||||
> **Legacy path:** earlier versions used a per-hub-UUID cert via an anonymous `/oic/sec/doxm` read escalation. That still works on the dryer-family firmware but isn't necessary: the cert minted here authenticates against every appliance and survives device resets. The old `bootstrap.py` for the legacy flow was removed when the package was renamed; see git history if you need it.
|
||||
|
||||
@@ -197,20 +428,21 @@ APPLIANCE_COUNT=2
|
||||
|
||||
# Appliance 1 — dryer
|
||||
APPLIANCE_1_CLASS=dryer
|
||||
APPLIANCE_1_IP=192.168.1.100
|
||||
APPLIANCE_1_IP=192.0.2.100
|
||||
APPLIANCE_1_OCF_PORT= # blank → auto-discover across the OCF band (dryer=49155)
|
||||
APPLIANCE_1_TOPIC=samsung_dryer
|
||||
APPLIANCE_1_NAME=Samsung Dryer
|
||||
|
||||
# Appliance 2 — oven
|
||||
APPLIANCE_2_CLASS=oven
|
||||
APPLIANCE_2_IP=192.168.1.101
|
||||
APPLIANCE_2_IP=192.0.2.101
|
||||
APPLIANCE_2_OCF_PORT= # blank → auto-discover across the OCF band (oven=49154)
|
||||
APPLIANCE_2_TOPIC=samsung_oven
|
||||
APPLIANCE_2_NAME=Samsung Oven
|
||||
```
|
||||
|
||||
Each `APPLIANCE_<n>_CLASS` must match a descriptor key in `mqtt_demo/samples/__init__.py::DESCRIPTORS`: currently `dryer`, `oven`, and `fridge`.
|
||||
Each `APPLIANCE_<n>_CLASS` must match a key in
|
||||
`mqtt_demo.samples.DESCRIPTORS`: currently `dryer`, `oven`, and `fridge`.
|
||||
|
||||
---
|
||||
|
||||
@@ -332,7 +564,7 @@ Notes specific to this firmware family:
|
||||
| `APPLIANCE_COUNT` | Number of `APPLIANCE_<n>_*` blocks to read (1-indexed) |
|
||||
| `APPLIANCE_<n>_CLASS` | Descriptor name: `dryer`, `oven`, `fridge` |
|
||||
| `APPLIANCE_<n>_IP` | LAN IP of the appliance |
|
||||
| `APPLIANCE_<n>_OCF_PORT` | Optional. Blank → auto-discover the DTLS port across the OCF band 49153–49156 (via a stateless ClientHello probe); set it to pin a specific port and skip discovery (dryer=49155, oven=49154, fridge=49155) |
|
||||
| `APPLIANCE_<n>_OCF_PORT` | Optional. Blank → probe standard port 5684 and the dynamic range 49152–49160 with a stateless ClientHello; set it to pin and gate one specific port (dryer=49155, oven=49154, fridge=49155) |
|
||||
| `APPLIANCE_<n>_TOPIC` | MQTT topic prefix (also the HA device identifier; changing it re-keys the device) |
|
||||
| `APPLIANCE_<n>_NAME` | Friendly name on the HA device card |
|
||||
| `MQTT_BROKER` / `MQTT_PORT` / `MQTT_USER` / `MQTT_PASS` | Broker config |
|
||||
@@ -398,8 +630,12 @@ smartthings_local/ The installable library — `pip install sm
|
||||
__init__.py
|
||||
protocol/ DTLS-CoAP transport (reusable by any consumer, not just MQTT)
|
||||
__init__.py
|
||||
auth.py Immutable DTLS authentication providers
|
||||
coap.py CoAP wire protocol: message encode/decode, token handling
|
||||
dtls_session.py DTLS session: handshake, client-cert auth (file or in-memory PEM), Block2, liveness
|
||||
dtls_probe.py Stateless DTLS liveness + opt-in stateful diagnostic
|
||||
dtls_handshake.py Shared memory-BIO handshake driver, bounded by a monotonic deadline (used by session + probe)
|
||||
owner_psk.py Pure manufacturer-certificate OwnerPSK derivation
|
||||
ocf_root_ca.pem Samsung OCF root CA, bundled for handshake verification
|
||||
ocf/ OCF resource + state layer (reusable)
|
||||
__init__.py
|
||||
@@ -426,7 +662,7 @@ mqtt_demo/ MQTT bridge demo (consumes smartthings_loca
|
||||
.env.example Template — copy to .env, fill in
|
||||
setup_cert.py One-shot cert minting script (live-fetches AC14K_M + UUID)
|
||||
pyproject.toml Packaging — PyPI dist `smartthings-local`, hatch-vcs versioning
|
||||
tests/ pytest suite (CoAP wire, state cache, import isolation, cert loading)
|
||||
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
|
||||
```
|
||||
|
||||
@@ -478,3 +714,13 @@ If reconnects become persistent (e.g. >10 in a minute) something's wrong: check
|
||||
## Contributing
|
||||
|
||||
If you submit a PR, please don't include real device UUIDs, MACs, serials, IPs, or bearer tokens. Use the placeholders from `.env.example`.
|
||||
|
||||
---
|
||||
|
||||
## Trademarks & disclaimer
|
||||
|
||||
This is an independent, unofficial project. It is **not affiliated with, authorised, endorsed, or sponsored by Samsung Electronics Co., Ltd.** or any of its subsidiaries.
|
||||
|
||||
"Samsung", "SmartThings", and any related names, marks, and logos are trademarks of Samsung Electronics Co., Ltd. They are used in this project **only nominatively** — to identify the hardware and protocols this software interoperates with — and no claim is made to any right in them. Use of these marks does not imply any affiliation with or endorsement by their owner.
|
||||
|
||||
The software is provided under the [MIT License](LICENSE) for interoperability with hardware you own, without warranty of any kind.
|
||||
|
||||
@@ -0,0 +1,240 @@
|
||||
# Newer OCF-PKI laundry: connection findings and current limits
|
||||
|
||||
This is a compatibility and implementation note, not an ownership-reset or
|
||||
onboarding guide. It records the sanitized protocol facts that made local
|
||||
control possible on two Samsung Bespoke AI Laundry Combo appliances and maps
|
||||
those facts to [issue #16](https://github.com/QuiteYellow/SmartThings-Local/issues/16)
|
||||
and [issue #20](https://github.com/QuiteYellow/SmartThings-Local/issues/20).
|
||||
|
||||
The important distinction is that endpoint reachability, DTLS authentication,
|
||||
resource authorization, and OCF ownership are four separate states. A response
|
||||
at one layer is not proof that the next layer is usable.
|
||||
|
||||
## Hardware and software validated
|
||||
|
||||
The locally validated appliances are two `WD53DBA900HZA1` all-in-one
|
||||
washer/dryers. Both report:
|
||||
|
||||
- model family `AWM-US-M64-24-WD80`;
|
||||
- Tizen 7 / One UI 7 Laundry Combo; and
|
||||
- primary software version `20260416.215549`.
|
||||
|
||||
Issue #16 reports `WD53DBA900HZ` and the same primary software version. The
|
||||
reported protocol behavior also matches, so it is the same appliance/software
|
||||
profile for the purposes of this library.
|
||||
|
||||
Issue #20 is different hardware: a `WW11BB534DAWS6` washer and
|
||||
`DV90BB5245AWS6` dryer. Only the washer has detailed protocol evidence in that
|
||||
issue, so nothing here claims that the dryer has the same profile.
|
||||
|
||||
## How the WD53 connection was established
|
||||
|
||||
### 1. Discover OCF instead of assuming a 4915x port
|
||||
|
||||
The WD53 exposes its public OCF surface on UDP 5683. `GET /oic/res` returns a
|
||||
multi-block resource directory and advertises secure endpoint data. The two
|
||||
validated units exposed the same 72 hrefs. They also have IPv4, IPv6 ULA, and
|
||||
IPv6 link-local endpoints, and the secure endpoint can move.
|
||||
|
||||
The practical rules are:
|
||||
|
||||
- include the standard CoAP-DTLS port 5684 as well as the 4915x appliance
|
||||
range;
|
||||
- preserve an IPv6 scope ID instead of flattening a link-local address into a
|
||||
host string;
|
||||
- rediscover the secure endpoint before authentication when the appliance has
|
||||
slept or restarted; and
|
||||
- prove a listener with a DTLS ClientHello instead of treating an Nmap
|
||||
`open|filtered` result as protocol evidence.
|
||||
|
||||
The production liveness probe must stop after the first
|
||||
HelloVerifyRequest/ServerHello/Alert. It never returns the cookie to the
|
||||
appliance, and packet-loss retries resend the exact same first flight. This
|
||||
avoids creating half-open DTLS associations while searching several candidate
|
||||
ports.
|
||||
|
||||
### 2. Treat the AC14K_M rejection as an authentication-profile result
|
||||
|
||||
An AC14K_M client chain reaches the WD53 DTLS server but is rejected with a
|
||||
fatal `unknown_ca` alert. RSA versus ECDSA client keys do not change that
|
||||
result. Re-signing only a leaf with SHA-256 cannot repair a trust chain the
|
||||
appliance does not accept.
|
||||
|
||||
That result does **not** mean local OCF was removed. It means the fleet
|
||||
certificate used by older SmartThings appliances is not the runtime principal
|
||||
for this profile. Repeated AC14K_M attempts, a broader cipher list, or disabling
|
||||
server verification do not produce authorization.
|
||||
|
||||
The accepted cipher for the authenticated paths below is exactly
|
||||
`ECDHE-ECDSA-AES128-GCM-SHA256`. The production sessions did not disable TLS
|
||||
verification. A first-flight diagnostic can classify an offered certificate
|
||||
without authenticating it, but that observation never grants authorization.
|
||||
|
||||
### 3. Read and classify the public security state
|
||||
|
||||
The public security resources expose enough redacted state to choose a safe
|
||||
next step:
|
||||
|
||||
- `/oic/sec/doxm` advertises standard manufacturer-certificate OTM `2` and
|
||||
Samsung manufacturer-certificate OTM `0xFF02` (`65282`);
|
||||
- the two validated units were observed with each of those methods selected;
|
||||
- `/oic/sec/pstat` distinguishes an operational owned device from a real
|
||||
manufacturer ownership-transfer window;
|
||||
- the provisioning nonce rotates on every read; and
|
||||
- this model declares that additional authorization is required.
|
||||
|
||||
Device, owner, and resource-owner UUIDs are sensitive identifiers and are not
|
||||
needed in a public fixture. They must be compared locally and replaced with
|
||||
synthetic values in tests or diagnostics.
|
||||
|
||||
### 4. Use the model's authorized, non-reset transition
|
||||
|
||||
During one-time research while the appliance was idle, the signed-in
|
||||
SmartThings Android path was used to invoke the model's signed same-account,
|
||||
non-factory-reset confirmation. This was a setup research carrier, not a
|
||||
runtime dependency. It intentionally moved the OCF security state from owned
|
||||
operation into a bounded, unowned manufacturer-OTM window; the later steps
|
||||
installed a new OCF owner. SmartThings pairing survived on the two tested
|
||||
units, but that does not make an ownership-changing operation generically safe.
|
||||
The fresh confirmation had to occur immediately before the manufacturer DTLS
|
||||
connection; a delayed confirmation missed the firmware's window.
|
||||
|
||||
The installed `5.0.47` appliance stack exposed provisioning feature `0x4000`
|
||||
and validated two proof requests in this exact order:
|
||||
|
||||
Here `serial_hash_ascii` is the 128-character lowercase hexadecimal SHA-512
|
||||
digest of the ASCII registration serial.
|
||||
|
||||
1. `TriggerSerialHashRequest` checks
|
||||
`SHA256(serial_hash_ascii || nonce_raw)`. The nonce is the current raw four
|
||||
bytes, not its eight-character hexadecimal text. This proof contains no
|
||||
account value.
|
||||
2. The appliance rotates its nonce. `TriggerAutoResetHashRequest` then checks
|
||||
`SHA256(serial_hash_ascii || SHA256(user_id_ascii) || fresh_nonce_raw)`,
|
||||
where the inner SHA-256 is its raw 32-byte digest and the same-account user
|
||||
ID is ten ASCII characters. There are no delimiters between fields.
|
||||
|
||||
The order is the inverse of the method names in a newer application helper.
|
||||
Reversing the requests caused the second stage to fail; matching the appliance
|
||||
order opened the clean manufacturer-certificate RFOTM state. Only a fresh
|
||||
public DOXM/PSTAT read—not an application callback—was accepted as proof of
|
||||
that transition. No serial, account ID, nonce, or computed proof is included
|
||||
here.
|
||||
|
||||
This authorization transition is the part that is **not yet a supported public
|
||||
workflow**. The public formulas explain the installed firmware's checks; they
|
||||
do not supply Samsung's signed request authority or disclose an account value.
|
||||
A stock SmartThings-paired appliance must not be reset, claimed, or have its
|
||||
owner replaced merely because its public OCF endpoint is reachable. A public
|
||||
implementation still needs a model-supported same-account grant that does not
|
||||
depend on private application state, captured credentials, or
|
||||
reverse-engineering tools.
|
||||
|
||||
### 5. Open manufacturer DTLS without a client identity leaf
|
||||
|
||||
Inside the confirmed manufacturer window, the successful carrier is
|
||||
server-authenticated DTLS using Samsung's manufacturer trust path. The client
|
||||
does not present an AC14K_M, TEST, or OneApp identity leaf. This trust-only
|
||||
connection can read the authenticated OCF security state needed for the
|
||||
selected manufacturer OTM.
|
||||
|
||||
No new Samsung CA private key is needed for this step. The earlier
|
||||
`unknown_ca` result and the successful manufacturer carrier are different
|
||||
authentication modes, not contradictory observations.
|
||||
|
||||
### 6. Derive, stage, prove, and finalize OwnerPSK
|
||||
|
||||
The standards-based OwnerPSK derivation uses the selected method's exact label:
|
||||
|
||||
- method `2`: `oic.sec.doxm.mfgcert`;
|
||||
- method `0xFF02`: `x.org.iotivity.conmfgcert`.
|
||||
|
||||
For the negotiated `ECDHE-ECDSA-AES128-GCM-SHA256` session, IoTivity computes:
|
||||
|
||||
1. `key_block = P_SHA256(master_secret, "key expansion" || server_random || client_random, 120)`;
|
||||
2. `OwnerPSK = P_SHA256(key_block, selected_otm_label || owner_uuid || appliance_uuid, 16)`.
|
||||
|
||||
The master secret is 48 bytes, each random is 32 bytes, and each UUID is its
|
||||
raw 16-byte value. The derivation is pure; obtaining the authenticated session
|
||||
and deciding that an ownership transaction is authorized are separate
|
||||
responsibilities.
|
||||
|
||||
The validated transaction stages the derived credential before the first
|
||||
security mutation, writes only the reviewed credential/ACL/DOXM/PSTAT shapes,
|
||||
then proves the new key on a fresh ECDHE-PSK session before publishing it as a
|
||||
usable runtime credential. Final DOXM/PSTAT and public postflight reads must
|
||||
all agree before the transaction is considered complete.
|
||||
|
||||
The resulting OwnerPSK is per appliance. It is never logged, returned by a
|
||||
diagnostic, embedded in a fixture, or committed to source control.
|
||||
|
||||
### 7. Run normal control over OwnerPSK
|
||||
|
||||
After finalization, normal reads and writes use ECDHE-PSK CoAP-DTLS over the
|
||||
currently advertised LAN endpoint. On each validated WD53, that path returned
|
||||
39 complete protected representations with no link stubs. Low-risk settings
|
||||
and power changes were accepted, verified by exact protected readback, and
|
||||
restored. The same changes remained visible through SmartThings, demonstrating
|
||||
coexistence for the tested transaction rather than a cloud replacement.
|
||||
|
||||
When the panel enters deep sleep, the secure endpoint can disappear. Runtime
|
||||
code therefore retains last-good state honestly, backs off, and rediscovers
|
||||
the endpoint when the panel returns; it does not use the cloud or an Android
|
||||
application as a wake or polling dependency.
|
||||
|
||||
## How this maps to issues #16 and #20
|
||||
|
||||
### Issue #16: exact WD53 profile
|
||||
|
||||
Issue #16 reproduces both halves of the initial diagnosis:
|
||||
|
||||
- standard OCF ports rather than a fixed 4915x-only assumption; and
|
||||
- AC14K_M client authentication rejected with `unknown_ca`.
|
||||
|
||||
The validated WD53 work demonstrates a path beyond that boundary:
|
||||
manufacturer OTM followed by per-appliance OwnerPSK runtime authentication.
|
||||
The remaining upstream gap is not proof that the protocol works; it is a safe,
|
||||
portable, owner-preserving authorization and credential setup flow.
|
||||
|
||||
### Issue #20: related `0xFF02` evidence, different models
|
||||
|
||||
The washer in issue #20 exposes public OCF on 5683, a DTLS listener on 49154,
|
||||
and reports `owned:false`, `isop:false`, with only OTM `0xFF02` advertised. That
|
||||
is consistent with a Samsung manufacturer-OTM window, and it makes the WD53
|
||||
`0xFF02` transport and OwnerPSK work directly relevant.
|
||||
|
||||
It is not yet proof of support. The issue reports `handshake_failure` rather
|
||||
than the WD53's `unknown_ca`, and the model-specific additional-authorization,
|
||||
nonce, confirmation timing, security payload, and protected-read behavior have
|
||||
not been validated. The dryer in the issue has not supplied equivalent
|
||||
evidence. Both devices need independent, non-destructive validation.
|
||||
|
||||
## What this pull request does and does not solve
|
||||
|
||||
This pull request implements the endpoint half of these reports:
|
||||
|
||||
- bounded, connected IPv4/IPv6 stateless probes;
|
||||
- byte-identical first-flight retransmission;
|
||||
- concurrent standard-port and 4915x probing;
|
||||
- deterministic listener selection; and
|
||||
- an explicit ambiguous result instead of first-responder guessing.
|
||||
|
||||
It does not make AC14K_M authenticate to either issue's appliance and does not
|
||||
perform OTM or write `/oic/sec/*`. Follow-up package work is still required for
|
||||
explicit authentication providers, PSK sessions, Samsung certificate profiles,
|
||||
OwnerPSK derivation, reviewed OCF security codecs, and the separately reviewed
|
||||
authorization/setup policy.
|
||||
|
||||
## Safe evidence for another device report
|
||||
|
||||
Useful public evidence is limited to:
|
||||
|
||||
- retail model without a serial number;
|
||||
- software version;
|
||||
- sanitized candidate ports and first-flight response classes;
|
||||
- redacted `/oic/res`, `/oic/sec/doxm`, and `/oic/sec/pstat` shapes; and
|
||||
- the fixed TLS alert number/name.
|
||||
|
||||
Do not post appliance or owner UUIDs, account identifiers, network addresses,
|
||||
registration values, nonces, certificate fingerprints, credentials, packet
|
||||
captures, or raw exception traces.
|
||||
@@ -166,7 +166,7 @@ def main():
|
||||
stopping.set()
|
||||
logger.info("shutting down…")
|
||||
for b in bridges:
|
||||
b.stop.set()
|
||||
b.request_stop()
|
||||
try: b.set_availability(False)
|
||||
except Exception: pass
|
||||
|
||||
|
||||
+95
-58
@@ -23,19 +23,21 @@ import time
|
||||
|
||||
import cbor2
|
||||
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession, fmt_code
|
||||
from smartthings_local.protocol.dtls_probe import probe
|
||||
|
||||
from smartthings_local.ocf.keepalive import KeepaliveTask
|
||||
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
|
||||
from smartthings_local.ocf.poll_scheduler import PollScheduler
|
||||
from smartthings_local.ocf.state_cache import StateCache
|
||||
from smartthings_local.protocol.dtls_probe import (
|
||||
AMBIGUOUS,
|
||||
probe_dtls_port,
|
||||
probe_dtls_ports,
|
||||
)
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession, fmt_code
|
||||
|
||||
from .descriptor import ApplianceDescriptor, bridge_diagnostic_discovery
|
||||
from .config import ApplianceConfig, SharedConfig
|
||||
from .descriptor import ApplianceDescriptor, bridge_diagnostic_discovery
|
||||
from .logger import bridge_logger
|
||||
|
||||
|
||||
DEBUG_BRIDGE = os.environ.get('DEBUG_BRIDGE') == '1'
|
||||
|
||||
|
||||
@@ -70,20 +72,22 @@ OBSERVE_REFRESH_INTERVAL_S = 6 * 3600.0
|
||||
# orphan otherwise lingers 5-15 min.
|
||||
DTLS_LOCAL_PORT_BASE = 49700
|
||||
|
||||
# SmartThings appliances bind their OCF CoAP-DTLS control port in this
|
||||
# dynamic band (dryer/fridge 49155, oven 49154). When OCF_PORT is unset we
|
||||
# race a stateless ClientHello across the band to find the live one instead
|
||||
# of trusting a single hardcoded default.
|
||||
OCF_PORT_BAND = range(49153, 49157)
|
||||
# Samsung's RT-OCF appliances commonly bind CoAP-DTLS in this dynamic band,
|
||||
# while full-Tizen OCF-PKI appliances also use the standard secure CoAP port.
|
||||
# When OCF_PORT is unset, probe both profiles instead of assuming one fleet-
|
||||
# wide port layout.
|
||||
OCF_PORT_BAND = range(49152, 49161)
|
||||
OCF_STANDARD_SECURE_PORT = 5684
|
||||
|
||||
# The pre-flight liveness gate tolerates one dropped ClientHello (retries=1
|
||||
# → ~1 RTT when the device answers, ~2.6 s to call a silent port DEAD),
|
||||
# → ~1 RTT when the device answers, ~4 s to call a silent port DEAD),
|
||||
# which is far cheaper than eating the 12 s HANDSHAKE_TIMEOUT_S on a
|
||||
# rebooting device or a wrong port. It is stateless (stops at
|
||||
# HelloVerifyRequest), so it leaves no association on the device and the
|
||||
# fixed-source-port reconnect invariant is untouched (see session_once).
|
||||
_GATE_RETRIES = 1
|
||||
_GATE_TIMEOUT_S = 4.0
|
||||
_WORKER_JOIN_TIMEOUT_S = 2.0
|
||||
|
||||
|
||||
class PushBridge:
|
||||
@@ -120,6 +124,8 @@ class PushBridge:
|
||||
self.last_cycle_pub = None
|
||||
self.last_avail_pub: str | None = None
|
||||
self.stop = threading.Event()
|
||||
self._session_stop_lock = threading.Lock()
|
||||
self._session_stop: threading.Event | None = None
|
||||
self.started_ts = time.time()
|
||||
self.session_started_ts = None
|
||||
self.last_change_ts = None
|
||||
@@ -172,6 +178,13 @@ class PushBridge:
|
||||
app.topic_prefix, shared.HA_DISCOVERY_PREFIX, app.device_name,
|
||||
model=descriptor.name.title()))
|
||||
|
||||
def request_stop(self) -> None:
|
||||
"""Stop the bridge and wake workers belonging to its current session."""
|
||||
self.stop.set()
|
||||
with self._session_stop_lock:
|
||||
if self._session_stop is not None:
|
||||
self._session_stop.set()
|
||||
|
||||
# ---- cache plumbing ---------------------------------------------
|
||||
|
||||
def _on_cache_change(self, changed: bool, source: str) -> None:
|
||||
@@ -263,35 +276,20 @@ class PushBridge:
|
||||
# ---- session lifecycle ------------------------------------------
|
||||
|
||||
def _candidate_ports(self) -> list[int]:
|
||||
"""The OCF band plus the descriptor's documented default, deduped
|
||||
and ordered — the search space when OCF_PORT is unset."""
|
||||
return sorted(set(OCF_PORT_BAND) | {self.descriptor.default_observe_port})
|
||||
"""Known OCF secure ports plus the descriptor default, in order."""
|
||||
return sorted(
|
||||
set(OCF_PORT_BAND)
|
||||
| {OCF_STANDARD_SECURE_PORT, self.descriptor.default_observe_port}
|
||||
)
|
||||
|
||||
def _race_probe(self, candidates: list[int]) -> int | None:
|
||||
"""Race a stateless ClientHello across all candidates in parallel
|
||||
and return the first port that answers LIVE — without waiting for
|
||||
the dead ones to burn their full retry budget. Returns None if none
|
||||
answer.
|
||||
|
||||
The winner comes back in ~1 RTT; the losing probes are abandoned
|
||||
(shutdown(wait=False)) and each just runs out its own ~timeout loop
|
||||
and closes its own socket in finally. This is a latency win, not a
|
||||
correctness need — unlike #212's full-handshake race the losers are
|
||||
bounded at a few seconds, not 12 s. Real appliances expose exactly
|
||||
one DTLS port, so first-to-answer is unambiguous."""
|
||||
import concurrent.futures as cf
|
||||
ex = cf.ThreadPoolExecutor(max_workers=len(candidates))
|
||||
try:
|
||||
futs = [ex.submit(probe, self.app.ip, p,
|
||||
retries=_GATE_RETRIES, timeout=_GATE_TIMEOUT_S)
|
||||
for p in candidates]
|
||||
for fut in cf.as_completed(futs):
|
||||
r = fut.result()
|
||||
if r.is_dtls_server:
|
||||
return r.port
|
||||
return None
|
||||
finally:
|
||||
ex.shutdown(wait=False)
|
||||
def _probe_candidates(self, candidates: list[int]):
|
||||
"""Probe all candidates inside one budget and preserve ambiguity."""
|
||||
return probe_dtls_ports(
|
||||
self.app.ip,
|
||||
tuple(candidates),
|
||||
retries=_GATE_RETRIES,
|
||||
timeout=_GATE_TIMEOUT_S,
|
||||
)
|
||||
|
||||
def _resolve_port(self) -> int:
|
||||
"""Return a port that just answered a stateless DTLS ClientHello,
|
||||
@@ -307,30 +305,40 @@ class PushBridge:
|
||||
first on the next reconnect and rediscovered only if it goes DEAD."""
|
||||
pinned = self.app.ocf_port
|
||||
if pinned is not None:
|
||||
r = probe(self.app.ip, pinned,
|
||||
retries=_GATE_RETRIES, timeout=_GATE_TIMEOUT_S)
|
||||
r = probe_dtls_port(
|
||||
self.app.ip,
|
||||
pinned,
|
||||
retries=_GATE_RETRIES,
|
||||
timeout=_GATE_TIMEOUT_S,
|
||||
)
|
||||
if not r.is_dtls_server:
|
||||
raise ConnectionError(
|
||||
f"port {pinned} not a live DTLS server ({r.outcome})")
|
||||
raise ConnectionError('configured port is not a DTLS server')
|
||||
return pinned
|
||||
|
||||
# A previously discovered port is almost certainly still the one —
|
||||
# try it alone first and only fall back to a full band re-race if
|
||||
# try it alone first and only fall back to the full candidate set if
|
||||
# it has gone silent (firmware moved it, or it was never right).
|
||||
if self._discovered_port is not None:
|
||||
r = probe(self.app.ip, self._discovered_port,
|
||||
retries=_GATE_RETRIES, timeout=_GATE_TIMEOUT_S)
|
||||
r = probe_dtls_port(
|
||||
self.app.ip,
|
||||
self._discovered_port,
|
||||
retries=_GATE_RETRIES,
|
||||
timeout=_GATE_TIMEOUT_S,
|
||||
)
|
||||
if r.is_dtls_server:
|
||||
return self._discovered_port
|
||||
self._discovered_port = None
|
||||
|
||||
candidates = self._candidate_ports()
|
||||
live = self._race_probe(candidates)
|
||||
if live is None:
|
||||
raise ConnectionError(f"no live DTLS server across {candidates}")
|
||||
self.log.info("discovered DTLS port %d", live)
|
||||
self._discovered_port = live
|
||||
return live
|
||||
selection = self._probe_candidates(candidates)
|
||||
if selection.outcome == AMBIGUOUS:
|
||||
raise ConnectionError(
|
||||
'multiple DTLS listeners answered; configure OCF_PORT')
|
||||
if selection.selected_port is None:
|
||||
raise ConnectionError('no live DTLS server found')
|
||||
self.log.info("discovered DTLS port %d", selection.selected_port)
|
||||
self._discovered_port = selection.selected_port
|
||||
return selection.selected_port
|
||||
|
||||
def session_once(self):
|
||||
port = self._resolve_port()
|
||||
@@ -425,25 +433,54 @@ class PushBridge:
|
||||
self.keepalive = keepalive
|
||||
self.observe_refresh = observe_refresh
|
||||
|
||||
# These workers belong to this DTLS session, not to the bridge
|
||||
# process. A reconnect must retire them before the replacement
|
||||
# session starts or they continue operating on the closed session.
|
||||
session_stop = threading.Event()
|
||||
with self._session_stop_lock:
|
||||
self._session_stop = session_stop
|
||||
# ``request_stop()`` sets the bridge event before taking this
|
||||
# lock. Checking it while publishing the handle prevents a lost
|
||||
# wakeup if shutdown races this session handoff.
|
||||
if self.stop.is_set():
|
||||
session_stop.set()
|
||||
|
||||
sched_t = threading.Thread(
|
||||
target=scheduler.run_forever, args=(self.stop,),
|
||||
target=scheduler.run_forever, args=(session_stop,),
|
||||
daemon=True, name=f'{self.app.klass}-poll')
|
||||
ka_t = threading.Thread(
|
||||
target=keepalive.run_forever, args=(self.stop,),
|
||||
target=keepalive.run_forever, args=(session_stop,),
|
||||
daemon=True, name=f'{self.app.klass}-ping')
|
||||
ref_t = threading.Thread(
|
||||
target=observe_refresh.run_forever, args=(self.stop,),
|
||||
target=observe_refresh.run_forever, args=(session_stop,),
|
||||
daemon=True, name=f'{self.app.klass}-obsref')
|
||||
sched_t.start()
|
||||
ka_t.start()
|
||||
ref_t.start()
|
||||
workers = (sched_t, ka_t, ref_t)
|
||||
started_workers = []
|
||||
|
||||
try:
|
||||
for worker in workers:
|
||||
worker.start()
|
||||
started_workers.append(worker)
|
||||
sess.join()
|
||||
finally:
|
||||
# A worker already inside a tick can finish after the reader
|
||||
# exits. Disable old-session reachability callbacks first so it
|
||||
# cannot change availability after a replacement takes over.
|
||||
keepalive.on_reachable = None
|
||||
keepalive.on_unreachable = None
|
||||
session_stop.set()
|
||||
join_deadline = time.monotonic() + _WORKER_JOIN_TIMEOUT_S
|
||||
for worker in started_workers:
|
||||
worker.join(max(0.0, join_deadline - time.monotonic()))
|
||||
if worker.is_alive():
|
||||
self.log.warning(
|
||||
"session worker did not stop: %s", worker.name)
|
||||
self.scheduler = None
|
||||
self.keepalive = None
|
||||
self.observe_refresh = None
|
||||
with self._session_stop_lock:
|
||||
if self._session_stop is session_stop:
|
||||
self._session_stop = None
|
||||
|
||||
def _seed_from_device0(self, sess):
|
||||
code, pl = sess.get(self.descriptor.seed_path, timeout=15.0)
|
||||
|
||||
@@ -166,6 +166,14 @@ def flatten(links):
|
||||
if temps_items:
|
||||
cur_c = _int(temps_items[0].get('x.com.samsung.da.current'))
|
||||
des_c = _int(temps_items[0].get('x.com.samsung.da.desired'))
|
||||
# With no cycle set the oven reports desired=0. That means "no
|
||||
# setpoint", not a 0 °C target, and HA rejects it against the Number
|
||||
# entity's 30-270 range on every publish. Anything outside the
|
||||
# settable band is absent, not a value: null lands as unknown on both
|
||||
# the Number and the Setpoint sensor, the way completion_minutes
|
||||
# already reads when idle. _setpoint applies the same bounds on write.
|
||||
if des_c is not None and not (SETPOINT_MIN_C <= des_c <= SETPOINT_MAX_C):
|
||||
des_c = None
|
||||
|
||||
# Door
|
||||
doors_items = g('/doors/vs/0', 'x.com.samsung.da.items') or []
|
||||
|
||||
+3
-2
@@ -9,7 +9,7 @@ description = "Local CoAP-over-DTLS (OCF) transport + polling layer for Samsung
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
license = "MIT"
|
||||
license-files = ["LICENSE"]
|
||||
license-files = ["LICENSE", "NOTICE"]
|
||||
authors = [{ name = "Jack Nagy" }]
|
||||
keywords = ["smartthings", "samsung", "ocf", "coap", "dtls", "home-assistant", "iot"]
|
||||
classifiers = [
|
||||
@@ -21,7 +21,7 @@ classifiers = [
|
||||
]
|
||||
dependencies = [
|
||||
"cbor2>=5.6",
|
||||
"pyOpenSSL>=23.0",
|
||||
"pyOpenSSL>=23.1",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
@@ -60,5 +60,6 @@ include = [
|
||||
"tests",
|
||||
"README.md",
|
||||
"LICENSE",
|
||||
"NOTICE",
|
||||
"pyproject.toml",
|
||||
]
|
||||
|
||||
+81
-7
@@ -184,8 +184,57 @@ def verify_cert_key_pair(cert_path, key_path):
|
||||
f"AC14K_M cert and key do not pair (cert modulus != key modulus)")
|
||||
|
||||
|
||||
# OpenSSL config that force-enables SHA-1 signatures. Fedora/RHEL (and some
|
||||
# other hardened OpenSSL 3.x builds) reject SHA-1 signing under the default
|
||||
# crypto policy, but the AC14K_M trust chain requires a SHA-1-signed leaf, so
|
||||
# we re-enable it just for the signing step via a scoped OPENSSL_CONF.
|
||||
SHA1_OVERRIDE_CONF = """\
|
||||
openssl_conf = openssl_init
|
||||
|
||||
[openssl_init]
|
||||
alg_section = evp_properties
|
||||
|
||||
[evp_properties]
|
||||
rh-allow-sha1-signatures = yes
|
||||
"""
|
||||
|
||||
|
||||
class CommandError(RuntimeError):
|
||||
"""A subprocess exited non-zero; carries the command and its output."""
|
||||
|
||||
|
||||
def run(cmd, **kw):
|
||||
return subprocess.run(cmd, check=True, capture_output=True, text=True, **kw)
|
||||
proc = subprocess.run(cmd, capture_output=True, text=True, **kw)
|
||||
if proc.returncode != 0:
|
||||
detail = (proc.stderr or proc.stdout or '').strip()
|
||||
raise CommandError(
|
||||
f"command failed (exit {proc.returncode}): {' '.join(cmd)}"
|
||||
+ (f"\n{detail}" if detail else ""))
|
||||
return proc
|
||||
|
||||
|
||||
def run_allow_sha1(cmd):
|
||||
"""Run an openssl command with SHA-1 signatures force-enabled, for
|
||||
distros whose crypto policy otherwise blocks SHA-1 signing."""
|
||||
version = run(['openssl', 'version']).stdout.strip()
|
||||
if not version.startswith('OpenSSL 3.'):
|
||||
# The provider configuration below is specific to OpenSSL 3.
|
||||
# LibreSSL can exit successfully without running the requested
|
||||
# command when it is given that configuration, leaving no output
|
||||
# certificate behind. Older OpenSSL releases do not need the
|
||||
# provider override either, so retry them with a clean environment.
|
||||
env = dict(os.environ)
|
||||
env.pop('OPENSSL_CONF', None)
|
||||
return run(cmd, env=env)
|
||||
|
||||
conf = tempfile.NamedTemporaryFile(
|
||||
'w', suffix='.cnf', prefix='sha1_ok_', delete=False)
|
||||
conf.write(SHA1_OVERRIDE_CONF)
|
||||
conf.close()
|
||||
try:
|
||||
return run(cmd, env=dict(os.environ, OPENSSL_CONF=conf.name))
|
||||
finally:
|
||||
os.unlink(conf.name)
|
||||
|
||||
|
||||
def mint_cert(uuid, ac14k_cert, ac14k_key, chain_files, out_dir):
|
||||
@@ -229,11 +278,23 @@ DNS.1 = {uuid}
|
||||
run(['openssl', 'req', '-new', '-key', str(paths['key']),
|
||||
'-out', str(paths['csr']), '-subj', subject])
|
||||
|
||||
run(['openssl', 'x509', '-req', '-in', str(paths['csr']),
|
||||
'-CA', str(ac14k_cert), '-CAkey', str(ac14k_key),
|
||||
'-CAcreateserial', '-CAserial', str(paths['srl']),
|
||||
'-out', str(paths['leaf']), '-days', '3650',
|
||||
'-extfile', str(paths['ext']), '-sha1'])
|
||||
sign_cmd = ['openssl', 'x509', '-req', '-in', str(paths['csr']),
|
||||
'-CA', str(ac14k_cert), '-CAkey', str(ac14k_key),
|
||||
'-CAcreateserial', '-CAserial', str(paths['srl']),
|
||||
'-out', str(paths['leaf']), '-days', '3650',
|
||||
'-extfile', str(paths['ext']), '-sha1']
|
||||
try:
|
||||
run(sign_cmd)
|
||||
except CommandError as first:
|
||||
# Most likely the local crypto policy blocks SHA-1 signing
|
||||
# (common on Fedora/RHEL). Retry once with SHA-1 force-enabled;
|
||||
# if that still fails, surface the original error.
|
||||
print(" SHA-1 signing was rejected by the local OpenSSL policy; "
|
||||
"retrying with a SHA-1 override...")
|
||||
try:
|
||||
run_allow_sha1(sign_cmd)
|
||||
except CommandError:
|
||||
raise first
|
||||
|
||||
parts = [paths['leaf'].read_text()]
|
||||
for p in chain_files:
|
||||
@@ -459,7 +520,20 @@ def main():
|
||||
print("=" * 60)
|
||||
print(f"Phase 3: mint client cert with UUID {uuid}")
|
||||
print("=" * 60)
|
||||
paths = mint_cert(uuid, ac14k_cert, ac14k_key, chain_files, out_dir)
|
||||
try:
|
||||
paths = mint_cert(uuid, ac14k_cert, ac14k_key, chain_files, out_dir)
|
||||
except CommandError as e:
|
||||
print(f"\n[!] Failed to mint the client cert:\n{e}", file=sys.stderr)
|
||||
print(
|
||||
"\n If the failure mentions SHA-1 / disabled digests, your "
|
||||
"OpenSSL build blocks SHA-1 signing (common on Fedora/RHEL).\n"
|
||||
" The AC14K_M chain requires SHA-1, so allow it and re-run:\n"
|
||||
" sudo update-crypto-policies --set DEFAULT:SHA1\n"
|
||||
" (or LEGACY). Undo afterwards with: "
|
||||
"sudo update-crypto-policies --set DEFAULT",
|
||||
file=sys.stderr)
|
||||
return 4
|
||||
|
||||
print(f" key: {paths['key']}")
|
||||
print(f" leaf: {paths['leaf']}")
|
||||
print(f" fullchain: {paths['fullchain']}")
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
"""Public, redacted exception types for smartthings-local."""
|
||||
|
||||
__all__ = [
|
||||
'AuthenticationError',
|
||||
'AuthorizationError',
|
||||
'BlockwiseError',
|
||||
'EndpointError',
|
||||
'MalformedMessageError',
|
||||
'ObserveError',
|
||||
'ProbeError',
|
||||
'SessionClosedError',
|
||||
'SessionError',
|
||||
'SessionTimeoutError',
|
||||
'SmartThingsLocalError',
|
||||
]
|
||||
|
||||
|
||||
class SmartThingsLocalError(Exception):
|
||||
"""Base class for classified library failures.
|
||||
|
||||
Subclasses expose a stable ``code`` and a fixed, non-sensitive message.
|
||||
They intentionally do not accept arbitrary detail because backend errors
|
||||
can contain remote endpoints, local paths, or credential metadata.
|
||||
"""
|
||||
|
||||
code = 'smartthings_local_error'
|
||||
message = 'SmartThings Local operation failed'
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(self.message)
|
||||
|
||||
def __repr__(self):
|
||||
return f'{type(self).__name__}(code={self.code!r})'
|
||||
|
||||
|
||||
class EndpointError(SmartThingsLocalError, OSError):
|
||||
"""An endpoint could not be resolved, bound, or connected."""
|
||||
|
||||
code = 'endpoint'
|
||||
message = 'endpoint operation failed'
|
||||
|
||||
|
||||
class ProbeError(SmartThingsLocalError, ConnectionError):
|
||||
"""A DTLS probe failed before producing a protocol result."""
|
||||
|
||||
code = 'probe'
|
||||
message = 'DTLS probe failed'
|
||||
|
||||
|
||||
class SessionError(SmartThingsLocalError, ConnectionError):
|
||||
"""A connected-session operation failed."""
|
||||
|
||||
code = 'session'
|
||||
message = 'session operation failed'
|
||||
|
||||
|
||||
class AuthenticationError(SessionError):
|
||||
"""The peer or local credentials could not be authenticated."""
|
||||
|
||||
code = 'authentication'
|
||||
message = 'authentication failed'
|
||||
|
||||
|
||||
class AuthorizationError(SmartThingsLocalError, PermissionError):
|
||||
"""The authenticated peer is not authorized for an operation."""
|
||||
|
||||
code = 'authorization'
|
||||
message = 'operation is not authorized'
|
||||
|
||||
|
||||
class SessionTimeoutError(SmartThingsLocalError, TimeoutError):
|
||||
"""A bounded session operation exceeded its deadline."""
|
||||
|
||||
code = 'timeout'
|
||||
message = 'session operation timed out'
|
||||
|
||||
|
||||
class SessionClosedError(SessionError):
|
||||
"""An operation was attempted on a closed session."""
|
||||
|
||||
code = 'session_closed'
|
||||
message = 'session is closed'
|
||||
|
||||
|
||||
class MalformedMessageError(SmartThingsLocalError, ValueError):
|
||||
"""A protocol message could not be decoded safely."""
|
||||
|
||||
code = 'malformed_message'
|
||||
message = 'malformed protocol message'
|
||||
|
||||
|
||||
class BlockwiseError(SessionError):
|
||||
"""A Block1 or Block2 transfer violated its bounded contract."""
|
||||
|
||||
code = 'blockwise'
|
||||
message = 'blockwise transfer failed'
|
||||
|
||||
|
||||
class ObserveError(SessionError):
|
||||
"""A CoAP Observe relation could not be established or maintained."""
|
||||
|
||||
code = 'observe'
|
||||
message = 'Observe relation failed'
|
||||
@@ -0,0 +1,571 @@
|
||||
"""Immutable authentication providers for DTLS sessions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import warnings
|
||||
from enum import Enum
|
||||
from os import PathLike
|
||||
from pathlib import Path
|
||||
from typing import Protocol, runtime_checkable
|
||||
from uuid import UUID
|
||||
|
||||
from cryptography.x509.oid import ExtensionOID
|
||||
from OpenSSL import SSL, _util, crypto
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_OCF_ROOT_CA = str(Path(__file__).with_name("ocf_root_ca.pem"))
|
||||
_DTLS_CIPHERS = b"ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0"
|
||||
_DTLS_PSK_CIPHERS = b"ECDHE-PSK-AES128-CBC-SHA256:@SECLEVEL=0"
|
||||
_SAMSUNG_SERVER_CURVES = b"prime256v1"
|
||||
_SAMSUNG_SERVER_SIGNATURE_ALGORITHMS = (
|
||||
b"RSA+SHA256:ECDSA+SHA256:RSA+SHA1:ECDSA+SHA1"
|
||||
)
|
||||
_SAMSUNG_SERVER_CN_RE = re.compile(
|
||||
r"\AOCF Device: [^()\r\n]{1,96} "
|
||||
r"\((?P<device_identity>[0-9a-f]{8}-[0-9a-f]{4}-"
|
||||
r"[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12})\)\Z",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_PSK_CLIENT_CALLBACK_CDEF = (
|
||||
"unsigned int (*)(SSL *, char *, char *, unsigned int, "
|
||||
"unsigned char *, unsigned int)"
|
||||
)
|
||||
_PEM_CERT_RE = re.compile(
|
||||
rb"-----BEGIN CERTIFICATE-----.*?-----END CERTIFICATE-----",
|
||||
re.DOTALL,
|
||||
)
|
||||
|
||||
|
||||
def _verify_peer(_connection, _certificate, _error, _depth, ok):
|
||||
"""Keep pyOpenSSL's existing verification result unchanged."""
|
||||
return ok
|
||||
|
||||
|
||||
def _load_pem_chain(ctx: SSL.Context, cert_pem: str, key_pem: str) -> None:
|
||||
"""Load a PEM certificate chain and private key into a context in memory."""
|
||||
certificates = _PEM_CERT_RE.findall(cert_pem.encode())
|
||||
if not certificates:
|
||||
raise ValueError("No certificates found in cert_pem")
|
||||
ctx.use_certificate(
|
||||
crypto.load_certificate(crypto.FILETYPE_PEM, certificates[0])
|
||||
)
|
||||
for extra in certificates[1:]:
|
||||
ctx.add_extra_chain_cert(
|
||||
crypto.load_certificate(crypto.FILETYPE_PEM, extra)
|
||||
)
|
||||
ctx.use_privatekey(
|
||||
crypto.load_privatekey(crypto.FILETYPE_PEM, key_pem.encode())
|
||||
)
|
||||
ctx.check_privatekey()
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class AuthenticationProvider(Protocol):
|
||||
"""Configure authentication for a newly created DTLS context."""
|
||||
|
||||
def configure_context(self, context: SSL.Context) -> None:
|
||||
"""Configure a context while this provider remains session-owned."""
|
||||
|
||||
|
||||
class SamsungServerRole(Enum):
|
||||
"""Known Samsung OCF hardware-certificate subject roles."""
|
||||
|
||||
HOME_APPLIANCE = "OCF HA Device"
|
||||
VD_DEVICE = "OCF VD Device"
|
||||
|
||||
|
||||
class SamsungServerProfile:
|
||||
"""Opt-in Samsung hardware-certificate verification profile."""
|
||||
|
||||
__slots__ = (
|
||||
"_additional_ca_certificates",
|
||||
"_expected_certificate_identity",
|
||||
"_role",
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
expected_certificate_identity: UUID | str,
|
||||
role: SamsungServerRole = SamsungServerRole.HOME_APPLIANCE,
|
||||
additional_ca_pem: str | None = None,
|
||||
) -> None:
|
||||
if type(expected_certificate_identity) is UUID:
|
||||
parsed_identity = expected_certificate_identity
|
||||
elif type(expected_certificate_identity) is str:
|
||||
try:
|
||||
parsed_identity = UUID(expected_certificate_identity)
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
"expected_certificate_identity must be a canonical "
|
||||
"non-zero UUID"
|
||||
) from None
|
||||
if expected_certificate_identity != str(parsed_identity):
|
||||
raise ValueError(
|
||||
"expected_certificate_identity must be a canonical "
|
||||
"non-zero UUID"
|
||||
)
|
||||
else:
|
||||
raise TypeError(
|
||||
"expected_certificate_identity must be a UUID or string"
|
||||
)
|
||||
if parsed_identity.int == 0:
|
||||
raise ValueError(
|
||||
"expected_certificate_identity must be a canonical "
|
||||
"non-zero UUID"
|
||||
)
|
||||
if type(role) is not SamsungServerRole:
|
||||
raise TypeError("role must be a SamsungServerRole")
|
||||
|
||||
certificates: tuple[bytes, ...] = ()
|
||||
if additional_ca_pem is not None:
|
||||
if type(additional_ca_pem) is not str:
|
||||
raise TypeError("additional_ca_pem must be a string")
|
||||
try:
|
||||
raw_ca_pem = additional_ca_pem.encode("ascii")
|
||||
except UnicodeEncodeError:
|
||||
raise ValueError(
|
||||
"additional_ca_pem must contain ASCII PEM certificates"
|
||||
) from None
|
||||
parsed_certificates = tuple(_PEM_CERT_RE.findall(raw_ca_pem))
|
||||
if (
|
||||
not 1 <= len(parsed_certificates) <= 4
|
||||
or len(raw_ca_pem) > 32 * 1024
|
||||
or _PEM_CERT_RE.sub(b"", raw_ca_pem).strip()
|
||||
):
|
||||
raise ValueError(
|
||||
"additional_ca_pem must contain one to four PEM certificates"
|
||||
)
|
||||
try:
|
||||
loaded_certificates = [
|
||||
crypto.load_certificate(crypto.FILETYPE_PEM, certificate)
|
||||
for certificate in parsed_certificates
|
||||
]
|
||||
basic_constraints = [
|
||||
[
|
||||
extension
|
||||
for extension in certificate.to_cryptography().extensions
|
||||
if extension.oid == ExtensionOID.BASIC_CONSTRAINTS
|
||||
]
|
||||
for certificate in loaded_certificates
|
||||
]
|
||||
except (crypto.Error, ValueError):
|
||||
raise ValueError(
|
||||
"additional_ca_pem contains an invalid certificate"
|
||||
) from None
|
||||
if any(
|
||||
len(constraints) != 1 or not constraints[0].value.ca
|
||||
for constraints in basic_constraints
|
||||
):
|
||||
raise ValueError(
|
||||
"additional_ca_pem must contain only CA certificates"
|
||||
)
|
||||
fingerprints = {
|
||||
crypto.dump_certificate(crypto.FILETYPE_ASN1, certificate)
|
||||
for certificate in loaded_certificates
|
||||
}
|
||||
if len(fingerprints) != len(loaded_certificates):
|
||||
raise ValueError(
|
||||
"additional_ca_pem must not contain duplicate certificates"
|
||||
)
|
||||
certificates = parsed_certificates
|
||||
|
||||
object.__setattr__(
|
||||
self,
|
||||
"_expected_certificate_identity",
|
||||
parsed_identity,
|
||||
)
|
||||
object.__setattr__(self, "_role", role)
|
||||
object.__setattr__(self, "_additional_ca_certificates", certificates)
|
||||
|
||||
def __setattr__(self, _name: str, _value: object) -> None:
|
||||
raise AttributeError("SamsungServerProfile is immutable")
|
||||
|
||||
def __delattr__(self, _name: str) -> None:
|
||||
raise AttributeError("SamsungServerProfile is immutable")
|
||||
|
||||
@classmethod
|
||||
def bound_device(
|
||||
cls,
|
||||
expected_certificate_identity: UUID | str,
|
||||
*,
|
||||
role: SamsungServerRole = SamsungServerRole.HOME_APPLIANCE,
|
||||
additional_ca_pem: str | None = None,
|
||||
) -> SamsungServerProfile:
|
||||
"""Bind a verified Samsung hardware leaf to its certificate UUID."""
|
||||
return cls(
|
||||
expected_certificate_identity=expected_certificate_identity,
|
||||
role=role,
|
||||
additional_ca_pem=additional_ca_pem,
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""Return a representation without device or trust-chain details."""
|
||||
return "SamsungServerProfile()"
|
||||
|
||||
def _configure_context(self, context: SSL.Context) -> None:
|
||||
curve_setter = getattr(_util.lib, "SSL_CTX_set1_curves_list", None)
|
||||
if curve_setter is not None:
|
||||
if curve_setter(context._context, _SAMSUNG_SERVER_CURVES) != 1:
|
||||
raise RuntimeError(
|
||||
"OpenSSL rejected the Samsung server certificate profile"
|
||||
)
|
||||
else:
|
||||
# pyOpenSSL 23.1 does not expose SSL_CTX_set1_curves_list. Its
|
||||
# public set_tmp_ecdh fallback produces the same single P-256
|
||||
# supported-groups ClientHello extension; wire-level tests protect
|
||||
# that compatibility path. Newer pyOpenSSL uses the exact setter.
|
||||
try:
|
||||
with warnings.catch_warnings():
|
||||
warnings.simplefilter("ignore", DeprecationWarning)
|
||||
curve = crypto.get_elliptic_curve("prime256v1")
|
||||
context.set_tmp_ecdh(curve)
|
||||
except (AttributeError, TypeError, ValueError, SSL.Error):
|
||||
raise RuntimeError(
|
||||
"OpenSSL rejected the Samsung server certificate profile"
|
||||
) from None
|
||||
|
||||
signature_setter = getattr(
|
||||
_util.lib,
|
||||
"SSL_CTX_set1_sigalgs_list",
|
||||
None,
|
||||
)
|
||||
if (
|
||||
signature_setter is None
|
||||
or signature_setter(
|
||||
context._context,
|
||||
_SAMSUNG_SERVER_SIGNATURE_ALGORITHMS,
|
||||
)
|
||||
!= 1
|
||||
):
|
||||
raise RuntimeError(
|
||||
"OpenSSL rejected the Samsung server certificate profile"
|
||||
)
|
||||
context.set_options(SSL.OP_NO_TICKET)
|
||||
|
||||
if self._additional_ca_certificates:
|
||||
store = context.get_cert_store()
|
||||
try:
|
||||
for certificate in self._additional_ca_certificates:
|
||||
store.add_cert(
|
||||
crypto.load_certificate(
|
||||
crypto.FILETYPE_PEM,
|
||||
certificate,
|
||||
)
|
||||
)
|
||||
except crypto.Error:
|
||||
raise RuntimeError(
|
||||
"OpenSSL rejected the Samsung server trust profile"
|
||||
) from None
|
||||
|
||||
def _verify_peer(
|
||||
self,
|
||||
_connection,
|
||||
certificate,
|
||||
_error,
|
||||
depth,
|
||||
ok,
|
||||
) -> bool:
|
||||
if not ok or certificate is None or depth < 0:
|
||||
return False
|
||||
if depth > 0:
|
||||
return True
|
||||
try:
|
||||
with warnings.catch_warnings():
|
||||
# pyOpenSSL deprecates this API in favor of cryptography, but
|
||||
# reparsing Samsung's non-DER factory leaves with cryptography
|
||||
# rejects certificates that OpenSSL has already verified.
|
||||
warnings.simplefilter("ignore", DeprecationWarning)
|
||||
components = [
|
||||
(name.decode("ascii"), value.decode("ascii"))
|
||||
for name, value in certificate.get_subject().get_components()
|
||||
]
|
||||
except (
|
||||
AttributeError,
|
||||
TypeError,
|
||||
UnicodeDecodeError,
|
||||
ValueError,
|
||||
crypto.Error,
|
||||
):
|
||||
logger.warning("Unable to parse Samsung server certificate subject")
|
||||
return False
|
||||
|
||||
common_names = [value for name, value in components if name == "CN"]
|
||||
organizational_units = [
|
||||
value for name, value in components if name == "OU"
|
||||
]
|
||||
organizations = [value for name, value in components if name == "O"]
|
||||
countries = [value for name, value in components if name == "C"]
|
||||
# This deliberately pins the complete Samsung subject role. The OCF
|
||||
# reference implementation reads only the UUID-bearing CN, but
|
||||
# relaxing C/O/OU here could accept a different certificate cohort.
|
||||
if (
|
||||
len(common_names) != 1
|
||||
or organizational_units != [self._role.value]
|
||||
or organizations != ["Samsung Electronics"]
|
||||
or countries != ["KR"]
|
||||
):
|
||||
return False
|
||||
match = _SAMSUNG_SERVER_CN_RE.fullmatch(common_names[0])
|
||||
return (
|
||||
match is not None
|
||||
and UUID(match.group("device_identity"))
|
||||
== self._expected_certificate_identity
|
||||
)
|
||||
|
||||
|
||||
def _configure_certificate_server(
|
||||
context: SSL.Context,
|
||||
server_profile: SamsungServerProfile | None,
|
||||
) -> None:
|
||||
"""Configure certificate-server verification for one DTLS context."""
|
||||
context.load_verify_locations(_OCF_ROOT_CA)
|
||||
if server_profile is None:
|
||||
context.set_verify(SSL.VERIFY_PEER, _verify_peer)
|
||||
else:
|
||||
server_profile._configure_context(context)
|
||||
context.set_verify(
|
||||
SSL.VERIFY_PEER,
|
||||
server_profile._verify_peer,
|
||||
)
|
||||
# @SECLEVEL=0 permits SHA-1 in Samsung's server cert chain (AC14K_M
|
||||
# intermediate is SHA-1 signed). This is the only channel that reaches
|
||||
# the OpenSSL instance cryptography bundles; ctypes and cffi bindings
|
||||
# do not expose SSL_CTX_set_security_level on this build.
|
||||
context.set_cipher_list(_DTLS_CIPHERS)
|
||||
|
||||
|
||||
class ServerCertificateAuth:
|
||||
"""Verify a pinned Samsung server without a client certificate."""
|
||||
|
||||
__slots__ = ("_server_profile",)
|
||||
|
||||
def __init__(self, *, server_profile: SamsungServerProfile) -> None:
|
||||
if type(server_profile) is not SamsungServerProfile:
|
||||
raise TypeError("server_profile must be a SamsungServerProfile")
|
||||
object.__setattr__(self, "_server_profile", server_profile)
|
||||
|
||||
def __setattr__(self, _name: str, _value: object) -> None:
|
||||
raise AttributeError("ServerCertificateAuth is immutable")
|
||||
|
||||
def __delattr__(self, _name: str) -> None:
|
||||
raise AttributeError("ServerCertificateAuth is immutable")
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""Return a representation without server identity or trust details."""
|
||||
return "ServerCertificateAuth()"
|
||||
|
||||
def configure_context(self, context: SSL.Context) -> None:
|
||||
"""Verify the selected server profile without loading client material."""
|
||||
_configure_certificate_server(context, self._server_profile)
|
||||
|
||||
|
||||
class CertificateAuth:
|
||||
"""Certificate authentication loaded from files or in-memory PEM data.
|
||||
|
||||
Use :meth:`from_files` or :meth:`from_memory` to create an instance.
|
||||
Credential sources are intentionally not exposed as public attributes.
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"_certificate_path",
|
||||
"_certificate_pem",
|
||||
"_private_key_path",
|
||||
"_private_key_pem",
|
||||
"_server_profile",
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
certificate_path: str | PathLike[str] | None = None,
|
||||
private_key_path: str | PathLike[str] | None = None,
|
||||
certificate_pem: str | None = None,
|
||||
private_key_pem: str | None = None,
|
||||
server_profile: SamsungServerProfile | None = None,
|
||||
) -> None:
|
||||
file_supplied = (
|
||||
certificate_path is not None or private_key_path is not None
|
||||
)
|
||||
memory_supplied = (
|
||||
certificate_pem is not None or private_key_pem is not None
|
||||
)
|
||||
if file_supplied and memory_supplied:
|
||||
raise ValueError(
|
||||
"pass either certificate_path/private_key_path or "
|
||||
"certificate_pem/private_key_pem, not both"
|
||||
)
|
||||
if file_supplied:
|
||||
if certificate_path is None or private_key_path is None:
|
||||
raise ValueError(
|
||||
"certificate_path and private_key_path must be passed together"
|
||||
)
|
||||
elif memory_supplied:
|
||||
if certificate_pem is None or private_key_pem is None:
|
||||
raise ValueError(
|
||||
"certificate_pem and private_key_pem must be passed together"
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"must pass either certificate_path/private_key_path or "
|
||||
"certificate_pem/private_key_pem"
|
||||
)
|
||||
if (
|
||||
server_profile is not None
|
||||
and type(server_profile) is not SamsungServerProfile
|
||||
):
|
||||
raise TypeError("server_profile must be a SamsungServerProfile")
|
||||
object.__setattr__(
|
||||
self,
|
||||
"_certificate_path",
|
||||
str(certificate_path) if certificate_path is not None else None,
|
||||
)
|
||||
object.__setattr__(
|
||||
self,
|
||||
"_private_key_path",
|
||||
str(private_key_path) if private_key_path is not None else None,
|
||||
)
|
||||
object.__setattr__(self, "_certificate_pem", certificate_pem)
|
||||
object.__setattr__(self, "_private_key_pem", private_key_pem)
|
||||
object.__setattr__(self, "_server_profile", server_profile)
|
||||
|
||||
def __setattr__(self, _name: str, _value: object) -> None:
|
||||
raise AttributeError("CertificateAuth is immutable")
|
||||
|
||||
def __delattr__(self, _name: str) -> None:
|
||||
raise AttributeError("CertificateAuth is immutable")
|
||||
|
||||
@classmethod
|
||||
def from_files(
|
||||
cls,
|
||||
certificate_path: str | PathLike[str],
|
||||
private_key_path: str | PathLike[str],
|
||||
*,
|
||||
server_profile: SamsungServerProfile | None = None,
|
||||
) -> CertificateAuth:
|
||||
"""Create a provider backed by certificate-chain and key files."""
|
||||
return cls(
|
||||
certificate_path=certificate_path,
|
||||
private_key_path=private_key_path,
|
||||
server_profile=server_profile,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_memory(
|
||||
cls,
|
||||
certificate_pem: str,
|
||||
private_key_pem: str,
|
||||
*,
|
||||
server_profile: SamsungServerProfile | None = None,
|
||||
) -> CertificateAuth:
|
||||
"""Create a provider backed by an in-memory PEM chain and key."""
|
||||
return cls(
|
||||
certificate_pem=certificate_pem,
|
||||
private_key_pem=private_key_pem,
|
||||
server_profile=server_profile,
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""Return a representation that never includes credential material."""
|
||||
return "CertificateAuth()"
|
||||
|
||||
def configure_context(self, context: SSL.Context) -> None:
|
||||
"""Apply the existing certificate authentication profile to a context."""
|
||||
_configure_certificate_server(context, self._server_profile)
|
||||
if self._certificate_pem is not None:
|
||||
_load_pem_chain(
|
||||
context,
|
||||
self._certificate_pem,
|
||||
self._private_key_pem,
|
||||
)
|
||||
else:
|
||||
context.use_certificate_chain_file(self._certificate_path)
|
||||
context.use_privatekey_file(self._private_key_path)
|
||||
context.check_privatekey()
|
||||
|
||||
|
||||
class PskAuth:
|
||||
"""DTLS authentication using an existing OCF PSK credential.
|
||||
|
||||
The identity must be a raw 16-byte OCF UUID. The key must contain 16 or
|
||||
32 bytes. Credential material is intentionally not exposed as public
|
||||
attributes and is never included in this provider's representation. A
|
||||
configured context must not outlive this provider; ``DtlsCoapSession``
|
||||
enforces that lifetime by retaining its provider.
|
||||
"""
|
||||
|
||||
__slots__ = ("_callback",)
|
||||
|
||||
def __init__(self, *, identity: bytes, key: bytes) -> None:
|
||||
if type(identity) is not bytes or type(key) is not bytes:
|
||||
raise TypeError("identity and key must be bytes")
|
||||
if len(identity) != 16:
|
||||
raise ValueError("identity must be a raw 16-byte OCF UUID")
|
||||
if b"\x00" in identity:
|
||||
raise ValueError("identity cannot contain a NUL byte")
|
||||
if len(key) not in (16, 32):
|
||||
raise ValueError("key must be 16 or 32 bytes")
|
||||
|
||||
ffi = _util.ffi
|
||||
|
||||
@ffi.callback(_PSK_CLIENT_CALLBACK_CDEF)
|
||||
def client_callback(
|
||||
_ssl,
|
||||
_identity_hint,
|
||||
identity_buffer,
|
||||
max_identity_length,
|
||||
key_buffer,
|
||||
max_key_length,
|
||||
):
|
||||
# OpenSSL callbacks cannot propagate Python exceptions. Fail
|
||||
# before touching either destination when a buffer is unavailable
|
||||
# or too small for the complete credential.
|
||||
if (
|
||||
identity_buffer == ffi.NULL
|
||||
or key_buffer == ffi.NULL
|
||||
or len(identity) + 1 > max_identity_length
|
||||
or len(key) > max_key_length
|
||||
):
|
||||
return 0
|
||||
ffi.memmove(
|
||||
identity_buffer,
|
||||
identity + b"\x00",
|
||||
len(identity) + 1,
|
||||
)
|
||||
ffi.memmove(key_buffer, key, len(key))
|
||||
return len(key)
|
||||
|
||||
object.__setattr__(self, "_callback", client_callback)
|
||||
|
||||
def __setattr__(self, _name: str, _value: object) -> None:
|
||||
raise AttributeError("PskAuth is immutable")
|
||||
|
||||
def __delattr__(self, _name: str) -> None:
|
||||
raise AttributeError("PskAuth is immutable")
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""Return a representation that never includes credential material."""
|
||||
return "PskAuth()"
|
||||
|
||||
def configure_context(self, context: SSL.Context) -> None:
|
||||
"""Configure one context for the narrow Samsung OCF PSK profile."""
|
||||
setter = getattr(_util.lib, "SSL_CTX_set_psk_client_callback", None)
|
||||
if setter is None:
|
||||
raise RuntimeError(
|
||||
"the installed OpenSSL binding does not support DTLS PSK"
|
||||
)
|
||||
context.set_cipher_list(_DTLS_PSK_CIPHERS)
|
||||
setter(context._context, self._callback)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AuthenticationProvider",
|
||||
"CertificateAuth",
|
||||
"PskAuth",
|
||||
"SamsungServerProfile",
|
||||
"SamsungServerRole",
|
||||
"ServerCertificateAuth",
|
||||
]
|
||||
@@ -7,10 +7,13 @@ independently.
|
||||
"""
|
||||
import struct
|
||||
|
||||
from ..errors import MalformedMessageError
|
||||
|
||||
# CoAP option numbers (RFC 7252 + 7641 + 7959)
|
||||
URI_PATH = 11
|
||||
URI_QUERY = 15
|
||||
OBSERVE = 6
|
||||
ETAG = 4
|
||||
CONTENT_FORMAT = 12
|
||||
ACCEPT = 17
|
||||
BLOCK2 = 23
|
||||
@@ -81,7 +84,7 @@ def parse_coap(data):
|
||||
elif d_nib == 14:
|
||||
delta = 269 + int.from_bytes(data[i:i + 2], 'big'); i += 2
|
||||
elif d_nib == 15:
|
||||
raise ValueError("reserved option delta nibble 15")
|
||||
raise MalformedMessageError()
|
||||
else:
|
||||
delta = d_nib
|
||||
if l_nib == 13:
|
||||
@@ -89,7 +92,7 @@ def parse_coap(data):
|
||||
elif l_nib == 14:
|
||||
length = 269 + int.from_bytes(data[i:i + 2], 'big'); i += 2
|
||||
elif l_nib == 15:
|
||||
raise ValueError("reserved option length nibble 15")
|
||||
raise MalformedMessageError()
|
||||
else:
|
||||
length = l_nib
|
||||
num = prev + delta
|
||||
@@ -119,6 +122,14 @@ def block_value(num, more, szx):
|
||||
return struct.pack('>I', v)[1:]
|
||||
|
||||
|
||||
def block_fields(value):
|
||||
"""Decode a CoAP Block-N option value. Inverse of block_value().
|
||||
Returns (num, more, szx). An empty value means block 0, no more,
|
||||
SZX=0 — RFC 7959 §2.2 allows a zero-length option to elide it."""
|
||||
v = int.from_bytes(value, 'big')
|
||||
return v >> 4, (v >> 3) & 1, v & 0x07
|
||||
|
||||
|
||||
def fmt_code(c):
|
||||
"""0x45 → '2.05', 0x84 → '4.04'. Used in log lines."""
|
||||
return f"{c >> 5}.{c & 0x1F:02d}"
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
"""Shared memory-BIO driver for bounded DTLS handshakes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import select
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
from OpenSSL import SSL
|
||||
|
||||
from .coap import split_dtls
|
||||
|
||||
_HANDSHAKE_POLL_S = 0.5
|
||||
_MAX_DATAGRAM_SIZE = 65535
|
||||
|
||||
|
||||
class _HandshakeCancelled(Exception):
|
||||
"""Internal signal that a handshake wake socket became readable."""
|
||||
|
||||
|
||||
def _drive_dtls_handshake(
|
||||
connection,
|
||||
sock,
|
||||
*,
|
||||
deadline: float,
|
||||
retries: int | None = None,
|
||||
on_datagram: Callable[[bytes], None] | None = None,
|
||||
wake_socket=None,
|
||||
) -> bool:
|
||||
"""Drive one memory-BIO DTLS handshake up to a monotonic deadline.
|
||||
|
||||
OpenSSL owns the retransmission schedule. ``retries`` optionally limits
|
||||
how many expired retransmission timers are serviced; a normal session is
|
||||
bounded only by its deadline, while the diagnostic probe retains its
|
||||
explicit retry budget.
|
||||
|
||||
Return ``True`` once OpenSSL reports the handshake complete. The deadline
|
||||
prevents another setup, retry, or network-wait iteration; it does not tear
|
||||
down a session that completed while ``do_handshake()`` was running. TLS and
|
||||
socket failures are left to the caller to classify.
|
||||
"""
|
||||
retransmits = 0
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
connection.do_handshake()
|
||||
return True
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
|
||||
try:
|
||||
output = connection.bio_read(_MAX_DATAGRAM_SIZE)
|
||||
except SSL.WantReadError:
|
||||
output = None
|
||||
if output:
|
||||
for record in split_dtls(output):
|
||||
if sock.send(record) != len(record):
|
||||
raise OSError("incomplete UDP send")
|
||||
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
break
|
||||
wait = min(_HANDSHAKE_POLL_S, remaining)
|
||||
timed_out = False
|
||||
if wake_socket is None:
|
||||
sock.settimeout(wait)
|
||||
try:
|
||||
datagram = sock.recv(_MAX_DATAGRAM_SIZE)
|
||||
except TimeoutError:
|
||||
timed_out = True
|
||||
else:
|
||||
readable, _, _ = select.select(
|
||||
(sock, wake_socket),
|
||||
(),
|
||||
(),
|
||||
wait,
|
||||
)
|
||||
if wake_socket in readable:
|
||||
raise _HandshakeCancelled()
|
||||
if sock not in readable:
|
||||
timed_out = True
|
||||
else:
|
||||
datagram = sock.recv(_MAX_DATAGRAM_SIZE)
|
||||
|
||||
if timed_out:
|
||||
timer = connection.DTLSv1_get_timeout()
|
||||
if timer is not None and timer <= 0:
|
||||
if retries is not None and retransmits >= retries:
|
||||
break
|
||||
connection.DTLSv1_handle_timeout()
|
||||
retransmits += 1
|
||||
continue
|
||||
|
||||
if not datagram:
|
||||
continue
|
||||
if on_datagram is not None:
|
||||
on_datagram(datagram)
|
||||
connection.bio_write(datagram)
|
||||
|
||||
return False
|
||||
@@ -23,18 +23,26 @@ Two problems this solves:
|
||||
how you tell an OCF-PKI-wall device (rejects at cert-verify) from a
|
||||
cipher/version mismatch without a cert it would ever accept.
|
||||
|
||||
Reuses split_dtls() (the record framer) and the same memory-BIO pump as
|
||||
DtlsCoapSession.connect(), so the ClientHello on the wire is byte-for-byte
|
||||
what our real client emits (same cipher list, same @SECLEVEL=0).
|
||||
The production probe generates its frozen ClientHello through the same OpenSSL
|
||||
memory-BIO profile as DtlsCoapSession.connect(), including the exact cipher
|
||||
list, security level, and MTU. The opt-in diagnostic drive retains the full
|
||||
memory-BIO pump for characterizing later server flights.
|
||||
"""
|
||||
|
||||
import concurrent.futures as cf
|
||||
import math
|
||||
import socket
|
||||
import time
|
||||
import warnings
|
||||
from dataclasses import dataclass
|
||||
|
||||
from OpenSSL import SSL
|
||||
|
||||
from ..errors import ProbeError
|
||||
from .auth import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
|
||||
from .coap import split_dtls
|
||||
from .dtls_session import _OCF_ROOT_CA, _load_pem_chain
|
||||
from .dtls_handshake import _drive_dtls_handshake
|
||||
from .endpoint import open_connected_udp_socket
|
||||
|
||||
# DTLS record content types (RFC 6347 §4.1)
|
||||
_CT_CHANGE_CIPHER_SPEC = 20
|
||||
@@ -90,6 +98,349 @@ LIVE = 'live' # DTLS server confirmed (HelloVerifyRequest/ServerHello
|
||||
COMPLETED = 'completed' # full handshake succeeded (cert accepted)
|
||||
REJECTED = 'rejected' # server sent a fatal Alert
|
||||
|
||||
# Aggregate stateless-probe outcomes.
|
||||
SELECTED = 'selected'
|
||||
UNREACHABLE = 'unreachable'
|
||||
AMBIGUOUS = 'ambiguous'
|
||||
|
||||
# First-flight response classes retained by the production liveness API.
|
||||
HELLO_VERIFY_REQUEST = 'hello_verify_request'
|
||||
SERVER_HELLO = 'server_hello'
|
||||
ALERT = 'alert'
|
||||
|
||||
_DTLS_VERSIONS = frozenset((b'\xfe\xff', b'\xfe\xfd'))
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DtlsLivenessResult:
|
||||
"""Bounded, non-sensitive result for one stateless port probe."""
|
||||
|
||||
port: int
|
||||
response_kind: str | None
|
||||
attempts: int
|
||||
rtt_s: float | None = None
|
||||
alert: tuple[int, str] | None = None
|
||||
error_code: str | None = None
|
||||
|
||||
@property
|
||||
def is_dtls_server(self):
|
||||
"""Return whether a structurally valid first-flight reply arrived."""
|
||||
return self.response_kind is not None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DtlsPortProbeResult:
|
||||
"""Selection result for one bounded concurrent probe set."""
|
||||
|
||||
outcome: str
|
||||
selected_port: int | None
|
||||
results: tuple[DtlsLivenessResult, ...]
|
||||
|
||||
@property
|
||||
def live_ports(self):
|
||||
"""Return proven listeners in caller-supplied order."""
|
||||
return tuple(
|
||||
result.port for result in self.results if result.is_dtls_server)
|
||||
|
||||
|
||||
def _validate_liveness_options(port, retries, timeout, mtu):
|
||||
if isinstance(port, bool) or not isinstance(port, int):
|
||||
raise TypeError('port must be an integer')
|
||||
if not 1 <= port <= 65535:
|
||||
raise ValueError('port must be between 1 and 65535')
|
||||
if isinstance(retries, bool) or not isinstance(retries, int):
|
||||
raise TypeError('retries must be an integer')
|
||||
if not 0 <= retries <= 4:
|
||||
raise ValueError('retries must be between zero and four')
|
||||
if isinstance(timeout, bool) or not isinstance(timeout, (int, float)):
|
||||
raise TypeError('timeout must be a number')
|
||||
if not math.isfinite(timeout) or not 0 < timeout <= 30:
|
||||
raise ValueError('timeout must be greater than zero and at most 30')
|
||||
if isinstance(mtu, bool) or not isinstance(mtu, int):
|
||||
raise TypeError('mtu must be an integer')
|
||||
if not 576 <= mtu <= 16384:
|
||||
raise ValueError('mtu is outside the safe UDP range')
|
||||
|
||||
|
||||
def _validate_probe_family(family):
|
||||
if isinstance(family, bool) or not isinstance(family, int):
|
||||
raise TypeError('family must be an address-family integer')
|
||||
if family not in (socket.AF_UNSPEC, socket.AF_INET, socket.AF_INET6):
|
||||
raise ValueError('family must be AF_UNSPEC, AF_INET, or AF_INET6')
|
||||
|
||||
|
||||
def _client_hello_flight(*, mtu):
|
||||
"""Build and freeze the same narrow first flight as a real session."""
|
||||
context = SSL.Context(SSL.DTLS_METHOD)
|
||||
context.load_verify_locations(_OCF_ROOT_CA)
|
||||
context.set_verify(SSL.VERIFY_PEER, lambda *args: True)
|
||||
context.set_cipher_list(_DTLS_CIPHERS)
|
||||
|
||||
connection = SSL.Connection(context, None)
|
||||
connection.set_connect_state()
|
||||
connection.set_ciphertext_mtu(mtu)
|
||||
try:
|
||||
connection.do_handshake()
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
|
||||
records = []
|
||||
while True:
|
||||
try:
|
||||
outbound = connection.bio_read(65535)
|
||||
except SSL.WantReadError:
|
||||
break
|
||||
if not outbound:
|
||||
break
|
||||
records.extend(split_dtls(outbound))
|
||||
if not records:
|
||||
raise ProbeError()
|
||||
return tuple(records)
|
||||
|
||||
|
||||
def _is_complete_hello_verify(body):
|
||||
"""Validate the DTLS version and length-prefixed cookie."""
|
||||
return (
|
||||
len(body) >= 3
|
||||
and body[:2] in _DTLS_VERSIONS
|
||||
and len(body) == 3 + body[2]
|
||||
)
|
||||
|
||||
|
||||
def _is_complete_server_hello(body):
|
||||
"""Validate the fixed fields, session ID, and optional extensions."""
|
||||
if len(body) < 38 or body[:2] not in _DTLS_VERSIONS:
|
||||
return False
|
||||
session_id_length = body[34]
|
||||
if session_id_length > 32:
|
||||
return False
|
||||
fixed_end = 38 + session_id_length
|
||||
if len(body) == fixed_end:
|
||||
return True
|
||||
if len(body) < fixed_end + 2:
|
||||
return False
|
||||
extensions_length = int.from_bytes(body[fixed_end:fixed_end + 2], 'big')
|
||||
return len(body) == fixed_end + 2 + extensions_length
|
||||
|
||||
|
||||
def _parse_liveness_response(datagram):
|
||||
"""Return the kind and validated alert for an epoch-zero first flight."""
|
||||
records = split_dtls(datagram)
|
||||
if not records or sum(map(len, records)) != len(datagram):
|
||||
return None, None
|
||||
|
||||
fallback_kind = None
|
||||
fallback_alert = None
|
||||
for record in records:
|
||||
if len(record) < 13 or record[1:3] not in _DTLS_VERSIONS:
|
||||
continue
|
||||
if record[3:5] != b'\x00\x00':
|
||||
continue
|
||||
fragment = record[13:]
|
||||
if record[0] == _CT_HANDSHAKE:
|
||||
offset = 0
|
||||
while offset + 12 <= len(fragment):
|
||||
header = fragment[offset:offset + 12]
|
||||
message_length = int.from_bytes(header[1:4], 'big')
|
||||
fragment_offset = int.from_bytes(header[6:9], 'big')
|
||||
fragment_length = int.from_bytes(header[9:12], 'big')
|
||||
end = offset + 12 + fragment_length
|
||||
if end > len(fragment):
|
||||
break
|
||||
if fragment_offset == 0 and fragment_length == message_length:
|
||||
body = fragment[offset + 12:end]
|
||||
if header[0] == 3 and _is_complete_hello_verify(body):
|
||||
return HELLO_VERIFY_REQUEST, None
|
||||
if header[0] == 2 and _is_complete_server_hello(body):
|
||||
if fallback_kind is None:
|
||||
fallback_kind = SERVER_HELLO
|
||||
offset = end
|
||||
elif record[0] == _CT_ALERT and len(fragment) == 2:
|
||||
level, description = fragment
|
||||
fallback_kind = ALERT
|
||||
fallback_alert = (
|
||||
level,
|
||||
_ALERT_NAMES.get(description, str(description)),
|
||||
)
|
||||
if level == 2:
|
||||
return fallback_kind, fallback_alert
|
||||
return fallback_kind, fallback_alert
|
||||
|
||||
|
||||
def _classify_liveness_response(datagram):
|
||||
"""Classify a structurally complete epoch-zero DTLS first flight."""
|
||||
return _parse_liveness_response(datagram)[0]
|
||||
|
||||
|
||||
def _probe_dtls_port_with_flight(
|
||||
host, port, *, flight, timeout, retries, family):
|
||||
"""Send one frozen ClientHello flight on a connected UDP socket."""
|
||||
attempt_budget = float(timeout) / (retries + 1)
|
||||
attempts = 0
|
||||
sock = None
|
||||
try:
|
||||
sock, _endpoint = open_connected_udp_socket(
|
||||
host,
|
||||
port,
|
||||
family=family,
|
||||
timeout=attempt_budget,
|
||||
)
|
||||
started = time.monotonic()
|
||||
for attempts in range(1, retries + 2):
|
||||
for record in flight:
|
||||
if sock.send(record) != len(record):
|
||||
raise OSError('short UDP send')
|
||||
attempt_deadline = started + attempts * attempt_budget
|
||||
while True:
|
||||
remaining = attempt_deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
break
|
||||
sock.settimeout(remaining)
|
||||
try:
|
||||
datagram = sock.recv(65535)
|
||||
except TimeoutError:
|
||||
break
|
||||
response_kind, alert = _parse_liveness_response(datagram)
|
||||
if response_kind is None:
|
||||
# A connected UDP socket already rejects other peers. An
|
||||
# unrelated or malformed datagram from the appliance must
|
||||
# not consume a retransmission or count as DTLS proof.
|
||||
continue
|
||||
return DtlsLivenessResult(
|
||||
port=port,
|
||||
response_kind=response_kind,
|
||||
attempts=attempts,
|
||||
rtt_s=time.monotonic() - started,
|
||||
alert=alert,
|
||||
)
|
||||
return DtlsLivenessResult(
|
||||
port=port,
|
||||
response_kind=None,
|
||||
attempts=attempts,
|
||||
error_code='no_dtls_response',
|
||||
)
|
||||
except OSError:
|
||||
return DtlsLivenessResult(
|
||||
port=port,
|
||||
response_kind=None,
|
||||
attempts=attempts,
|
||||
error_code='endpoint_unavailable',
|
||||
)
|
||||
finally:
|
||||
if sock is not None:
|
||||
try:
|
||||
sock.close()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def probe_dtls_port(
|
||||
host, port, *, timeout=3.0, retries=2, mtu=1200,
|
||||
family=socket.AF_UNSPEC):
|
||||
"""Prove one DTLS listener without sending a cookie-bearing flight.
|
||||
|
||||
The ClientHello is generated once. Packet-loss retries resend those exact
|
||||
bytes and no response is ever fed back into OpenSSL, so this function
|
||||
cannot emit a second ClientHello or allocate a server association.
|
||||
|
||||
``timeout`` bounds socket I/O after synchronous platform name resolution;
|
||||
resolver timing remains controlled by the operating system.
|
||||
"""
|
||||
_validate_liveness_options(port, retries, timeout, mtu)
|
||||
_validate_probe_family(family)
|
||||
try:
|
||||
flight = _client_hello_flight(mtu=mtu)
|
||||
except Exception: # noqa: BLE001 - return only a fixed failure code
|
||||
return DtlsLivenessResult(
|
||||
port=port,
|
||||
response_kind=None,
|
||||
attempts=0,
|
||||
error_code='client_hello_unavailable',
|
||||
)
|
||||
return _probe_dtls_port_with_flight(
|
||||
host,
|
||||
port,
|
||||
flight=flight,
|
||||
timeout=timeout,
|
||||
retries=retries,
|
||||
family=family,
|
||||
)
|
||||
|
||||
|
||||
def probe_dtls_ports(
|
||||
host, ports, *, preferred_port=None, timeout=3.0, retries=2,
|
||||
mtu=1200, family=socket.AF_UNSPEC):
|
||||
"""Probe a bounded port set concurrently and select without guessing.
|
||||
|
||||
One proven listener is selected. If multiple listeners answer, a proven
|
||||
``preferred_port`` wins; otherwise the explicit outcome is ``ambiguous``.
|
||||
Results preserve the caller's de-duplicated port order. Each worker's
|
||||
``timeout`` starts after synchronous platform name resolution.
|
||||
"""
|
||||
_validate_probe_family(family)
|
||||
ordered_ports = tuple(dict.fromkeys(ports))
|
||||
if not ordered_ports:
|
||||
return DtlsPortProbeResult(UNREACHABLE, None, ())
|
||||
if len(ordered_ports) > 32:
|
||||
raise ValueError('at most 32 DTLS ports may be probed')
|
||||
for port in ordered_ports:
|
||||
_validate_liveness_options(port, retries, timeout, mtu)
|
||||
if preferred_port is not None:
|
||||
_validate_liveness_options(preferred_port, retries, timeout, mtu)
|
||||
|
||||
try:
|
||||
flight = _client_hello_flight(mtu=mtu)
|
||||
except Exception: # noqa: BLE001 - duplicate one fixed result per port
|
||||
results = tuple(
|
||||
DtlsLivenessResult(
|
||||
port=port,
|
||||
response_kind=None,
|
||||
attempts=0,
|
||||
error_code='client_hello_unavailable',
|
||||
)
|
||||
for port in ordered_ports
|
||||
)
|
||||
return DtlsPortProbeResult(UNREACHABLE, None, results)
|
||||
|
||||
by_port = {}
|
||||
with cf.ThreadPoolExecutor(
|
||||
max_workers=len(ordered_ports),
|
||||
thread_name_prefix='smartthings-dtls-probe') as executor:
|
||||
futures = {
|
||||
executor.submit(
|
||||
_probe_dtls_port_with_flight,
|
||||
host,
|
||||
port,
|
||||
flight=flight,
|
||||
timeout=timeout,
|
||||
retries=retries,
|
||||
family=family,
|
||||
): port
|
||||
for port in ordered_ports
|
||||
}
|
||||
for future in cf.as_completed(futures):
|
||||
port = futures[future]
|
||||
try:
|
||||
by_port[port] = future.result()
|
||||
except Exception: # noqa: BLE001 - isolate one bounded worker
|
||||
by_port[port] = DtlsLivenessResult(
|
||||
port=port,
|
||||
response_kind=None,
|
||||
attempts=0,
|
||||
error_code='probe_worker_failed',
|
||||
)
|
||||
|
||||
results = tuple(by_port[port] for port in ordered_ports)
|
||||
live_ports = tuple(
|
||||
result.port for result in results if result.is_dtls_server)
|
||||
if preferred_port is not None and preferred_port in live_ports:
|
||||
return DtlsPortProbeResult(SELECTED, preferred_port, results)
|
||||
if len(live_ports) == 1:
|
||||
return DtlsPortProbeResult(SELECTED, live_ports[0], results)
|
||||
if live_ports:
|
||||
return DtlsPortProbeResult(AMBIGUOUS, None, results)
|
||||
return DtlsPortProbeResult(UNREACHABLE, None, results)
|
||||
|
||||
|
||||
class ProbeResult:
|
||||
"""What a single ClientHello probe learned about one host:port."""
|
||||
@@ -146,38 +497,80 @@ def classify_datagram(dgram):
|
||||
|
||||
def probe(host, port, *, cert_pem=None, key_pem=None,
|
||||
cert_path=None, key_path=None,
|
||||
stateless=True, retries=2, timeout=3.0, mtu=1280):
|
||||
"""Send a DTLS ClientHello to host:port and classify the server's
|
||||
first flight.
|
||||
stateless=True, retries=2, timeout=3.0, mtu=1200,
|
||||
family=socket.AF_UNSPEC):
|
||||
"""Run the backward-compatible stateless liveness probe.
|
||||
|
||||
Two modes:
|
||||
|
||||
stateless=True (default) — a *liveness* gate. Stop the instant the
|
||||
server proves itself with a HelloVerifyRequest (or ServerHello),
|
||||
and never send the cookie'd second ClientHello. By RFC 6347
|
||||
§4.2.1 the server answers the first ClientHello WITHOUT allocating
|
||||
association state, so a stateless probe leaves the device
|
||||
completely untouched — no orphaned association, no ~8 s §4.2.8
|
||||
cooldown for a later real connect from a different source port.
|
||||
Outcome is DEAD or LIVE. This is the mode a discovery/reconnect
|
||||
loop should use in front of a real handshake.
|
||||
|
||||
stateless=False — a *diagnostic* drive. Continue the handshake as
|
||||
far as the server's own flight goes (up to ServerHelloDone, or to
|
||||
COMPLETED with a client cert), capturing its cipher, cert chain,
|
||||
CertificateRequest, or a fatal Alert. This deliberately commits
|
||||
association state on the device, so keep it out of hot reconnect
|
||||
paths; it is the tool for characterizing an OCF-PKI-wall device
|
||||
(#16) — trust rejection vs cipher/version mismatch.
|
||||
|
||||
A single dropped ClientHello would otherwise read as a false DEAD, so
|
||||
the silent path services OpenSSL's DTLS retransmit timer and re-sends
|
||||
up to `retries` times before giving up. A live server still answers
|
||||
on the first RTT — retransmit only lengthens the silent path.
|
||||
Production callers should prefer :func:`probe_dtls_port`, whose immutable
|
||||
result cannot retain remote datagrams or host names. This adapter preserves
|
||||
the original ``ProbeResult`` shape. ``stateless=False`` remains only as a
|
||||
deprecated compatibility path to the explicitly named stateful diagnostic.
|
||||
|
||||
Never raises on a network/handshake failure — those are folded into
|
||||
the ProbeResult so a discovery loop can race many ports safely.
|
||||
"""
|
||||
if not stateless:
|
||||
warnings.warn(
|
||||
'probe(stateless=False) is deprecated; use '
|
||||
'diagnose_dtls_handshake() explicitly',
|
||||
DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
return diagnose_dtls_handshake(
|
||||
host,
|
||||
port,
|
||||
cert_pem=cert_pem,
|
||||
key_pem=key_pem,
|
||||
cert_path=cert_path,
|
||||
key_path=key_path,
|
||||
retries=retries,
|
||||
timeout=timeout,
|
||||
mtu=mtu,
|
||||
family=family,
|
||||
)
|
||||
|
||||
result = ProbeResult(host, port)
|
||||
liveness = probe_dtls_port(
|
||||
host,
|
||||
port,
|
||||
timeout=timeout,
|
||||
retries=retries,
|
||||
mtu=mtu,
|
||||
family=family,
|
||||
)
|
||||
if liveness.response_kind == HELLO_VERIFY_REQUEST:
|
||||
result.outcome = LIVE
|
||||
result.handshake_msgs.append('HelloVerifyRequest')
|
||||
elif liveness.response_kind == SERVER_HELLO:
|
||||
result.outcome = LIVE
|
||||
result.handshake_msgs.append('ServerHello')
|
||||
elif liveness.response_kind == ALERT:
|
||||
result.alert = liveness.alert
|
||||
result.outcome = (
|
||||
REJECTED
|
||||
if liveness.alert is not None and liveness.alert[0] == 2
|
||||
else LIVE
|
||||
)
|
||||
result.rtt_s = liveness.rtt_s
|
||||
if liveness.error_code not in (None, 'no_dtls_response'):
|
||||
result.error = ProbeError()
|
||||
return result
|
||||
|
||||
|
||||
def diagnose_dtls_handshake(
|
||||
host, port, *, cert_pem=None, key_pem=None,
|
||||
cert_path=None, key_path=None,
|
||||
retries=2, timeout=3.0, mtu=1200,
|
||||
family=socket.AF_UNSPEC):
|
||||
"""Opt in to a stateful DTLS handshake for protocol diagnosis.
|
||||
|
||||
Unlike :func:`probe_dtls_port`, this function feeds the server flight back
|
||||
into OpenSSL. It can therefore emit a cookie-bearing second ClientHello and
|
||||
allocate appliance-side association state. Keep it out of discovery,
|
||||
reconnect, and other production liveness paths.
|
||||
"""
|
||||
_validate_liveness_options(port, retries, timeout, mtu)
|
||||
_validate_probe_family(family)
|
||||
result = ProbeResult(host, port)
|
||||
|
||||
ctx = SSL.Context(SSL.DTLS_METHOD)
|
||||
@@ -185,7 +578,7 @@ def probe(host, port, *, cert_pem=None, key_pem=None,
|
||||
# Accept the chain unconditionally: a probe classifies what the server
|
||||
# sends, it does not gate on our trust decision.
|
||||
ctx.set_verify(SSL.VERIFY_PEER, lambda *a: True)
|
||||
ctx.set_cipher_list(b'ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0')
|
||||
ctx.set_cipher_list(_DTLS_CIPHERS)
|
||||
if cert_pem is not None:
|
||||
_load_pem_chain(ctx, cert_pem, key_pem)
|
||||
elif cert_path is not None:
|
||||
@@ -197,86 +590,56 @@ def probe(host, port, *, cert_pem=None, key_pem=None,
|
||||
conn.set_connect_state()
|
||||
conn.set_ciphertext_mtu(mtu)
|
||||
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
sock.settimeout(0.5)
|
||||
dest = (host, port)
|
||||
|
||||
t0 = time.time()
|
||||
seen = set()
|
||||
retransmits = 0
|
||||
try:
|
||||
while time.time() - t0 < timeout:
|
||||
try:
|
||||
conn.do_handshake()
|
||||
result.outcome = COMPLETED
|
||||
if result.rtt_s is None:
|
||||
result.rtt_s = time.time() - t0
|
||||
break
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
except SSL.Error as e:
|
||||
# A fatal Alert lands here; the alert record was already
|
||||
# captured below, so classification still works.
|
||||
result.error = str(e)
|
||||
break
|
||||
sock, _endpoint = open_connected_udp_socket(
|
||||
host,
|
||||
port,
|
||||
family=family,
|
||||
timeout=min(0.5, timeout),
|
||||
)
|
||||
except OSError:
|
||||
result.error = ProbeError()
|
||||
return result
|
||||
|
||||
try:
|
||||
o = conn.bio_read(65535)
|
||||
if o:
|
||||
for r in split_dtls(o):
|
||||
sock.sendto(r, dest)
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
started = time.monotonic()
|
||||
deadline = started + timeout
|
||||
seen = set()
|
||||
|
||||
try:
|
||||
d, _ = sock.recvfrom(65535)
|
||||
except socket.timeout:
|
||||
# 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
|
||||
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:
|
||||
completed = _drive_dtls_handshake(
|
||||
conn,
|
||||
sock,
|
||||
deadline=deadline,
|
||||
retries=retries,
|
||||
on_datagram=record_datagram,
|
||||
)
|
||||
if completed:
|
||||
result.outcome = COMPLETED
|
||||
if result.rtt_s is None:
|
||||
result.rtt_s = time.time() - t0
|
||||
result.datagrams.append(d)
|
||||
server_flight = False
|
||||
for ct, detail in classify_datagram(d):
|
||||
if ct == _CT_HANDSHAKE:
|
||||
if detail not in seen:
|
||||
seen.add(detail)
|
||||
result.handshake_msgs.append(detail)
|
||||
if result.outcome == DEAD:
|
||||
result.outcome = LIVE
|
||||
if detail in ('HelloVerifyRequest', 'ServerHello'):
|
||||
server_flight = True
|
||||
elif ct == _CT_ALERT and detail is not None:
|
||||
level, name = detail
|
||||
result.alert = (level, name)
|
||||
if level == 2: # fatal
|
||||
result.outcome = REJECTED
|
||||
# Stateless liveness: the server proved itself with a
|
||||
# HelloVerifyRequest/ServerHello, which it answered without
|
||||
# allocating state. Stop before feeding this flight back to
|
||||
# OpenSSL — doing so would make it emit the cookie'd second
|
||||
# ClientHello, the message that actually commits association
|
||||
# state on the device. Not writing it keeps the probe
|
||||
# zero-footprint.
|
||||
if stateless and server_flight:
|
||||
break
|
||||
conn.bio_write(d)
|
||||
result.rtt_s = time.monotonic() - started
|
||||
except SSL.Error:
|
||||
# A fatal Alert lands here; record_datagram() has already classified
|
||||
# the alert record before it is fed back into OpenSSL.
|
||||
result.error = ProbeError()
|
||||
except OSError:
|
||||
result.error = ProbeError()
|
||||
finally:
|
||||
sock.close()
|
||||
|
||||
@@ -288,14 +651,11 @@ def _main(argv):
|
||||
|
||||
if len(argv) < 2:
|
||||
print('usage: python -m smartthings_local.protocol.dtls_probe '
|
||||
'HOST PORT [PORT...] [--cert FILE --key FILE] [--stateless]')
|
||||
'HOST PORT [PORT...] [--diagnostic --cert FILE --key FILE]')
|
||||
return 2
|
||||
host = argv[0]
|
||||
cert_path = key_path = None
|
||||
# CLI defaults to the diagnostic drive so `HOST PORT` characterizes a
|
||||
# device (cipher/cert/Alert). Pass --stateless for the zero-footprint
|
||||
# liveness gate a reconnect loop would use.
|
||||
stateless = False
|
||||
diagnostic = False
|
||||
ports = []
|
||||
it = iter(argv[1:])
|
||||
for a in it:
|
||||
@@ -303,16 +663,31 @@ def _main(argv):
|
||||
cert_path = next(it)
|
||||
elif a == '--key':
|
||||
key_path = next(it)
|
||||
elif a == '--diagnostic':
|
||||
diagnostic = True
|
||||
elif a == '--stateless':
|
||||
stateless = True
|
||||
# Compatibility no-op: stateless is now the fail-safe default.
|
||||
pass
|
||||
else:
|
||||
ports.append(int(a))
|
||||
if not ports:
|
||||
print('at least one PORT is required')
|
||||
return 2
|
||||
ports = list(dict.fromkeys(ports))
|
||||
if len(ports) > 32:
|
||||
print('at most 32 PORT values may be probed')
|
||||
return 2
|
||||
if (cert_path is None) != (key_path is None):
|
||||
print('--cert and --key must be supplied together')
|
||||
return 2
|
||||
if not diagnostic and (cert_path is not None or key_path is not None):
|
||||
print('--cert/--key require the explicit --diagnostic mode')
|
||||
return 2
|
||||
|
||||
# Race the ports: a ClientHello probe is cheap, so fan out and let the
|
||||
# live one answer in ~1 RTT instead of serializing 12 s timeouts.
|
||||
target = diagnose_dtls_handshake if diagnostic else probe
|
||||
with cf.ThreadPoolExecutor(max_workers=max(1, len(ports))) as ex:
|
||||
futs = {ex.submit(probe, host, p, cert_path=cert_path,
|
||||
key_path=key_path, stateless=stateless): p
|
||||
futs = {ex.submit(target, host, p, cert_path=cert_path,
|
||||
key_path=key_path): p
|
||||
for p in ports}
|
||||
results = [f.result() for f in cf.as_completed(futs)]
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,164 @@
|
||||
"""Deterministic UDP endpoint resolution and connected socket setup."""
|
||||
|
||||
import math
|
||||
import socket
|
||||
from dataclasses import dataclass
|
||||
|
||||
from ..errors import EndpointError
|
||||
|
||||
__all__ = [
|
||||
'ResolvedUdpEndpoint',
|
||||
'open_connected_udp_socket',
|
||||
'resolve_udp_endpoint',
|
||||
'resolve_udp_endpoints',
|
||||
]
|
||||
|
||||
_SUPPORTED_FAMILIES = (socket.AF_UNSPEC, socket.AF_INET, socket.AF_INET6)
|
||||
|
||||
|
||||
def _validate_port(port, *, allow_zero=False):
|
||||
lower_bound = 0 if allow_zero else 1
|
||||
if isinstance(port, bool) or not isinstance(port, int):
|
||||
raise TypeError('port must be an integer')
|
||||
if not lower_bound <= port <= 65535:
|
||||
raise ValueError(f'port must be between {lower_bound} and 65535')
|
||||
|
||||
|
||||
def _validate_family(family):
|
||||
if isinstance(family, bool) or not isinstance(family, int):
|
||||
raise TypeError('family must be an address-family integer')
|
||||
if family not in _SUPPORTED_FAMILIES:
|
||||
raise ValueError('family must be AF_UNSPEC, AF_INET, or AF_INET6')
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, repr=False)
|
||||
class ResolvedUdpEndpoint:
|
||||
"""One concrete IPv4 or IPv6 UDP destination.
|
||||
|
||||
``sockaddr`` is the exact tuple returned by ``getaddrinfo`` and therefore
|
||||
retains IPv6 flow and scope IDs. The custom representation omits the
|
||||
address, port, and scope so exception and diagnostic output does not leak
|
||||
a remote endpoint accidentally.
|
||||
"""
|
||||
|
||||
family: int
|
||||
sockaddr: tuple
|
||||
|
||||
def __post_init__(self):
|
||||
if self.family not in (socket.AF_INET, socket.AF_INET6):
|
||||
raise ValueError('resolved family must be AF_INET or AF_INET6')
|
||||
expected_length = 2 if self.family == socket.AF_INET else 4
|
||||
if not isinstance(self.sockaddr, tuple) or \
|
||||
len(self.sockaddr) != expected_length:
|
||||
raise ValueError('sockaddr does not match its address family')
|
||||
_validate_port(self.sockaddr[1])
|
||||
|
||||
@property
|
||||
def host(self):
|
||||
return self.sockaddr[0]
|
||||
|
||||
@property
|
||||
def port(self):
|
||||
return self.sockaddr[1]
|
||||
|
||||
@property
|
||||
def scope_id(self):
|
||||
return self.sockaddr[3] if self.family == socket.AF_INET6 else 0
|
||||
|
||||
@property
|
||||
def family_name(self):
|
||||
return socket.AddressFamily(self.family).name
|
||||
|
||||
def bind_address(self, local_port):
|
||||
"""Return the wildcard bind tuple matching this endpoint's family."""
|
||||
_validate_port(local_port, allow_zero=True)
|
||||
if self.family == socket.AF_INET6:
|
||||
return ('::', local_port, 0, 0)
|
||||
return ('', local_port)
|
||||
|
||||
def __repr__(self):
|
||||
return f'ResolvedUdpEndpoint(family={self.family_name})'
|
||||
|
||||
|
||||
def resolve_udp_endpoints(host, port, *, family=socket.AF_UNSPEC):
|
||||
"""Resolve all unique IPv4/IPv6 UDP candidates in resolver order."""
|
||||
if not isinstance(host, (str, bytes)):
|
||||
raise TypeError('host must be a string or bytes value')
|
||||
if not host:
|
||||
raise ValueError('host must be a non-empty string or bytes value')
|
||||
_validate_port(port)
|
||||
_validate_family(family)
|
||||
|
||||
infos = None
|
||||
try:
|
||||
infos = socket.getaddrinfo(
|
||||
host, port, family, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
|
||||
except (OSError, UnicodeError):
|
||||
pass
|
||||
if infos is None:
|
||||
raise EndpointError() from OSError('UDP endpoint resolution failed')
|
||||
|
||||
endpoints = []
|
||||
seen = set()
|
||||
for resolved_family, socktype, protocol, _canonname, sockaddr in infos:
|
||||
if resolved_family not in (socket.AF_INET, socket.AF_INET6):
|
||||
continue
|
||||
if socktype not in (0, socket.SOCK_DGRAM):
|
||||
continue
|
||||
if protocol not in (0, socket.IPPROTO_UDP):
|
||||
continue
|
||||
expected_length = 2 if resolved_family == socket.AF_INET else 4
|
||||
if not isinstance(sockaddr, tuple) or len(sockaddr) != expected_length:
|
||||
continue
|
||||
key = (resolved_family, sockaddr)
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
endpoints.append(ResolvedUdpEndpoint(resolved_family, sockaddr))
|
||||
|
||||
if not endpoints:
|
||||
raise EndpointError()
|
||||
return tuple(endpoints)
|
||||
|
||||
|
||||
def resolve_udp_endpoint(host, port, *, family=socket.AF_UNSPEC):
|
||||
"""Resolve the first usable UDP candidate."""
|
||||
return resolve_udp_endpoints(host, port, family=family)[0]
|
||||
|
||||
|
||||
def open_connected_udp_socket(
|
||||
host, port, *, family=socket.AF_UNSPEC, local_port=None, timeout=None):
|
||||
"""Create, optionally bind, and connect a UDP socket.
|
||||
|
||||
Candidates are tried in resolver order. A connected UDP socket accepts
|
||||
datagrams only from its exact remote peer and lets the caller use
|
||||
``send``/``recv`` instead of passing an address on every operation.
|
||||
"""
|
||||
if local_port is not None:
|
||||
_validate_port(local_port, allow_zero=True)
|
||||
if timeout is not None:
|
||||
if isinstance(timeout, bool) or not isinstance(timeout, (int, float)):
|
||||
raise TypeError('timeout must be a number or None')
|
||||
if not math.isfinite(timeout) or timeout < 0:
|
||||
raise ValueError('timeout must be a non-negative number or None')
|
||||
|
||||
endpoints = resolve_udp_endpoints(host, port, family=family)
|
||||
for endpoint in endpoints:
|
||||
sock = None
|
||||
try:
|
||||
sock = socket.socket(
|
||||
endpoint.family, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
|
||||
if local_port is not None:
|
||||
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
sock.bind(endpoint.bind_address(local_port))
|
||||
sock.connect(endpoint.sockaddr)
|
||||
sock.settimeout(timeout)
|
||||
return sock, endpoint
|
||||
except OSError:
|
||||
if sock is not None:
|
||||
try:
|
||||
sock.close()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
raise EndpointError() from OSError('UDP socket setup failed')
|
||||
@@ -0,0 +1,128 @@
|
||||
"""Pure IoTivity manufacturer-certificate OwnerPSK derivation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
import hashlib
|
||||
import hmac
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
|
||||
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL: Final = b"x.org.iotivity.conmfgcert"
|
||||
STANDARD_MFG_CERTIFICATE_OXM_LABEL: Final = b"oic.sec.doxm.mfgcert"
|
||||
|
||||
# OpenSSL cipher names mapped to the key-block lengths used by IoTivity's
|
||||
# CAGenerateOwnerPSK implementation.
|
||||
MFG_CERTIFICATE_KEY_BLOCK_LENGTHS: Final[Mapping[str, int]] = MappingProxyType(
|
||||
{
|
||||
"ECDHE-ECDSA-AES128-SHA256": 96,
|
||||
"ECDHE-ECDSA-AES128-CCM": 40,
|
||||
"ECDHE-ECDSA-AES128-CCM8": 40,
|
||||
"ECDHE-ECDSA-AES128-GCM-SHA256": 120,
|
||||
"AES256-SHA256": 128,
|
||||
"ECDHE-ECDSA-AES256-SHA384": 160,
|
||||
"ECDHE-ECDSA-AES256-GCM-SHA384": 184,
|
||||
"AES128-GCM-SHA256": 120,
|
||||
}
|
||||
)
|
||||
|
||||
_TLS_MASTER_SECRET_BYTES: Final = 48
|
||||
_TLS_RANDOM_BYTES: Final = 32
|
||||
_OCF_UUID_BYTES: Final = 16
|
||||
_OWNER_PSK_BYTES: Final = 16
|
||||
|
||||
|
||||
def _require_bytes(name: str, value: bytes, length: int) -> bytes:
|
||||
if not isinstance(value, bytes):
|
||||
raise TypeError(f"{name} must be bytes")
|
||||
if len(value) != length:
|
||||
raise ValueError(f"{name} must be exactly {length} bytes")
|
||||
return value
|
||||
|
||||
|
||||
def _require_uuid(name: str, value: bytes) -> bytes:
|
||||
value = _require_bytes(name, value, _OCF_UUID_BYTES)
|
||||
if not any(value):
|
||||
raise ValueError(f"{name} must not be the nil UUID")
|
||||
return value
|
||||
|
||||
|
||||
def _tls12_p_hash_sha256(
|
||||
key: bytes,
|
||||
label: bytes,
|
||||
random1: bytes,
|
||||
random2: bytes,
|
||||
length: int,
|
||||
) -> bytes:
|
||||
seed = label + random1 + random2
|
||||
a_value = hmac.new(key, seed, hashlib.sha256).digest()
|
||||
output = bytearray()
|
||||
while len(output) < length:
|
||||
output.extend(hmac.new(key, a_value + seed, hashlib.sha256).digest())
|
||||
a_value = hmac.new(key, a_value, hashlib.sha256).digest()
|
||||
return bytes(output[:length])
|
||||
|
||||
|
||||
def derive_mfg_certificate_owner_psk(
|
||||
*,
|
||||
master_secret: bytes,
|
||||
client_random: bytes,
|
||||
server_random: bytes,
|
||||
owner_uuid: bytes,
|
||||
device_uuid: bytes,
|
||||
cipher_name: str,
|
||||
oxm_label: bytes,
|
||||
) -> bytes:
|
||||
"""Derive a 128-bit OwnerPSK from caller-supplied DTLS state.
|
||||
|
||||
This implements IoTivity's two-stage TLS 1.2 SHA-256 P_hash operation.
|
||||
It performs no session access, network I/O, ownership writes, or storage.
|
||||
The caller must supply state from an authenticated manufacturer-certificate
|
||||
session and explicitly select the OXM label used by that transaction.
|
||||
IoTivity's other 96-byte ECDH_ANON, ECDHE_PSK, and ECDHE_RSA mappings are
|
||||
intentionally outside this helper's manufacturer-certificate allowlist.
|
||||
"""
|
||||
|
||||
if not isinstance(cipher_name, str):
|
||||
raise TypeError("cipher_name must be a string")
|
||||
key_block_bytes = MFG_CERTIFICATE_KEY_BLOCK_LENGTHS.get(cipher_name)
|
||||
if key_block_bytes is None:
|
||||
raise ValueError("unexpected manufacturer-certificate DTLS cipher")
|
||||
|
||||
master_secret = _require_bytes(
|
||||
"master_secret", master_secret, _TLS_MASTER_SECRET_BYTES
|
||||
)
|
||||
client_random = _require_bytes(
|
||||
"client_random", client_random, _TLS_RANDOM_BYTES
|
||||
)
|
||||
server_random = _require_bytes(
|
||||
"server_random", server_random, _TLS_RANDOM_BYTES
|
||||
)
|
||||
owner_uuid = _require_uuid("owner_uuid", owner_uuid)
|
||||
device_uuid = _require_uuid("device_uuid", device_uuid)
|
||||
if not isinstance(oxm_label, bytes):
|
||||
raise TypeError("oxm_label must be bytes")
|
||||
if oxm_label not in {
|
||||
CONFIRMED_MFG_CERTIFICATE_OXM_LABEL,
|
||||
STANDARD_MFG_CERTIFICATE_OXM_LABEL,
|
||||
}:
|
||||
raise ValueError("unexpected manufacturer-certificate OXM label")
|
||||
|
||||
key_block = _tls12_p_hash_sha256(
|
||||
master_secret,
|
||||
b"key expansion",
|
||||
server_random,
|
||||
client_random,
|
||||
key_block_bytes,
|
||||
)
|
||||
# IoTivity's OTM callers pass the owner UUID first and the target device
|
||||
# UUID second. The lower adapter's historical rsrc/prov parameter names
|
||||
# describe those arguments inconsistently, so preserve the caller order.
|
||||
return _tls12_p_hash_sha256(
|
||||
key_block,
|
||||
oxm_label,
|
||||
owner_uuid,
|
||||
device_uuid,
|
||||
_OWNER_PSK_BYTES,
|
||||
)
|
||||
@@ -1,9 +1,8 @@
|
||||
"""Port-resolution logic for the MQTT bridge: the stateless pre-flight
|
||||
gate and OCF-band autodiscovery in PushBridge. The DTLS probe is faked so
|
||||
these run without hardware — only the routing/gating/caching is exercised.
|
||||
gate and standard/dynamic OCF port discovery in PushBridge. The DTLS probe is
|
||||
faked so these run without hardware; only routing and selection are exercised.
|
||||
"""
|
||||
import logging
|
||||
import time
|
||||
import types
|
||||
|
||||
import pytest
|
||||
@@ -15,31 +14,52 @@ def _mk_bridge(ocf_port, default=49155, discovered=None):
|
||||
"""A PushBridge shell with only the attributes _resolve_port touches,
|
||||
bypassing the heavyweight __init__ (MQTT client, cert paths, …)."""
|
||||
b = bridge.PushBridge.__new__(bridge.PushBridge)
|
||||
b.app = types.SimpleNamespace(ip='10.0.0.9', ocf_port=ocf_port, index=0)
|
||||
b.app = types.SimpleNamespace(ip='192.0.2.9', ocf_port=ocf_port, index=0)
|
||||
b.descriptor = types.SimpleNamespace(default_observe_port=default)
|
||||
b._discovered_port = discovered
|
||||
b.log = logging.getLogger('test-bridge')
|
||||
return b
|
||||
|
||||
|
||||
def _fake_probe(live_ports):
|
||||
"""Return a probe() stand-in reporting is_dtls_server for live_ports."""
|
||||
def _fake_port_probe(live_ports):
|
||||
"""Return a one-port probe stand-in for the selected live ports."""
|
||||
def fake(ip, port, **kw):
|
||||
alive = port in live_ports
|
||||
return types.SimpleNamespace(
|
||||
port=port, is_dtls_server=alive,
|
||||
outcome='live' if alive else 'dead')
|
||||
port=port,
|
||||
is_dtls_server=alive,
|
||||
)
|
||||
return fake
|
||||
|
||||
|
||||
def _fake_port_set(live_ports):
|
||||
"""Return an aggregate probe stand-in with explicit ambiguity."""
|
||||
def fake(ip, ports, **kw):
|
||||
live = tuple(port for port in ports if port in live_ports)
|
||||
if len(live) == 1:
|
||||
outcome = 'selected'
|
||||
selected_port = live[0]
|
||||
elif live:
|
||||
outcome = 'ambiguous'
|
||||
selected_port = None
|
||||
else:
|
||||
outcome = 'unreachable'
|
||||
selected_port = None
|
||||
return types.SimpleNamespace(
|
||||
outcome=outcome,
|
||||
selected_port=selected_port,
|
||||
)
|
||||
return fake
|
||||
|
||||
|
||||
def test_pinned_live_port_is_gated_and_returned(monkeypatch):
|
||||
monkeypatch.setattr(bridge, 'probe', _fake_probe({49155}))
|
||||
monkeypatch.setattr(bridge, 'probe_dtls_port', _fake_port_probe({49155}))
|
||||
b = _mk_bridge(ocf_port=49155)
|
||||
assert b._resolve_port() == 49155
|
||||
|
||||
|
||||
def test_pinned_dead_port_raises_for_backoff(monkeypatch):
|
||||
monkeypatch.setattr(bridge, 'probe', _fake_probe(set()))
|
||||
monkeypatch.setattr(bridge, 'probe_dtls_port', _fake_port_probe(set()))
|
||||
b = _mk_bridge(ocf_port=49155)
|
||||
with pytest.raises(ConnectionError):
|
||||
b._resolve_port()
|
||||
@@ -48,50 +68,39 @@ def test_pinned_dead_port_raises_for_backoff(monkeypatch):
|
||||
def test_autodiscovery_finds_and_caches_live_port(monkeypatch):
|
||||
# Only 49154 answers; it isn't the descriptor default, so discovery is
|
||||
# what finds it — and it must be cached for the next reconnect.
|
||||
monkeypatch.setattr(bridge, 'probe', _fake_probe({49154}))
|
||||
monkeypatch.setattr(bridge, 'probe_dtls_ports', _fake_port_set({49154}))
|
||||
b = _mk_bridge(ocf_port=None, default=49155)
|
||||
assert b._resolve_port() == 49154
|
||||
assert b._discovered_port == 49154
|
||||
|
||||
|
||||
def test_autodiscovery_returns_a_live_port(monkeypatch):
|
||||
# Early-exit: the first candidate to answer LIVE wins. Real devices
|
||||
# expose exactly one DTLS port; if several answer, any live one is a
|
||||
# correct result.
|
||||
monkeypatch.setattr(bridge, 'probe', _fake_probe({49153, 49155}))
|
||||
def test_autodiscovery_refuses_ambiguous_live_ports(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
bridge,
|
||||
'probe_dtls_ports',
|
||||
_fake_port_set({5684, 49154}),
|
||||
)
|
||||
b = _mk_bridge(ocf_port=None, default=49155)
|
||||
assert b._resolve_port() in {49153, 49155}
|
||||
|
||||
|
||||
def test_autodiscovery_early_exits_before_dead_ports_finish(monkeypatch):
|
||||
# The live port answers immediately; the dead ports "hang" on their
|
||||
# retry budget. Discovery must return at the live port's speed, not
|
||||
# block on the slow dead probes.
|
||||
def slow_probe(ip, port, **kw):
|
||||
if port == 49154:
|
||||
return types.SimpleNamespace(
|
||||
port=port, is_dtls_server=True, outcome='live')
|
||||
time.sleep(0.5) # a dead port burning its retry budget
|
||||
return types.SimpleNamespace(
|
||||
port=port, is_dtls_server=False, outcome='dead')
|
||||
monkeypatch.setattr(bridge, 'probe', slow_probe)
|
||||
b = _mk_bridge(ocf_port=None, default=49155)
|
||||
t0 = time.time()
|
||||
assert b._resolve_port() == 49154
|
||||
assert time.time() - t0 < 0.25 # did not wait out the 0.5s dead probes
|
||||
with pytest.raises(ConnectionError, match='multiple DTLS listeners'):
|
||||
b._resolve_port()
|
||||
|
||||
|
||||
def test_cached_live_port_is_reused_without_rediscovery(monkeypatch):
|
||||
# Cached 49156 and the default 49155 are both live; the cache-first
|
||||
# path must return the cached port, not re-race the band (which would
|
||||
# tie-break to the default).
|
||||
monkeypatch.setattr(bridge, 'probe', _fake_probe({49155, 49156}))
|
||||
# path must return the previously proven port without an ambiguous
|
||||
# full-set probe.
|
||||
monkeypatch.setattr(
|
||||
bridge,
|
||||
'probe_dtls_port',
|
||||
_fake_port_probe({49155, 49156}),
|
||||
)
|
||||
b = _mk_bridge(ocf_port=None, default=49155, discovered=49156)
|
||||
assert b._resolve_port() == 49156
|
||||
|
||||
|
||||
def test_autodiscovery_all_dead_raises_and_clears_cache(monkeypatch):
|
||||
monkeypatch.setattr(bridge, 'probe', _fake_probe(set()))
|
||||
monkeypatch.setattr(bridge, 'probe_dtls_port', _fake_port_probe(set()))
|
||||
monkeypatch.setattr(bridge, 'probe_dtls_ports', _fake_port_set(set()))
|
||||
b = _mk_bridge(ocf_port=None, discovered=49154)
|
||||
with pytest.raises(ConnectionError):
|
||||
b._resolve_port()
|
||||
@@ -102,5 +111,6 @@ def test_candidate_ports_cover_band_plus_default(monkeypatch):
|
||||
b = _mk_bridge(ocf_port=None, default=49200)
|
||||
cands = b._candidate_ports()
|
||||
assert set(bridge.OCF_PORT_BAND) <= set(cands)
|
||||
assert bridge.OCF_STANDARD_SECURE_PORT in cands
|
||||
assert 49200 in cands
|
||||
assert cands == sorted(cands)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+32
-1
@@ -1,5 +1,9 @@
|
||||
import pytest
|
||||
|
||||
from smartthings_local.errors import MalformedMessageError
|
||||
from smartthings_local.protocol.coap import (
|
||||
build_coap, parse_coap, encode_options, block_value, fmt_code,
|
||||
build_coap, parse_coap, encode_options, block_value, block_fields,
|
||||
fmt_code,
|
||||
TYPE_CON, METHOD_GET, URI_PATH, ACCEPT, CF_CBOR, BLOCK2,
|
||||
)
|
||||
|
||||
@@ -40,6 +44,33 @@ def test_block_value_promotes_to_two_bytes_when_num_is_large():
|
||||
assert len(v) == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize('num, more, szx', [
|
||||
(0, 0, 0),
|
||||
(0, 1, 6),
|
||||
(2, 1, 6),
|
||||
(1, 0, 4),
|
||||
(0xFFF, 0, 0),
|
||||
(0xFFFFF, 1, 7),
|
||||
])
|
||||
def test_block_fields_inverts_block_value(num, more, szx):
|
||||
assert block_fields(block_value(num, more, szx)) == (num, more, szx)
|
||||
|
||||
|
||||
def test_block_fields_treats_empty_value_as_block_zero():
|
||||
# RFC 7959 §2.2: a zero-length Block option means num=0, m=0, szx=0.
|
||||
assert block_fields(b'') == (0, 0, 0)
|
||||
|
||||
|
||||
def test_fmt_code_formats_class_dot_detail():
|
||||
assert fmt_code(0x45) == '2.05'
|
||||
assert fmt_code(0x84) == '4.04'
|
||||
|
||||
|
||||
@pytest.mark.parametrize('option_header', (b'\xf0', b'\x0f'))
|
||||
def test_reserved_option_nibbles_raise_classified_value_error(option_header):
|
||||
datagram = b'\x40\x01\x00\x01' + option_header
|
||||
|
||||
with pytest.raises(MalformedMessageError) as exc:
|
||||
parse_coap(datagram)
|
||||
|
||||
assert isinstance(exc.value, ValueError)
|
||||
|
||||
+317
-24
@@ -1,35 +1,60 @@
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from smartthings_local.protocol import dtls_probe as p
|
||||
|
||||
|
||||
def _rec(content_type, frag):
|
||||
def _rec(content_type, frag, *, epoch=0):
|
||||
"""Build one DTLS record: 13-byte header + fragment."""
|
||||
return (bytes([content_type])
|
||||
+ b'\xfe\xfd' # DTLS 1.2
|
||||
+ b'\x00\x00' # epoch
|
||||
+ epoch.to_bytes(2, 'big') # epoch
|
||||
+ b'\x00\x00\x00\x00\x00\x00' # sequence number
|
||||
+ len(frag).to_bytes(2, 'big')
|
||||
+ frag)
|
||||
|
||||
|
||||
def _hs(msg_type, body=b''):
|
||||
return _rec(p._CT_HANDSHAKE, bytes([msg_type]) + body)
|
||||
header = (
|
||||
bytes([msg_type])
|
||||
+ len(body).to_bytes(3, 'big')
|
||||
+ b'\x00\x00' # message sequence
|
||||
+ b'\x00\x00\x00' # fragment offset
|
||||
+ len(body).to_bytes(3, 'big')
|
||||
)
|
||||
return _rec(p._CT_HANDSHAKE, header + body)
|
||||
|
||||
|
||||
def _alert(level, desc):
|
||||
return _rec(p._CT_ALERT, bytes([level, desc]))
|
||||
def _hvr(cookie=b'cookie'):
|
||||
return _hs(3, b'\xfe\xfd' + bytes([len(cookie)]) + cookie)
|
||||
|
||||
|
||||
def _server_hello():
|
||||
body = (
|
||||
b'\xfe\xfd'
|
||||
+ b'\x00' * 32
|
||||
+ b'\x00' # session ID length
|
||||
+ b'\xc0\x2b' # ECDHE-ECDSA-AES128-GCM-SHA256
|
||||
+ b'\x00' # null compression
|
||||
)
|
||||
return _hs(2, body)
|
||||
|
||||
|
||||
def _alert(level, desc, *, epoch=0):
|
||||
return _rec(p._CT_ALERT, bytes([level, desc]), epoch=epoch)
|
||||
|
||||
|
||||
def test_classify_hello_verify_request():
|
||||
assert p.classify_datagram(_hs(3, b'\x00' * 20)) == [
|
||||
assert p.classify_datagram(_hvr()) == [
|
||||
(p._CT_HANDSHAKE, 'HelloVerifyRequest')]
|
||||
|
||||
|
||||
def test_classify_coalesced_server_flight():
|
||||
# OpenSSL commonly hands back ServerHello+Certificate back-to-back.
|
||||
dgram = _hs(2, b'\x00' * 30) + _hs(11, b'\x00' * 40)
|
||||
dgram = _server_hello() + _hs(11, b'\x00' * 40)
|
||||
assert p.classify_datagram(dgram) == [
|
||||
(p._CT_HANDSHAKE, 'ServerHello'),
|
||||
(p._CT_HANDSHAKE, 'Certificate')]
|
||||
@@ -48,7 +73,7 @@ def test_classify_unknown_handshake_type_is_not_lost():
|
||||
def test_dead_port_probe_is_dead_and_never_raises():
|
||||
# Nothing listens here; the probe must fold the silence into a DEAD
|
||||
# result within the timeout rather than raise.
|
||||
r = p.probe('127.0.0.1', 5684, timeout=1.0)
|
||||
r = p.probe('127.0.0.1', 5684, timeout=0.1)
|
||||
assert r.outcome == p.DEAD
|
||||
assert not r.is_dtls_server
|
||||
assert r.datagrams == []
|
||||
@@ -79,6 +104,7 @@ class _FakeSock:
|
||||
self.sends = []
|
||||
self.recv_calls = 0
|
||||
self.closed = False
|
||||
self.destination = None
|
||||
|
||||
def settimeout(self, t):
|
||||
self._timeout = t
|
||||
@@ -89,16 +115,31 @@ class _FakeSock:
|
||||
def bind(self, *a):
|
||||
pass
|
||||
|
||||
def connect(self, destination):
|
||||
self.destination = destination
|
||||
|
||||
def send(self, data):
|
||||
self.sends.append(data)
|
||||
return len(data)
|
||||
|
||||
def sendto(self, data, dest):
|
||||
self.sends.append(data)
|
||||
return len(data)
|
||||
|
||||
def recv(self, n):
|
||||
self.recv_calls += 1
|
||||
resp = self._responder(self)
|
||||
if resp is None:
|
||||
time.sleep(self._timeout)
|
||||
raise TimeoutError()
|
||||
return resp
|
||||
|
||||
def recvfrom(self, n):
|
||||
self.recv_calls += 1
|
||||
resp = self._responder(self)
|
||||
if resp is None:
|
||||
time.sleep(self._timeout)
|
||||
raise socket.timeout()
|
||||
raise TimeoutError()
|
||||
return resp, ('127.0.0.1', 5684)
|
||||
|
||||
def close(self):
|
||||
@@ -113,25 +154,49 @@ def test_stateless_probe_sends_exactly_one_clienthello(monkeypatch):
|
||||
# The §4.2.8 regression guard: a HelloVerifyRequest proves liveness,
|
||||
# and the stateless gate must stop there — never emitting the cookie'd
|
||||
# second ClientHello that would commit association state on the device.
|
||||
fake = _FakeSock(lambda f: _hs(3, b'\x00' * 20))
|
||||
fake = _FakeSock(lambda _fake: _hvr())
|
||||
_patch_sock(monkeypatch, fake)
|
||||
r = p.probe('127.0.0.1', 5684, stateless=True, timeout=2.0)
|
||||
r = p.probe('127.0.0.1', 5684, stateless=True, timeout=0.2)
|
||||
assert r.outcome == p.LIVE
|
||||
assert len(fake.sends) == 1 # only the initial ClientHello
|
||||
assert fake.recv_calls == 1 # stopped on the first flight
|
||||
assert fake.closed
|
||||
|
||||
|
||||
def test_stateless_probe_preserves_first_flight_alert(monkeypatch):
|
||||
fake = _FakeSock(lambda _fake: _alert(2, 48))
|
||||
_patch_sock(monkeypatch, fake)
|
||||
|
||||
result = p.probe('127.0.0.1', 5684, stateless=True, timeout=0.2)
|
||||
|
||||
assert result.outcome == p.REJECTED
|
||||
assert result.alert == (2, 'unknown_ca')
|
||||
assert len(fake.sends) == 1
|
||||
|
||||
|
||||
def test_stateless_warning_alert_proves_liveness_without_fatal_rejection(
|
||||
monkeypatch):
|
||||
fake = _FakeSock(lambda _fake: _alert(1, 90))
|
||||
_patch_sock(monkeypatch, fake)
|
||||
|
||||
result = p.probe('127.0.0.1', 5684, stateless=True, timeout=0.2)
|
||||
|
||||
assert result.outcome == p.LIVE
|
||||
assert result.is_dtls_server
|
||||
assert result.alert == (1, 'user_canceled')
|
||||
|
||||
|
||||
def test_retransmit_recovers_from_dropped_first_flight(monkeypatch):
|
||||
# The first ClientHello is "lost" (recvfrom times out) until OpenSSL's
|
||||
# retransmit timer fires a second flight; only then does the server
|
||||
# answer. A single dropped datagram must NOT read as DEAD.
|
||||
fake = _FakeSock(lambda f: _hs(3, b'\x00' * 20) if len(f.sends) >= 2
|
||||
fake = _FakeSock(lambda f: _hvr() if len(f.sends) >= 2
|
||||
else None)
|
||||
_patch_sock(monkeypatch, fake)
|
||||
r = p.probe('127.0.0.1', 5684, stateless=True, retries=2, timeout=5.0)
|
||||
r = p.probe('127.0.0.1', 5684, stateless=True, retries=2, timeout=0.3)
|
||||
assert r.outcome == p.LIVE
|
||||
assert len(fake.sends) == 2 # initial + one retransmit
|
||||
assert fake.sends[0] == fake.sends[1]
|
||||
|
||||
|
||||
def test_silent_port_is_dead_only_after_flight_budget(monkeypatch):
|
||||
@@ -139,21 +204,249 @@ def test_silent_port_is_dead_only_after_flight_budget(monkeypatch):
|
||||
# `retries` retransmits — not on the first unanswered datagram.
|
||||
fake = _FakeSock(lambda f: None)
|
||||
_patch_sock(monkeypatch, fake)
|
||||
r = p.probe('127.0.0.1', 5684, stateless=True, retries=1, timeout=6.0)
|
||||
r = p.probe('127.0.0.1', 5684, stateless=True, retries=1, timeout=0.2)
|
||||
assert r.outcome == p.DEAD
|
||||
assert not r.is_dtls_server
|
||||
assert len(fake.sends) == 2 # initial + retries(1) retransmit
|
||||
|
||||
|
||||
def test_diagnostic_mode_feeds_server_flight_back(monkeypatch):
|
||||
# The inverse of the stateless guard: stateless=False must NOT stop at
|
||||
# the HelloVerifyRequest — it feeds the flight back into OpenSSL to
|
||||
# drive the handshake onward (the #16 characterization path). The
|
||||
# fed-back record here is a stub, so OpenSSL surfaces an error the
|
||||
# moment it processes it, which is precisely what proves the probe did
|
||||
# not short-circuit before the write.
|
||||
fake = _FakeSock(lambda f: _hs(3, b'\x00' * 20))
|
||||
def test_explicit_diagnostic_feeds_server_flight_back(monkeypatch):
|
||||
# The explicitly named diagnostic must NOT stop at the
|
||||
# HelloVerifyRequest: it feeds the flight back into OpenSSL to drive the
|
||||
# handshake onward (the #16 characterization path). The
|
||||
# fed-back record makes OpenSSL emit a cookie-bearing second ClientHello,
|
||||
# which is precisely what proves the diagnostic did not short-circuit.
|
||||
fake = _FakeSock(
|
||||
lambda f: _hvr() if f.recv_calls == 1 else None)
|
||||
_patch_sock(monkeypatch, fake)
|
||||
r = p.probe('127.0.0.1', 5684, stateless=False, timeout=3.0)
|
||||
r = p.diagnose_dtls_handshake('127.0.0.1', 5684, timeout=0.3)
|
||||
assert r.outcome == p.LIVE # HVR still proved liveness
|
||||
assert r.error is not None # OpenSSL processed the fed-back flight
|
||||
assert len(fake.sends) >= 2 # OpenSSL processed the flight
|
||||
|
||||
|
||||
def test_stateless_probe_ignores_unrelated_datagram_without_retransmit(
|
||||
monkeypatch):
|
||||
responses = iter((
|
||||
_rec(p._CT_APP_DATA, b'unrelated'),
|
||||
_hvr(),
|
||||
))
|
||||
fake = _FakeSock(lambda _fake: next(responses))
|
||||
_patch_sock(monkeypatch, fake)
|
||||
|
||||
result = p.probe_dtls_port(
|
||||
'127.0.0.1', 5684, retries=1, timeout=0.2)
|
||||
|
||||
assert result.response_kind == p.HELLO_VERIFY_REQUEST
|
||||
assert result.attempts == 1
|
||||
assert len(fake.sends) == 1
|
||||
assert fake.recv_calls == 2
|
||||
|
||||
|
||||
def test_stateless_probe_forwards_explicit_address_family(monkeypatch):
|
||||
fake = _FakeSock(lambda _fake: _hvr())
|
||||
calls = []
|
||||
|
||||
def open_socket(host, port, *, family, timeout):
|
||||
calls.append((host, port, family, timeout))
|
||||
fake.settimeout(timeout)
|
||||
return fake, object()
|
||||
|
||||
monkeypatch.setattr(p, 'open_connected_udp_socket', open_socket)
|
||||
|
||||
result = p.probe_dtls_port(
|
||||
'appliance.invalid', 5684, family=socket.AF_INET6, timeout=0.2)
|
||||
|
||||
assert result.is_dtls_server
|
||||
assert calls == [('appliance.invalid', 5684, socket.AF_INET6, 0.2 / 3)]
|
||||
|
||||
|
||||
def test_client_hello_flight_is_complete_epoch_zero_dtls():
|
||||
flight = p._client_hello_flight(mtu=1200)
|
||||
|
||||
assert flight
|
||||
assert all(len(record) <= 1200 for record in flight)
|
||||
assert all(record[1:3] in p._DTLS_VERSIONS for record in flight)
|
||||
assert all(record[3:5] == b'\x00\x00' for record in flight)
|
||||
assert any(
|
||||
record[0] == p._CT_HANDSHAKE and record[13] == 1
|
||||
for record in flight
|
||||
)
|
||||
|
||||
|
||||
def test_liveness_classifier_accepts_first_flight_response_classes():
|
||||
assert p._classify_liveness_response(_hvr()) == \
|
||||
p.HELLO_VERIFY_REQUEST
|
||||
assert p._classify_liveness_response(_server_hello()) == \
|
||||
p.SERVER_HELLO
|
||||
assert p._classify_liveness_response(_alert(2, 48)) == p.ALERT
|
||||
|
||||
|
||||
def test_liveness_classifier_rejects_truncated_or_nonzero_epoch():
|
||||
assert p._classify_liveness_response(_hvr()[:-1]) is None
|
||||
assert p._classify_liveness_response(_hs(3)) is None
|
||||
assert p._classify_liveness_response(_hs(2, b'\x00' * 20)) is None
|
||||
nonzero_epoch = bytearray(_hvr())
|
||||
nonzero_epoch[4] = 1
|
||||
assert p._classify_liveness_response(bytes(nonzero_epoch)) is None
|
||||
|
||||
|
||||
def test_liveness_alert_detail_comes_from_valid_epoch_zero_record(monkeypatch):
|
||||
datagram = _alert(2, 40, epoch=1) + _alert(2, 48)
|
||||
fake = _FakeSock(lambda _fake: datagram)
|
||||
_patch_sock(monkeypatch, fake)
|
||||
|
||||
result = p.probe_dtls_port('127.0.0.1', 5684, timeout=0.2)
|
||||
|
||||
assert result.response_kind == p.ALERT
|
||||
assert result.alert == (2, 'unknown_ca')
|
||||
|
||||
|
||||
def _liveness(port, *, live=True, error_code=None):
|
||||
return p.DtlsLivenessResult(
|
||||
port=port,
|
||||
response_kind=p.HELLO_VERIFY_REQUEST if live else None,
|
||||
attempts=1,
|
||||
error_code=error_code,
|
||||
)
|
||||
|
||||
|
||||
def test_multi_port_probe_runs_concurrently_and_preserves_order(monkeypatch):
|
||||
ports = (5684, 49154, 49155)
|
||||
barrier = threading.Barrier(len(ports))
|
||||
|
||||
def fake_probe(_host, port, **_kwargs):
|
||||
barrier.wait(timeout=2.0)
|
||||
return _liveness(port, live=port == 5684)
|
||||
|
||||
monkeypatch.setattr(p, '_client_hello_flight', lambda **_kwargs: (b'hello',))
|
||||
monkeypatch.setattr(p, '_probe_dtls_port_with_flight', fake_probe)
|
||||
|
||||
result = p.probe_dtls_ports('appliance.invalid', ports)
|
||||
|
||||
assert result.outcome == p.SELECTED
|
||||
assert result.selected_port == 5684
|
||||
assert tuple(item.port for item in result.results) == ports
|
||||
assert not any(
|
||||
thread.name.startswith('smartthings-dtls-probe')
|
||||
for thread in threading.enumerate()
|
||||
)
|
||||
|
||||
|
||||
def test_multi_port_probe_reports_ambiguity_without_guessing(monkeypatch):
|
||||
monkeypatch.setattr(p, '_client_hello_flight', lambda **_kwargs: (b'hello',))
|
||||
monkeypatch.setattr(
|
||||
p,
|
||||
'_probe_dtls_port_with_flight',
|
||||
lambda _host, port, **_kwargs: _liveness(port),
|
||||
)
|
||||
|
||||
result = p.probe_dtls_ports('appliance.invalid', (5684, 49154))
|
||||
|
||||
assert result.outcome == p.AMBIGUOUS
|
||||
assert result.selected_port is None
|
||||
assert result.live_ports == (5684, 49154)
|
||||
|
||||
|
||||
def test_multi_port_probe_prefers_previously_proven_listener(monkeypatch):
|
||||
monkeypatch.setattr(p, '_client_hello_flight', lambda **_kwargs: (b'hello',))
|
||||
monkeypatch.setattr(
|
||||
p,
|
||||
'_probe_dtls_port_with_flight',
|
||||
lambda _host, port, **_kwargs: _liveness(port),
|
||||
)
|
||||
|
||||
result = p.probe_dtls_ports(
|
||||
'appliance.invalid',
|
||||
(5684, 49154),
|
||||
preferred_port=49154,
|
||||
)
|
||||
|
||||
assert result.outcome == p.SELECTED
|
||||
assert result.selected_port == 49154
|
||||
|
||||
|
||||
def test_multi_port_probe_folds_worker_failure_into_redacted_result(monkeypatch):
|
||||
monkeypatch.setattr(p, '_client_hello_flight', lambda **_kwargs: (b'hello',))
|
||||
monkeypatch.setattr(
|
||||
p,
|
||||
'_probe_dtls_port_with_flight',
|
||||
lambda *_args, **_kwargs: (_ for _ in ()).throw(RuntimeError('private')),
|
||||
)
|
||||
|
||||
result = p.probe_dtls_ports('private-host.invalid', (5684,))
|
||||
|
||||
assert result.outcome == p.UNREACHABLE
|
||||
assert result.results[0].error_code == 'probe_worker_failed'
|
||||
assert 'private-host' not in repr(result)
|
||||
assert 'private' not in repr(result)
|
||||
|
||||
|
||||
def test_multi_port_probe_bounds_candidate_count():
|
||||
with pytest.raises(ValueError, match='at most 32'):
|
||||
p.probe_dtls_ports('appliance.invalid', tuple(range(1, 34)))
|
||||
|
||||
|
||||
def test_multi_port_probe_rejects_invalid_family_before_starting_workers():
|
||||
with pytest.raises(ValueError, match='family'):
|
||||
p.probe_dtls_ports(
|
||||
'appliance.invalid',
|
||||
(5684, 49154),
|
||||
family=9999,
|
||||
)
|
||||
|
||||
|
||||
def test_diagnostic_honors_timeout_below_half_second(monkeypatch):
|
||||
now = [10.0]
|
||||
|
||||
class BudgetSocket:
|
||||
def __init__(self):
|
||||
self.timeout = None
|
||||
self.timeouts = []
|
||||
|
||||
def settimeout(self, timeout):
|
||||
self.timeout = timeout
|
||||
self.timeouts.append(timeout)
|
||||
|
||||
def send(self, data):
|
||||
return len(data)
|
||||
|
||||
def recv(self, _size):
|
||||
now[0] += self.timeout
|
||||
raise TimeoutError()
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
sock = BudgetSocket()
|
||||
open_timeouts = []
|
||||
|
||||
def open_socket(_host, _port, *, family, timeout):
|
||||
assert family == socket.AF_UNSPEC
|
||||
open_timeouts.append(timeout)
|
||||
sock.settimeout(timeout)
|
||||
return sock, object()
|
||||
|
||||
monkeypatch.setattr(p, 'open_connected_udp_socket', open_socket)
|
||||
monkeypatch.setattr(p.time, 'monotonic', lambda: now[0])
|
||||
|
||||
result = p.diagnose_dtls_handshake(
|
||||
'appliance.invalid',
|
||||
5684,
|
||||
timeout=0.1,
|
||||
retries=0,
|
||||
)
|
||||
|
||||
assert result.outcome == p.DEAD
|
||||
assert open_timeouts == [0.1]
|
||||
assert sock.timeouts and max(sock.timeouts) <= 0.1
|
||||
assert now[0] <= 10.1
|
||||
|
||||
|
||||
def test_cli_bounds_port_fanout(capsys):
|
||||
result = p._main([
|
||||
'appliance.invalid',
|
||||
*(str(port) for port in range(1, 34)),
|
||||
])
|
||||
|
||||
assert result == 2
|
||||
assert 'at most 32 PORT values' in capsys.readouterr().out
|
||||
|
||||
@@ -1,31 +1,69 @@
|
||||
import gc
|
||||
import traceback
|
||||
import weakref
|
||||
from dataclasses import asdict
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
from cryptography import x509
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from cryptography.x509.oid import NameOID
|
||||
from OpenSSL import SSL, crypto
|
||||
|
||||
from smartthings_local.protocol.auth import CertificateAuth
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession, _load_pem_chain
|
||||
|
||||
|
||||
def _make_self_signed_pem_pair():
|
||||
"""A throwaway self-signed cert + key, just to exercise PEM loading —
|
||||
not meant to resemble a real Samsung client cert."""
|
||||
key = crypto.PKey()
|
||||
key.generate_key(crypto.TYPE_RSA, 2048)
|
||||
def _make_generated_pem_chain():
|
||||
"""Create a throwaway leaf + root chain unrelated to Samsung devices."""
|
||||
now = datetime.now(timezone.utc)
|
||||
root_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
root_name = x509.Name(
|
||||
[x509.NameAttribute(NameOID.COMMON_NAME, "Synthetic test root")]
|
||||
)
|
||||
root_cert = (
|
||||
x509.CertificateBuilder()
|
||||
.subject_name(root_name)
|
||||
.issuer_name(root_name)
|
||||
.public_key(root_key.public_key())
|
||||
.serial_number(1)
|
||||
.not_valid_before(now - timedelta(minutes=1))
|
||||
.not_valid_after(now + timedelta(hours=1))
|
||||
.add_extension(x509.BasicConstraints(ca=True, path_length=None), True)
|
||||
.sign(root_key, hashes.SHA256())
|
||||
)
|
||||
|
||||
cert = crypto.X509()
|
||||
cert.get_subject().CN = "test"
|
||||
cert.set_serial_number(1)
|
||||
cert.gmtime_adj_notBefore(0)
|
||||
cert.gmtime_adj_notAfter(3600)
|
||||
cert.set_issuer(cert.get_subject())
|
||||
cert.set_pubkey(key)
|
||||
cert.sign(key, "sha256")
|
||||
leaf_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
leaf_name = x509.Name(
|
||||
[x509.NameAttribute(NameOID.COMMON_NAME, "Synthetic test client")]
|
||||
)
|
||||
leaf_cert = (
|
||||
x509.CertificateBuilder()
|
||||
.subject_name(leaf_name)
|
||||
.issuer_name(root_name)
|
||||
.public_key(leaf_key.public_key())
|
||||
.serial_number(2)
|
||||
.not_valid_before(now - timedelta(minutes=1))
|
||||
.not_valid_after(now + timedelta(hours=1))
|
||||
.add_extension(x509.BasicConstraints(ca=False, path_length=None), True)
|
||||
.sign(root_key, hashes.SHA256())
|
||||
)
|
||||
|
||||
cert_pem = crypto.dump_certificate(crypto.FILETYPE_PEM, cert).decode()
|
||||
key_pem = crypto.dump_privatekey(crypto.FILETYPE_PEM, key).decode()
|
||||
cert_pem = (
|
||||
leaf_cert.public_bytes(serialization.Encoding.PEM)
|
||||
+ root_cert.public_bytes(serialization.Encoding.PEM)
|
||||
).decode()
|
||||
key_pem = leaf_key.private_bytes(
|
||||
serialization.Encoding.PEM,
|
||||
serialization.PrivateFormat.PKCS8,
|
||||
serialization.NoEncryption(),
|
||||
).decode()
|
||||
return cert_pem, key_pem
|
||||
|
||||
|
||||
def test_load_pem_chain_loads_cert_and_key_in_memory():
|
||||
cert_pem, key_pem = _make_self_signed_pem_pair()
|
||||
cert_pem, key_pem = _make_generated_pem_chain()
|
||||
ctx = SSL.Context(SSL.DTLS_METHOD)
|
||||
_load_pem_chain(ctx, cert_pem, key_pem)
|
||||
ctx.check_privatekey() # raises if cert/key don't match
|
||||
@@ -37,7 +75,7 @@ def test_load_pem_chain_rejects_cert_pem_with_no_certificates():
|
||||
|
||||
|
||||
def test_session_requires_exactly_one_cert_source():
|
||||
cert_pem, key_pem = _make_self_signed_pem_pair()
|
||||
cert_pem, key_pem = _make_generated_pem_chain()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
DtlsCoapSession("host", 1234) # neither pair given
|
||||
@@ -49,10 +87,226 @@ def test_session_requires_exactly_one_cert_source():
|
||||
with pytest.raises(ValueError):
|
||||
DtlsCoapSession("host", 1234, cert_pem=cert_pem) # key_pem missing
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
DtlsCoapSession("host", 1234, cert_path="/a") # key_path missing
|
||||
|
||||
|
||||
def test_session_rejects_provider_with_legacy_certificate_arguments():
|
||||
cert_pem, key_pem = _make_generated_pem_chain()
|
||||
auth = CertificateAuth.from_memory(cert_pem, key_pem)
|
||||
|
||||
with pytest.raises(ValueError, match="auth or legacy certificate"):
|
||||
DtlsCoapSession(
|
||||
"host",
|
||||
1234,
|
||||
cert_pem=cert_pem,
|
||||
key_pem=key_pem,
|
||||
auth=auth,
|
||||
)
|
||||
|
||||
|
||||
def test_session_rejects_object_that_is_not_an_authentication_provider():
|
||||
with pytest.raises(TypeError, match="AuthenticationProvider"):
|
||||
DtlsCoapSession("host", 1234, auth=object())
|
||||
|
||||
|
||||
def test_session_accepts_explicit_certificate_provider():
|
||||
cert_pem, key_pem = _make_generated_pem_chain()
|
||||
auth = CertificateAuth.from_memory(cert_pem, key_pem)
|
||||
session = DtlsCoapSession("host", 1234, auth=auth)
|
||||
|
||||
assert session.auth is auth
|
||||
assert session.cert_path is None
|
||||
assert session.key_path is None
|
||||
assert session.cert_pem is None
|
||||
assert session.key_pem is None
|
||||
|
||||
|
||||
def test_session_retains_authentication_provider_for_its_lifetime():
|
||||
class RetainedProvider:
|
||||
def configure_context(self, _context):
|
||||
return None
|
||||
|
||||
auth = RetainedProvider()
|
||||
reference = weakref.ref(auth)
|
||||
session = DtlsCoapSession("host", 1234, auth=auth)
|
||||
|
||||
del auth
|
||||
gc.collect()
|
||||
assert reference() is session.auth
|
||||
|
||||
|
||||
def test_session_accepts_pem_pair():
|
||||
cert_pem, key_pem = _make_self_signed_pem_pair()
|
||||
cert_pem, key_pem = _make_generated_pem_chain()
|
||||
sess = DtlsCoapSession("host", 1234, cert_pem=cert_pem, key_pem=key_pem)
|
||||
assert sess.cert_path is None
|
||||
assert sess.key_path is None
|
||||
assert sess.cert_pem == cert_pem
|
||||
assert isinstance(sess.auth, CertificateAuth)
|
||||
|
||||
|
||||
def test_session_routes_legacy_file_pair_through_certificate_auth(tmp_path):
|
||||
cert_path = tmp_path / "client.pem"
|
||||
key_path = tmp_path / "client-key.pem"
|
||||
session = DtlsCoapSession(
|
||||
"host",
|
||||
1234,
|
||||
cert_path=cert_path,
|
||||
key_path=key_path,
|
||||
)
|
||||
|
||||
assert isinstance(session.auth, CertificateAuth)
|
||||
assert session.cert_path == str(cert_path)
|
||||
assert session.key_path == str(key_path)
|
||||
|
||||
|
||||
def test_certificate_auth_loads_generated_chain_from_memory_and_files(tmp_path):
|
||||
cert_pem, key_pem = _make_generated_pem_chain()
|
||||
|
||||
memory_context = SSL.Context(SSL.DTLS_METHOD)
|
||||
CertificateAuth.from_memory(cert_pem, key_pem).configure_context(
|
||||
memory_context
|
||||
)
|
||||
memory_context.check_privatekey()
|
||||
|
||||
cert_path = tmp_path / "client.pem"
|
||||
key_path = tmp_path / "client-key.pem"
|
||||
cert_path.write_text(cert_pem)
|
||||
key_path.write_text(key_pem)
|
||||
file_context = SSL.Context(SSL.DTLS_METHOD)
|
||||
CertificateAuth.from_files(cert_path, key_path).configure_context(file_context)
|
||||
file_context.check_privatekey()
|
||||
|
||||
|
||||
def test_certificate_auth_rejects_invalid_memory_material():
|
||||
auth = CertificateAuth.from_memory("not a certificate", "not a key")
|
||||
with pytest.raises(ValueError, match="No certificates found"):
|
||||
auth.configure_context(SSL.Context(SSL.DTLS_METHOD))
|
||||
|
||||
|
||||
def test_invalid_certificate_error_does_not_include_credential_material():
|
||||
marker = "credential" + "-marker"
|
||||
certificate_blob = (
|
||||
"-----BEGIN CERTIFICATE-----\n"
|
||||
f"{marker}\n"
|
||||
"-----END CERTIFICATE-----\n"
|
||||
)
|
||||
key_blob = "invalid-" + marker
|
||||
auth = CertificateAuth.from_memory(certificate_blob, key_blob)
|
||||
|
||||
with pytest.raises(crypto.Error) as captured:
|
||||
auth.configure_context(SSL.Context(SSL.DTLS_METHOD))
|
||||
rendered = (
|
||||
str(captured.value)
|
||||
+ repr(captured.value)
|
||||
+ "".join(traceback.format_exception(captured.value))
|
||||
)
|
||||
assert marker not in rendered
|
||||
|
||||
|
||||
def test_certificate_auth_rejects_invalid_file_material(tmp_path):
|
||||
cert_path = tmp_path / "invalid.pem"
|
||||
key_path = tmp_path / "invalid-key.pem"
|
||||
cert_path.write_text("invalid")
|
||||
key_path.write_text("invalid")
|
||||
|
||||
with pytest.raises(SSL.Error):
|
||||
CertificateAuth.from_files(cert_path, key_path).configure_context(
|
||||
SSL.Context(SSL.DTLS_METHOD)
|
||||
)
|
||||
|
||||
|
||||
def test_certificate_auth_rejects_incomplete_or_mixed_sources():
|
||||
with pytest.raises(ValueError):
|
||||
CertificateAuth()
|
||||
with pytest.raises(ValueError):
|
||||
CertificateAuth(certificate_path="/synthetic/client.pem")
|
||||
with pytest.raises(ValueError):
|
||||
CertificateAuth(certificate_pem="certificate")
|
||||
|
||||
certificate_path = "/synthetic/client.pem"
|
||||
key_path = "/synthetic/client-key.pem"
|
||||
certificate_data = "certificate"
|
||||
key_data = "key"
|
||||
with pytest.raises(ValueError):
|
||||
CertificateAuth(
|
||||
certificate_path=certificate_path,
|
||||
private_key_path=key_path,
|
||||
certificate_pem=certificate_data,
|
||||
private_key_pem=key_data,
|
||||
)
|
||||
|
||||
|
||||
def test_certificate_auth_is_immutable_and_has_secret_safe_repr():
|
||||
cert_pem, key_pem = _make_generated_pem_chain()
|
||||
auth = CertificateAuth.from_memory(cert_pem, key_pem)
|
||||
|
||||
rendered = repr(auth)
|
||||
assert rendered == "CertificateAuth()"
|
||||
assert cert_pem not in rendered
|
||||
assert key_pem not in rendered
|
||||
with pytest.raises(AttributeError, match="immutable"):
|
||||
auth.certificate_pem = None
|
||||
with pytest.raises(AttributeError, match="immutable"):
|
||||
del auth._certificate_pem
|
||||
|
||||
|
||||
def test_certificate_auth_has_no_public_or_dataclass_credential_surface():
|
||||
cert_pem, key_pem = _make_generated_pem_chain()
|
||||
auth = CertificateAuth.from_memory(cert_pem, key_pem)
|
||||
|
||||
assert not hasattr(auth, "certificate_pem")
|
||||
assert not hasattr(auth, "private_key_pem")
|
||||
with pytest.raises(TypeError):
|
||||
vars(auth)
|
||||
with pytest.raises(TypeError):
|
||||
asdict(auth)
|
||||
|
||||
|
||||
def test_certificate_auth_context_setup_matches_legacy_happy_path():
|
||||
class RecordingContext:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
self.verify_callback = None
|
||||
|
||||
def load_verify_locations(self, path):
|
||||
self.calls.append(("load_verify_locations", path))
|
||||
|
||||
def set_verify(self, mode, callback):
|
||||
self.calls.append(("set_verify", mode))
|
||||
self.verify_callback = callback
|
||||
|
||||
def set_cipher_list(self, ciphers):
|
||||
self.calls.append(("set_cipher_list", ciphers))
|
||||
|
||||
def use_certificate_chain_file(self, path):
|
||||
self.calls.append(("use_certificate_chain_file", path))
|
||||
|
||||
def use_privatekey_file(self, path):
|
||||
self.calls.append(("use_privatekey_file", path))
|
||||
|
||||
def check_privatekey(self):
|
||||
self.calls.append(("check_privatekey",))
|
||||
|
||||
context = RecordingContext()
|
||||
CertificateAuth.from_files(
|
||||
"/synthetic/client.pem",
|
||||
"/synthetic/client-key.pem",
|
||||
).configure_context(context)
|
||||
|
||||
assert [call[0] for call in context.calls] == [
|
||||
"load_verify_locations",
|
||||
"set_verify",
|
||||
"set_cipher_list",
|
||||
"use_certificate_chain_file",
|
||||
"use_privatekey_file",
|
||||
"check_privatekey",
|
||||
]
|
||||
assert context.calls[1] == ("set_verify", SSL.VERIFY_PEER)
|
||||
assert context.calls[2] == (
|
||||
"set_cipher_list",
|
||||
b"ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0",
|
||||
)
|
||||
callback = context.verify_callback
|
||||
assert callback(None, None, 0, 0, True) is True
|
||||
assert callback(None, None, 1, 0, False) is False
|
||||
|
||||
@@ -0,0 +1,173 @@
|
||||
"""Reader-thread death visibility (QuiteYellow/SmartThings-Local#37).
|
||||
|
||||
A connected UDP socket surfaces ICMP errors on recv; before this the
|
||||
reader exited silently on the first one and every later request waited
|
||||
out its full timeout against a session nobody was reading. These tests
|
||||
pin the three behaviours that fixed it: advisory ICMP errnos keep the
|
||||
reader alive, a real socket error exits with a WARNING and clears
|
||||
_reader_running, and callers then fail fast with SessionClosedError.
|
||||
"""
|
||||
import errno
|
||||
import logging
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
|
||||
import pytest
|
||||
from OpenSSL import SSL
|
||||
|
||||
from smartthings_local.errors import SessionClosedError
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
||||
|
||||
_LOGGER_NAME = "smartthings_local.protocol.dtls_session"
|
||||
|
||||
|
||||
class _NullAuth:
|
||||
"""Structural AuthenticationProvider — never configured, we skip connect()."""
|
||||
|
||||
def configure_context(self, _context):
|
||||
return None
|
||||
|
||||
|
||||
class _FakeConn:
|
||||
"""Minimal stand-in for SSL.Connection: each datagram written to the
|
||||
BIO surfaces as one decrypted packet on the next recv(), then
|
||||
WantReadError like a drained DTLS record buffer."""
|
||||
|
||||
def __init__(self):
|
||||
self._decrypted = []
|
||||
|
||||
def bio_write(self, datagram):
|
||||
self._decrypted.append(datagram)
|
||||
|
||||
def recv(self, _n):
|
||||
if self._decrypted:
|
||||
return self._decrypted.pop(0)
|
||||
raise SSL.WantReadError()
|
||||
|
||||
def bio_read(self, _n):
|
||||
return b""
|
||||
|
||||
def send(self, _datagram):
|
||||
return None
|
||||
|
||||
def shutdown(self):
|
||||
return None
|
||||
|
||||
|
||||
class _FakeSock:
|
||||
"""Scripted UDP socket. Each step is bytes to return, an exception to
|
||||
raise, or a callable to run (then a timeout, so the loop re-checks
|
||||
_stop). An exhausted script blocks like a real recv timeout; once
|
||||
close()d it raises EBADF the way a closed fd does."""
|
||||
|
||||
def __init__(self, steps=()):
|
||||
self._steps = list(steps)
|
||||
self.closed = False
|
||||
self.timeout = None
|
||||
|
||||
def settimeout(self, value):
|
||||
self.timeout = value
|
||||
|
||||
def recv(self, _n):
|
||||
if self.closed:
|
||||
raise OSError(errno.EBADF, "bad file descriptor")
|
||||
if self._steps:
|
||||
step = self._steps.pop(0)
|
||||
if callable(step):
|
||||
step()
|
||||
raise socket.timeout()
|
||||
if isinstance(step, BaseException):
|
||||
raise step
|
||||
return step
|
||||
time.sleep(0.01)
|
||||
raise socket.timeout()
|
||||
|
||||
def send(self, data):
|
||||
return len(data)
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
def _make_session():
|
||||
sess = DtlsCoapSession("host", 1234, auth=_NullAuth())
|
||||
sess.conn = _FakeConn()
|
||||
return sess
|
||||
|
||||
|
||||
def _run_reader(sess, steps, timeout=2.0):
|
||||
sess.sock = _FakeSock(steps)
|
||||
sess.start_reader()
|
||||
sess._reader_thread.join(timeout)
|
||||
assert not sess._reader_thread.is_alive(), "reader thread did not exit"
|
||||
|
||||
|
||||
def test_advisory_icmp_error_does_not_kill_reader(caplog):
|
||||
sess = _make_session()
|
||||
dispatched = []
|
||||
sess._dispatch_coap = dispatched.append
|
||||
|
||||
with caplog.at_level(logging.DEBUG, logger=_LOGGER_NAME):
|
||||
_run_reader(sess, [
|
||||
OSError(errno.ECONNREFUSED, "connection refused"),
|
||||
b"\x60\x00\x00\x00", # survives, gets dispatched
|
||||
lambda: sess._stop.set(), # end the loop cleanly
|
||||
])
|
||||
|
||||
assert dispatched == [b"\x60\x00\x00\x00"]
|
||||
assert not sess._reader_running.is_set()
|
||||
assert any(r.levelno == logging.DEBUG and "advisory" in r.getMessage()
|
||||
for r in caplog.records)
|
||||
# An advisory errno is not a real exit — no WARNING.
|
||||
assert not any(r.levelno >= logging.WARNING for r in caplog.records)
|
||||
|
||||
|
||||
def test_fatal_socket_error_exits_with_warning(caplog):
|
||||
sess = _make_session()
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger=_LOGGER_NAME):
|
||||
_run_reader(sess, [OSError(errno.EBADF, "bad file descriptor")])
|
||||
|
||||
assert not sess._reader_running.is_set()
|
||||
warnings = [r for r in caplog.records if r.levelno == logging.WARNING]
|
||||
assert len(warnings) == 1
|
||||
assert "reader exiting" in warnings[0].getMessage()
|
||||
|
||||
|
||||
def test_request_fails_fast_after_reader_death():
|
||||
sess = _make_session()
|
||||
_run_reader(sess, [OSError(errno.EBADF, "bad file descriptor")])
|
||||
assert not sess._reader_running.is_set()
|
||||
|
||||
start = time.monotonic()
|
||||
with pytest.raises(SessionClosedError):
|
||||
sess.get(["oic", "d"], timeout=10.0)
|
||||
elapsed = time.monotonic() - start
|
||||
# The whole point: no waiting out the request timeout.
|
||||
assert elapsed < 1.0, f"get() waited {elapsed:.2f}s instead of failing fast"
|
||||
|
||||
|
||||
def test_close_does_not_log_warning_on_teardown(caplog):
|
||||
sess = _make_session()
|
||||
sess.sock = _FakeSock() # empty script: blocks on recv
|
||||
sess.start_reader()
|
||||
time.sleep(0.05) # let the reader reach recv
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger=_LOGGER_NAME):
|
||||
sess.close()
|
||||
sess._reader_thread.join(2.0)
|
||||
|
||||
assert not sess._reader_thread.is_alive()
|
||||
assert not sess._reader_running.is_set()
|
||||
assert not any(r.levelno >= logging.WARNING for r in caplog.records)
|
||||
|
||||
|
||||
def test_check_live_without_reader_matches_old_conn_guard():
|
||||
sess = _make_session() # conn set, reader never started
|
||||
assert sess._reader_thread is None
|
||||
sess._check_live() # must not raise — config-flow behaviour
|
||||
|
||||
sess.conn = None
|
||||
with pytest.raises(SessionClosedError):
|
||||
sess._check_live()
|
||||
@@ -0,0 +1,423 @@
|
||||
import socket
|
||||
import traceback
|
||||
|
||||
import pytest
|
||||
from OpenSSL import SSL
|
||||
|
||||
from smartthings_local.errors import EndpointError
|
||||
from smartthings_local.protocol import dtls_session
|
||||
from smartthings_local.protocol.endpoint import (
|
||||
ResolvedUdpEndpoint,
|
||||
open_connected_udp_socket,
|
||||
resolve_udp_endpoint,
|
||||
resolve_udp_endpoints,
|
||||
)
|
||||
|
||||
|
||||
def _addrinfo(family, sockaddr, *, socktype=socket.SOCK_DGRAM,
|
||||
protocol=socket.IPPROTO_UDP):
|
||||
return family, socktype, protocol, '', sockaddr
|
||||
|
||||
|
||||
class FakeSocket:
|
||||
def __init__(self, family, *, fail_bind=False, fail_connect=False):
|
||||
self.family = family
|
||||
self.fail_bind = fail_bind
|
||||
self.fail_connect = fail_connect
|
||||
self.options = []
|
||||
self.bound = None
|
||||
self.peer = None
|
||||
self.timeout = None
|
||||
self.closed = False
|
||||
self.sent = []
|
||||
self.inbound = []
|
||||
|
||||
def setsockopt(self, *args):
|
||||
self.options.append(args)
|
||||
|
||||
def bind(self, address):
|
||||
if self.fail_bind:
|
||||
raise OSError('synthetic bind failure')
|
||||
self.bound = address
|
||||
|
||||
def connect(self, address):
|
||||
if self.fail_connect:
|
||||
raise OSError('synthetic connect failure')
|
||||
self.peer = address
|
||||
|
||||
def settimeout(self, timeout):
|
||||
self.timeout = timeout
|
||||
|
||||
def send(self, data):
|
||||
self.sent.append(data)
|
||||
return len(data)
|
||||
|
||||
def recv(self, _size):
|
||||
return self.inbound.pop(0)
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
def _socket_factory(monkeypatch, failures=()):
|
||||
created = []
|
||||
remaining = list(failures)
|
||||
|
||||
def factory(family, _socktype, _protocol):
|
||||
failure = remaining.pop(0) if remaining else None
|
||||
sock = FakeSocket(
|
||||
family,
|
||||
fail_bind=failure == 'bind',
|
||||
fail_connect=failure == 'connect',
|
||||
)
|
||||
created.append(sock)
|
||||
return sock
|
||||
|
||||
monkeypatch.setattr(
|
||||
'smartthings_local.protocol.endpoint.socket.socket', factory)
|
||||
return created
|
||||
|
||||
|
||||
def test_resolve_ipv4_endpoint_and_redacted_repr(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
|
||||
lambda *args: [_addrinfo(socket.AF_INET, ('192.0.2.10', 5684))],
|
||||
)
|
||||
|
||||
endpoint = resolve_udp_endpoint('device.example', 5684)
|
||||
|
||||
assert endpoint.family == socket.AF_INET
|
||||
assert endpoint.host == '192.0.2.10'
|
||||
assert endpoint.port == 5684
|
||||
assert endpoint.scope_id == 0
|
||||
assert repr(endpoint) == 'ResolvedUdpEndpoint(family=AF_INET)'
|
||||
assert '192.0.2.10' not in repr(endpoint)
|
||||
assert '5684' not in repr(endpoint)
|
||||
|
||||
|
||||
def test_resolve_scoped_ipv6_preserves_flow_and_scope(monkeypatch):
|
||||
sockaddr = ('2001:db8::1', 5684, 3, 7)
|
||||
calls = []
|
||||
|
||||
def resolve(*args):
|
||||
calls.append(args)
|
||||
return [_addrinfo(socket.AF_INET6, sockaddr)]
|
||||
|
||||
monkeypatch.setattr(
|
||||
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
|
||||
resolve,
|
||||
)
|
||||
|
||||
endpoint = resolve_udp_endpoint(
|
||||
'device.example', 5684, family=socket.AF_INET6)
|
||||
|
||||
assert endpoint.sockaddr == sockaddr
|
||||
assert endpoint.scope_id == 7
|
||||
assert endpoint.bind_address(55000) == ('::', 55000, 0, 0)
|
||||
assert '2001:db8::1' not in repr(endpoint)
|
||||
assert '7' not in repr(endpoint)
|
||||
assert calls == [(
|
||||
'device.example', 5684, socket.AF_INET6,
|
||||
socket.SOCK_DGRAM, socket.IPPROTO_UDP)]
|
||||
|
||||
|
||||
def test_resolver_order_is_stable_and_duplicates_are_removed(monkeypatch):
|
||||
first = _addrinfo(socket.AF_INET6, ('2001:db8::10', 5684, 0, 0))
|
||||
second = _addrinfo(socket.AF_INET, ('198.51.100.20', 5684))
|
||||
ignored = _addrinfo(
|
||||
socket.AF_INET, ('203.0.113.30', 5684), socktype=socket.SOCK_STREAM)
|
||||
monkeypatch.setattr(
|
||||
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
|
||||
lambda *args: [first, first, ignored, second],
|
||||
)
|
||||
|
||||
endpoints = resolve_udp_endpoints('device.example', 5684)
|
||||
|
||||
assert [endpoint.sockaddr for endpoint in endpoints] == [
|
||||
first[-1], second[-1]]
|
||||
|
||||
|
||||
def test_resolver_failure_raises_redacted_endpoint_error(monkeypatch):
|
||||
def fail(*args):
|
||||
raise socket.gaierror('credential-value at device.example')
|
||||
|
||||
monkeypatch.setattr(
|
||||
'smartthings_local.protocol.endpoint.socket.getaddrinfo', fail)
|
||||
remote_host = 'device.example'
|
||||
|
||||
with pytest.raises(EndpointError) as exc:
|
||||
resolve_udp_endpoint(remote_host, 5684)
|
||||
|
||||
formatted = ''.join(traceback.format_exception(exc.value))
|
||||
assert isinstance(exc.value, OSError)
|
||||
assert exc.value.__context__ is None
|
||||
assert 'UDP endpoint resolution failed' in formatted
|
||||
assert 'credential-value' not in formatted
|
||||
assert 'device.example' not in formatted
|
||||
|
||||
|
||||
def test_socket_setup_tries_next_candidate_after_bind_failure(monkeypatch):
|
||||
candidates = [
|
||||
_addrinfo(socket.AF_INET6, ('2001:db8::10', 5684, 0, 0)),
|
||||
_addrinfo(socket.AF_INET, ('192.0.2.10', 5684)),
|
||||
]
|
||||
monkeypatch.setattr(
|
||||
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
|
||||
lambda *args: candidates,
|
||||
)
|
||||
created = _socket_factory(monkeypatch, failures=('bind', None))
|
||||
|
||||
sock, endpoint = open_connected_udp_socket(
|
||||
'device.example', 5684, local_port=55000, timeout=1.5)
|
||||
|
||||
assert created[0].closed
|
||||
assert sock is created[1]
|
||||
assert endpoint.family == socket.AF_INET
|
||||
assert sock.bound == ('', 55000)
|
||||
assert sock.peer == ('192.0.2.10', 5684)
|
||||
assert sock.timeout == 1.5
|
||||
assert (socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) in sock.options
|
||||
|
||||
|
||||
def test_socket_setup_failure_is_redacted_and_closes_candidates(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
|
||||
lambda *args: [
|
||||
_addrinfo(socket.AF_INET, ('192.0.2.10', 5684)),
|
||||
_addrinfo(socket.AF_INET, ('198.51.100.20', 5684)),
|
||||
],
|
||||
)
|
||||
created = _socket_factory(monkeypatch, failures=('connect', 'connect'))
|
||||
|
||||
with pytest.raises(EndpointError) as exc:
|
||||
open_connected_udp_socket('device.example', 5684)
|
||||
|
||||
formatted = ''.join(traceback.format_exception(exc.value))
|
||||
assert all(sock.closed for sock in created)
|
||||
assert exc.value.__context__ is None
|
||||
assert 'UDP socket setup failed' in formatted
|
||||
assert 'synthetic connect failure' not in formatted
|
||||
assert '192.0.2.10' not in formatted
|
||||
assert '198.51.100.20' not in formatted
|
||||
|
||||
|
||||
def test_same_port_different_hosts_remain_distinct_with_source_reuse(
|
||||
monkeypatch):
|
||||
addresses = {
|
||||
'first.example': '192.0.2.10',
|
||||
'second.example': '198.51.100.20',
|
||||
}
|
||||
|
||||
def resolve(host, port, *_args):
|
||||
return [_addrinfo(socket.AF_INET, (addresses[host], port))]
|
||||
|
||||
monkeypatch.setattr(
|
||||
'smartthings_local.protocol.endpoint.socket.getaddrinfo', resolve)
|
||||
created = _socket_factory(monkeypatch)
|
||||
|
||||
first, _ = open_connected_udp_socket(
|
||||
'first.example', 5684, local_port=55000)
|
||||
second, _ = open_connected_udp_socket(
|
||||
'second.example', 5684, local_port=55000)
|
||||
|
||||
assert first.peer == ('192.0.2.10', 5684)
|
||||
assert second.peer == ('198.51.100.20', 5684)
|
||||
assert first.peer != second.peer
|
||||
assert [sock.bound for sock in created] == [('', 55000), ('', 55000)]
|
||||
|
||||
|
||||
def test_connected_udp_socket_filters_datagrams_from_another_peer():
|
||||
expected_peer = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
other_peer = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
client = None
|
||||
try:
|
||||
expected_peer.bind(('127.0.0.1', 0))
|
||||
other_peer.bind(('127.0.0.1', 0))
|
||||
expected_peer.settimeout(0.5)
|
||||
client, endpoint = open_connected_udp_socket(
|
||||
'127.0.0.1',
|
||||
expected_peer.getsockname()[1],
|
||||
family=socket.AF_INET,
|
||||
timeout=0.5,
|
||||
)
|
||||
|
||||
assert endpoint.sockaddr == expected_peer.getsockname()
|
||||
assert client.send(b'client datagram') == len(b'client datagram')
|
||||
payload, source = expected_peer.recvfrom(64)
|
||||
assert payload == b'client datagram'
|
||||
assert source == client.getsockname()
|
||||
|
||||
other_peer.sendto(b'unrelated', client.getsockname())
|
||||
expected_peer.sendto(b'expected', client.getsockname())
|
||||
assert client.recv(64) == b'expected'
|
||||
finally:
|
||||
if client is not None:
|
||||
client.close()
|
||||
expected_peer.close()
|
||||
other_peer.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
('host', 'port', 'family'),
|
||||
(
|
||||
('', 5684, socket.AF_UNSPEC),
|
||||
('device.example', 0, socket.AF_UNSPEC),
|
||||
('device.example', 65536, socket.AF_UNSPEC),
|
||||
('device.example', 5684, 9999),
|
||||
),
|
||||
)
|
||||
def test_invalid_endpoint_inputs_fail_before_resolution(host, port, family):
|
||||
with pytest.raises(ValueError):
|
||||
resolve_udp_endpoint(host, port, family=family)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
('host', 'port', 'family'),
|
||||
(
|
||||
(None, 5684, socket.AF_UNSPEC),
|
||||
('device.example', '5684', socket.AF_UNSPEC),
|
||||
('device.example', 5684, 'AF_INET'),
|
||||
),
|
||||
)
|
||||
def test_endpoint_input_types_are_explicit(host, port, family):
|
||||
with pytest.raises(TypeError):
|
||||
resolve_udp_endpoint(host, port, family=family)
|
||||
|
||||
|
||||
@pytest.mark.parametrize('timeout', (-1, float('nan'), float('inf')))
|
||||
def test_socket_timeout_must_be_finite_and_non_negative(timeout):
|
||||
with pytest.raises(ValueError):
|
||||
open_connected_udp_socket('device.example', 5684, timeout=timeout)
|
||||
|
||||
|
||||
def test_session_uses_connected_socket_send_and_recv(monkeypatch):
|
||||
endpoint = ResolvedUdpEndpoint(
|
||||
socket.AF_INET6, ('2001:db8::10', 5684, 0, 0))
|
||||
sock = FakeSocket(socket.AF_INET6)
|
||||
sock.peer = endpoint.sockaddr
|
||||
sock.inbound.append(b'synthetic server flight')
|
||||
outbound_record = (
|
||||
b'\x16\xfe\xfd' + b'\x00' * 8 + b'\x00\x01' + b'x')
|
||||
|
||||
class FakeContext:
|
||||
def load_verify_locations(self, *args):
|
||||
pass
|
||||
|
||||
def set_verify(self, *args):
|
||||
pass
|
||||
|
||||
def set_cipher_list(self, *args):
|
||||
pass
|
||||
|
||||
def use_certificate_chain_file(self, *args):
|
||||
pass
|
||||
|
||||
def use_privatekey_file(self, *args):
|
||||
pass
|
||||
|
||||
def check_privatekey(self):
|
||||
pass
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(self):
|
||||
self.handshake_calls = 0
|
||||
self.bio_reads = 0
|
||||
self.bio_writes = []
|
||||
|
||||
def set_connect_state(self):
|
||||
pass
|
||||
|
||||
def set_ciphertext_mtu(self, *args):
|
||||
pass
|
||||
|
||||
def do_handshake(self):
|
||||
self.handshake_calls += 1
|
||||
if self.handshake_calls == 1:
|
||||
raise SSL.WantReadError()
|
||||
|
||||
def bio_read(self, _size):
|
||||
self.bio_reads += 1
|
||||
if self.bio_reads == 1:
|
||||
return outbound_record
|
||||
raise SSL.WantReadError()
|
||||
|
||||
def bio_write(self, data):
|
||||
self.bio_writes.append(data)
|
||||
|
||||
connection = FakeConnection()
|
||||
open_calls = []
|
||||
|
||||
def open_socket(*args, **kwargs):
|
||||
open_calls.append((args, kwargs))
|
||||
return sock, endpoint
|
||||
|
||||
monkeypatch.setattr(dtls_session.SSL, 'Context', lambda *args: FakeContext())
|
||||
monkeypatch.setattr(
|
||||
dtls_session.SSL, 'Connection', lambda *args: connection)
|
||||
monkeypatch.setattr(
|
||||
dtls_session,
|
||||
'open_connected_udp_socket',
|
||||
open_socket,
|
||||
)
|
||||
monkeypatch.setattr(dtls_session.time, 'sleep', lambda _delay: None)
|
||||
|
||||
session = dtls_session.DtlsCoapSession(
|
||||
'device.example', 5684,
|
||||
cert_path='/synthetic/client.pem',
|
||||
key_path='/synthetic/client.key',
|
||||
family=socket.AF_INET6,
|
||||
)
|
||||
session.connect()
|
||||
|
||||
assert sock.sent == [outbound_record]
|
||||
assert connection.bio_writes == [b'synthetic server flight']
|
||||
assert session.endpoint is endpoint
|
||||
assert session.dest == endpoint.sockaddr
|
||||
assert open_calls == [(('device.example', 5684), {
|
||||
'family': socket.AF_INET6,
|
||||
'local_port': None,
|
||||
'timeout': 0.5,
|
||||
})]
|
||||
|
||||
session.close()
|
||||
assert session.endpoint is None
|
||||
assert session.dest is None
|
||||
|
||||
|
||||
def test_session_send_failure_has_no_raw_exception_context():
|
||||
outbound_record = (
|
||||
b'\x16\xfe\xfd' + b'\x00' * 8 + b'\x00\x01' + b'x')
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(self):
|
||||
self.bio_reads = 0
|
||||
|
||||
def send(self, _data):
|
||||
pass
|
||||
|
||||
def bio_read(self, _size):
|
||||
self.bio_reads += 1
|
||||
if self.bio_reads == 1:
|
||||
return outbound_record
|
||||
raise SSL.WantReadError()
|
||||
|
||||
class FailingSocket:
|
||||
def send(self, _data):
|
||||
raise OSError('credential-value at device.example')
|
||||
|
||||
session = dtls_session.DtlsCoapSession(
|
||||
'device.example', 5684,
|
||||
cert_path='/synthetic/client.pem',
|
||||
key_path='/synthetic/client.key',
|
||||
)
|
||||
session.conn = FakeConnection()
|
||||
session.sock = FailingSocket()
|
||||
|
||||
with pytest.raises(EndpointError) as exc:
|
||||
session._send_dgram(b'payload')
|
||||
|
||||
formatted = ''.join(traceback.format_exception(exc.value))
|
||||
assert exc.value.__context__ is None
|
||||
assert 'UDP send failed' in formatted
|
||||
assert 'credential-value' not in formatted
|
||||
assert 'device.example' not in formatted
|
||||
@@ -0,0 +1,172 @@
|
||||
import traceback
|
||||
|
||||
import pytest
|
||||
from OpenSSL import SSL
|
||||
|
||||
from smartthings_local.errors import (
|
||||
AuthenticationError,
|
||||
AuthorizationError,
|
||||
BlockwiseError,
|
||||
EndpointError,
|
||||
MalformedMessageError,
|
||||
ObserveError,
|
||||
ProbeError,
|
||||
SessionClosedError,
|
||||
SessionError,
|
||||
SessionTimeoutError,
|
||||
SmartThingsLocalError,
|
||||
)
|
||||
from smartthings_local.protocol import dtls_session
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
||||
from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint
|
||||
|
||||
ERROR_TYPES = (
|
||||
EndpointError,
|
||||
ProbeError,
|
||||
SessionError,
|
||||
AuthenticationError,
|
||||
AuthorizationError,
|
||||
SessionTimeoutError,
|
||||
SessionClosedError,
|
||||
MalformedMessageError,
|
||||
BlockwiseError,
|
||||
ObserveError,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
('error_type', 'legacy_type'),
|
||||
(
|
||||
(EndpointError, OSError),
|
||||
(ProbeError, ConnectionError),
|
||||
(SessionError, ConnectionError),
|
||||
(AuthenticationError, ConnectionError),
|
||||
(AuthorizationError, PermissionError),
|
||||
(SessionTimeoutError, TimeoutError),
|
||||
(SessionClosedError, ConnectionError),
|
||||
(MalformedMessageError, ValueError),
|
||||
(BlockwiseError, ConnectionError),
|
||||
(ObserveError, ConnectionError),
|
||||
),
|
||||
)
|
||||
def test_errors_preserve_legacy_builtin_catches(error_type, legacy_type):
|
||||
error = error_type()
|
||||
assert isinstance(error, SmartThingsLocalError)
|
||||
assert isinstance(error, legacy_type)
|
||||
|
||||
|
||||
@pytest.mark.parametrize('error_type', ERROR_TYPES)
|
||||
def test_error_text_is_fixed_and_redacted(error_type):
|
||||
error = error_type()
|
||||
text = f'{error!s} {error!r}'
|
||||
assert error.code
|
||||
assert str(error) == error.message
|
||||
assert repr(error) == f'{error_type.__name__}(code={error.code!r})'
|
||||
assert 'device.example' not in text
|
||||
assert '/synthetic/client.key' not in text
|
||||
assert 'credential-value' not in text
|
||||
with pytest.raises(TypeError):
|
||||
error_type('credential-value')
|
||||
|
||||
|
||||
def test_chained_backend_error_is_not_copied_into_public_text():
|
||||
backend = ConnectionError('DTLS backend failed')
|
||||
|
||||
try:
|
||||
raise SessionError() from backend
|
||||
except SessionError as error:
|
||||
assert error.__cause__ is backend
|
||||
formatted = ''.join(traceback.format_exception(error))
|
||||
assert 'DTLS backend failed' in formatted
|
||||
assert 'device.example' not in formatted
|
||||
assert 'credential-value' not in formatted
|
||||
|
||||
|
||||
def test_handshake_error_is_classified_without_backend_text(monkeypatch):
|
||||
class FakeContext:
|
||||
def load_verify_locations(self, *args):
|
||||
pass
|
||||
|
||||
def set_verify(self, *args):
|
||||
pass
|
||||
|
||||
def set_cipher_list(self, *args):
|
||||
pass
|
||||
|
||||
def use_certificate_chain_file(self, *args):
|
||||
pass
|
||||
|
||||
def use_privatekey_file(self, *args):
|
||||
pass
|
||||
|
||||
def check_privatekey(self):
|
||||
pass
|
||||
|
||||
class FakeConnection:
|
||||
def set_connect_state(self):
|
||||
pass
|
||||
|
||||
def set_ciphertext_mtu(self, *args):
|
||||
pass
|
||||
|
||||
def do_handshake(self):
|
||||
raise SSL.Error('credential-value at device.example')
|
||||
|
||||
class FakeSocket:
|
||||
closed = False
|
||||
|
||||
def settimeout(self, *args):
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
fake_socket = FakeSocket()
|
||||
monkeypatch.setattr(dtls_session.SSL, 'Context', lambda *args: FakeContext())
|
||||
monkeypatch.setattr(
|
||||
dtls_session.SSL, 'Connection', lambda *args: FakeConnection())
|
||||
endpoint = ResolvedUdpEndpoint(
|
||||
dtls_session.socket.AF_INET, ('192.0.2.10', 5684))
|
||||
monkeypatch.setattr(
|
||||
dtls_session,
|
||||
'open_connected_udp_socket',
|
||||
lambda *args, **kwargs: (fake_socket, endpoint),
|
||||
)
|
||||
|
||||
session = DtlsCoapSession(
|
||||
'device.example', 5684,
|
||||
cert_path='/synthetic/client.pem',
|
||||
key_path='/synthetic/client.key',
|
||||
)
|
||||
with pytest.raises(SessionError) as exc:
|
||||
session.connect()
|
||||
|
||||
assert fake_socket.closed
|
||||
assert isinstance(exc.value, ConnectionError)
|
||||
assert isinstance(exc.value.__cause__, ConnectionError)
|
||||
assert exc.value.__context__ is None
|
||||
formatted = ''.join(traceback.format_exception(exc.value))
|
||||
assert 'DTLS backend failed' in formatted
|
||||
assert 'device.example' not in formatted
|
||||
assert 'credential-value' not in formatted
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
'operation',
|
||||
(
|
||||
lambda session: session.get(['resource']),
|
||||
lambda session: session.post(['resource'], b'payload'),
|
||||
lambda session: session.ping(),
|
||||
lambda session: session.refresh_observes([]),
|
||||
lambda session: session.subscribe(['resource']),
|
||||
),
|
||||
)
|
||||
def test_closed_session_operations_raise_classified_error(operation):
|
||||
session = DtlsCoapSession(
|
||||
'device.example', 5684,
|
||||
cert_path='/synthetic/client.pem',
|
||||
key_path='/synthetic/client.key',
|
||||
)
|
||||
|
||||
with pytest.raises(SessionClosedError):
|
||||
operation(session)
|
||||
@@ -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,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)
|
||||
@@ -0,0 +1,380 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import traceback
|
||||
import weakref
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import asdict
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from OpenSSL import SSL
|
||||
|
||||
from smartthings_local.errors import SessionError
|
||||
from smartthings_local.protocol import auth as auth_module
|
||||
from smartthings_local.protocol import dtls_session as session_module
|
||||
from smartthings_local.protocol.auth import PskAuth
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
||||
|
||||
_IDENTITY = b"i" * 16
|
||||
_KEY = b"k" * 16
|
||||
_OTHER_IDENTITY = b"j" * 16
|
||||
_OTHER_KEY = b"l" * 32
|
||||
|
||||
|
||||
class _BytesSubclass(bytes):
|
||||
pass
|
||||
|
||||
|
||||
def _fake_openssl_util(setter):
|
||||
return SimpleNamespace(
|
||||
ffi=auth_module._util.ffi,
|
||||
lib=SimpleNamespace(SSL_CTX_set_psk_client_callback=setter),
|
||||
)
|
||||
|
||||
|
||||
def _invoke_callback(callback, identity_size: int, key_size: int):
|
||||
ffi = auth_module._util.ffi
|
||||
identity_buffer = ffi.new("char[]", max(identity_size, 1))
|
||||
key_buffer = ffi.new("unsigned char[]", max(key_size, 1))
|
||||
copied = callback(
|
||||
ffi.NULL,
|
||||
ffi.NULL,
|
||||
identity_buffer,
|
||||
identity_size,
|
||||
key_buffer,
|
||||
key_size,
|
||||
)
|
||||
return (
|
||||
copied,
|
||||
bytes(ffi.buffer(identity_buffer, max(identity_size, 1))),
|
||||
bytes(ffi.buffer(key_buffer, max(key_size, 1))),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("key_length", [16, 32])
|
||||
def test_psk_auth_accepts_exact_supported_credential_lengths(key_length):
|
||||
provider = PskAuth(identity=_IDENTITY, key=b"k" * key_length)
|
||||
|
||||
assert repr(provider) == "PskAuth()"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("identity", "key"),
|
||||
[
|
||||
("i" * 16, _KEY),
|
||||
(bytearray(_IDENTITY), _KEY),
|
||||
(memoryview(_IDENTITY), _KEY),
|
||||
(_BytesSubclass(_IDENTITY), _KEY),
|
||||
(_IDENTITY, "k" * 16),
|
||||
(_IDENTITY, bytearray(_KEY)),
|
||||
(_IDENTITY, memoryview(_KEY)),
|
||||
(_IDENTITY, _BytesSubclass(_KEY)),
|
||||
],
|
||||
)
|
||||
def test_psk_auth_rejects_non_bytes_credentials(identity, key):
|
||||
with pytest.raises(TypeError, match="identity and key must be bytes"):
|
||||
PskAuth(identity=identity, key=key)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("identity_length", [0, 15, 17])
|
||||
def test_psk_auth_rejects_invalid_identity_lengths(identity_length):
|
||||
with pytest.raises(ValueError, match="raw 16-byte OCF UUID"):
|
||||
PskAuth(identity=b"i" * identity_length, key=_KEY)
|
||||
|
||||
|
||||
def test_psk_auth_rejects_identity_with_nul_byte():
|
||||
with pytest.raises(ValueError, match="cannot contain a NUL"):
|
||||
PskAuth(identity=b"i" * 15 + b"\x00", key=_KEY)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("key_length", [0, 15, 17, 31, 33])
|
||||
def test_psk_auth_rejects_invalid_key_lengths(key_length):
|
||||
with pytest.raises(ValueError, match="16 or 32 bytes"):
|
||||
PskAuth(identity=_IDENTITY, key=b"k" * key_length)
|
||||
|
||||
|
||||
def test_psk_auth_is_immutable_and_has_no_public_credential_surface():
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
|
||||
rendered = repr(provider)
|
||||
assert rendered == "PskAuth()"
|
||||
assert str(provider) == rendered
|
||||
assert _IDENTITY.decode() not in rendered
|
||||
assert _KEY.decode() not in rendered
|
||||
assert not hasattr(provider, "identity")
|
||||
assert not hasattr(provider, "key")
|
||||
assert not hasattr(provider, "_identity")
|
||||
assert not hasattr(provider, "_key")
|
||||
with pytest.raises(TypeError):
|
||||
vars(provider)
|
||||
with pytest.raises(TypeError):
|
||||
asdict(provider)
|
||||
with pytest.raises(AttributeError, match="immutable"):
|
||||
provider.identity = _OTHER_IDENTITY
|
||||
with pytest.raises(AttributeError, match="immutable"):
|
||||
del provider._callback
|
||||
|
||||
|
||||
def test_psk_auth_identity_equality_does_not_compare_credentials():
|
||||
first = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
second = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
|
||||
assert first != second
|
||||
assert len({first, second}) == 2
|
||||
|
||||
|
||||
def test_psk_callback_copies_exact_identity_and_key():
|
||||
installed = {}
|
||||
|
||||
def setter(context_handle, callback):
|
||||
installed["context"] = context_handle
|
||||
installed["callback"] = callback
|
||||
|
||||
context_handle = object()
|
||||
context = MagicMock()
|
||||
context._context = context_handle
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
|
||||
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
|
||||
provider.configure_context(context)
|
||||
|
||||
assert installed["context"] is context_handle
|
||||
callback = installed["callback"]
|
||||
copied, identity_bytes, key_bytes = _invoke_callback(callback, 17, 16)
|
||||
assert copied == 16
|
||||
assert identity_bytes == _IDENTITY + b"\x00"
|
||||
assert key_bytes == _KEY
|
||||
context.set_cipher_list.assert_called_once_with(
|
||||
b"ECDHE-PSK-AES128-CBC-SHA256:@SECLEVEL=0"
|
||||
)
|
||||
context.load_verify_locations.assert_not_called()
|
||||
context.set_verify.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("identity_size", "key_size"),
|
||||
[(16, 16), (17, 15)],
|
||||
)
|
||||
def test_psk_callback_rejects_short_buffers_without_partial_copy(
|
||||
identity_size,
|
||||
key_size,
|
||||
):
|
||||
installed = {}
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
context = MagicMock()
|
||||
context._context = object()
|
||||
|
||||
with patch.object(
|
||||
auth_module,
|
||||
"_util",
|
||||
_fake_openssl_util(
|
||||
lambda _context, callback: installed.setdefault(
|
||||
"callback", callback
|
||||
)
|
||||
),
|
||||
):
|
||||
provider.configure_context(context)
|
||||
|
||||
ffi = auth_module._util.ffi
|
||||
identity_buffer = ffi.new("char[]", 17)
|
||||
key_buffer = ffi.new("unsigned char[]", 16)
|
||||
ffi.memmove(identity_buffer, b"I" * 17, 17)
|
||||
ffi.memmove(key_buffer, b"K" * 16, 16)
|
||||
copied = installed["callback"](
|
||||
ffi.NULL,
|
||||
ffi.NULL,
|
||||
identity_buffer,
|
||||
identity_size,
|
||||
key_buffer,
|
||||
key_size,
|
||||
)
|
||||
|
||||
assert copied == 0
|
||||
assert bytes(ffi.buffer(identity_buffer, 17)) == b"I" * 17
|
||||
assert bytes(ffi.buffer(key_buffer, 16)) == b"K" * 16
|
||||
|
||||
|
||||
@pytest.mark.parametrize("null_buffer", ["identity", "key"])
|
||||
def test_psk_callback_rejects_null_buffers(null_buffer):
|
||||
installed = {}
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
context = MagicMock()
|
||||
context._context = object()
|
||||
|
||||
with patch.object(
|
||||
auth_module,
|
||||
"_util",
|
||||
_fake_openssl_util(
|
||||
lambda _context, callback: installed.setdefault(
|
||||
"callback", callback
|
||||
)
|
||||
),
|
||||
):
|
||||
provider.configure_context(context)
|
||||
|
||||
ffi = auth_module._util.ffi
|
||||
identity_buffer = ffi.new("char[]", 17)
|
||||
key_buffer = ffi.new("unsigned char[]", 16)
|
||||
ffi.memmove(identity_buffer, b"I" * 17, 17)
|
||||
ffi.memmove(key_buffer, b"K" * 16, 16)
|
||||
if null_buffer == "identity":
|
||||
identity_buffer = ffi.NULL
|
||||
else:
|
||||
key_buffer = ffi.NULL
|
||||
|
||||
copied = installed["callback"](
|
||||
ffi.NULL,
|
||||
ffi.NULL,
|
||||
identity_buffer,
|
||||
17,
|
||||
key_buffer,
|
||||
16,
|
||||
)
|
||||
assert copied == 0
|
||||
if null_buffer == "identity":
|
||||
assert bytes(ffi.buffer(key_buffer, 16)) == b"K" * 16
|
||||
else:
|
||||
assert bytes(ffi.buffer(identity_buffer, 17)) == b"I" * 17
|
||||
|
||||
|
||||
def test_psk_auth_unsupported_binding_error_contains_no_credentials():
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
context = MagicMock()
|
||||
context._context = object()
|
||||
unsupported_util = SimpleNamespace(
|
||||
ffi=auth_module._util.ffi,
|
||||
lib=SimpleNamespace(),
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(auth_module, "_util", unsupported_util),
|
||||
pytest.raises(RuntimeError) as captured,
|
||||
):
|
||||
provider.configure_context(context)
|
||||
|
||||
rendered = (
|
||||
str(captured.value)
|
||||
+ repr(captured.value)
|
||||
+ "".join(traceback.format_exception(captured.value))
|
||||
)
|
||||
assert _IDENTITY.decode() not in rendered
|
||||
assert _KEY.decode() not in rendered
|
||||
context.set_cipher_list.assert_not_called()
|
||||
|
||||
|
||||
def test_psk_auth_configures_real_openssl_context():
|
||||
context = SSL.Context(SSL.DTLS_METHOD)
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
|
||||
assert provider.configure_context(context) is None
|
||||
|
||||
|
||||
def test_distinct_psk_providers_do_not_share_callback_credentials():
|
||||
callbacks = []
|
||||
|
||||
def setter(_context, callback):
|
||||
callbacks.append(callback)
|
||||
|
||||
first = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
second = PskAuth(identity=_OTHER_IDENTITY, key=_OTHER_KEY)
|
||||
first_context = MagicMock()
|
||||
first_context._context = object()
|
||||
second_context = MagicMock()
|
||||
second_context._context = object()
|
||||
|
||||
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
|
||||
first.configure_context(first_context)
|
||||
second.configure_context(second_context)
|
||||
|
||||
assert callbacks[0] is not callbacks[1]
|
||||
with ThreadPoolExecutor(max_workers=2) as executor:
|
||||
first_future = executor.submit(_invoke_callback, callbacks[0], 17, 16)
|
||||
second_future = executor.submit(
|
||||
_invoke_callback,
|
||||
callbacks[1],
|
||||
17,
|
||||
32,
|
||||
)
|
||||
first_result = first_future.result()
|
||||
second_result = second_future.result()
|
||||
assert first_result == (16, _IDENTITY + b"\x00", _KEY)
|
||||
assert second_result == (32, _OTHER_IDENTITY + b"\x00", _OTHER_KEY)
|
||||
|
||||
|
||||
def test_session_retains_psk_callback_only_with_provider_lifetime():
|
||||
callback_reference = None
|
||||
|
||||
def setter(_context, callback):
|
||||
nonlocal callback_reference
|
||||
callback_reference = weakref.ref(callback)
|
||||
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
session = DtlsCoapSession(
|
||||
"appliance.invalid",
|
||||
49154,
|
||||
auth=provider,
|
||||
)
|
||||
context = MagicMock()
|
||||
context._context = object()
|
||||
with patch.object(auth_module, "_util", _fake_openssl_util(setter)):
|
||||
session.auth.configure_context(context)
|
||||
|
||||
del provider
|
||||
gc.collect()
|
||||
assert callback_reference is not None
|
||||
assert callback_reference() is not None
|
||||
assert _invoke_callback(callback_reference(), 17, 16) == (
|
||||
16,
|
||||
_IDENTITY + b"\x00",
|
||||
_KEY,
|
||||
)
|
||||
|
||||
del session
|
||||
gc.collect()
|
||||
assert callback_reference() is None
|
||||
|
||||
|
||||
def test_session_accepts_psk_provider_without_legacy_certificate_material():
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
session = DtlsCoapSession("appliance.invalid", 49154, auth=provider)
|
||||
|
||||
assert session.auth is provider
|
||||
assert session.cert_path is None
|
||||
assert session.key_path is None
|
||||
assert session.cert_pem is None
|
||||
assert session.key_pem is None
|
||||
|
||||
|
||||
def test_psk_handshake_rejection_does_not_expose_credentials():
|
||||
provider = PskAuth(identity=_IDENTITY, key=_KEY)
|
||||
session = DtlsCoapSession("appliance.invalid", 49154, auth=provider)
|
||||
context = MagicMock()
|
||||
context._context = object()
|
||||
connection = MagicMock()
|
||||
connection.do_handshake.side_effect = SSL.Error()
|
||||
udp_socket = MagicMock()
|
||||
endpoint = SimpleNamespace(sockaddr=("192.0.2.100", 49154))
|
||||
|
||||
with (
|
||||
patch.object(auth_module, "_util", _fake_openssl_util(lambda *_: None)),
|
||||
patch.object(session_module.SSL, "Context", return_value=context),
|
||||
patch.object(session_module.SSL, "Connection", return_value=connection),
|
||||
patch.object(
|
||||
session_module,
|
||||
"open_connected_udp_socket",
|
||||
return_value=(udp_socket, endpoint),
|
||||
),
|
||||
pytest.raises(SessionError) as captured,
|
||||
):
|
||||
session.connect()
|
||||
|
||||
rendered = (
|
||||
str(captured.value)
|
||||
+ repr(captured.value)
|
||||
+ "".join(traceback.format_exception(captured.value))
|
||||
)
|
||||
assert _IDENTITY.decode() not in rendered
|
||||
assert _KEY.decode() not in rendered
|
||||
udp_socket.close.assert_called_once_with()
|
||||
@@ -0,0 +1,218 @@
|
||||
"""Compatibility baseline for the published API and LocalThings consumer."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
|
||||
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
|
||||
from smartthings_local.ocf.state_cache import StateCache
|
||||
from smartthings_local.protocol.auth import (
|
||||
AuthenticationProvider,
|
||||
CertificateAuth,
|
||||
PskAuth,
|
||||
SamsungServerProfile,
|
||||
SamsungServerRole,
|
||||
ServerCertificateAuth,
|
||||
)
|
||||
from smartthings_local.protocol.dtls_session import (
|
||||
ConnectCancellation,
|
||||
DtlsCoapSession,
|
||||
)
|
||||
from smartthings_local.protocol.owner_psk import derive_mfg_certificate_owner_psk
|
||||
|
||||
|
||||
def _assert_compatible_signature(callable_object, expected: list[str]) -> None:
|
||||
"""Require the existing call surface while allowing safe extensions."""
|
||||
parameters = list(inspect.signature(callable_object).parameters.values())
|
||||
assert [parameter.name for parameter in parameters[: len(expected)]] == expected
|
||||
for parameter in parameters[len(expected) :]:
|
||||
assert (
|
||||
parameter.kind
|
||||
in (
|
||||
inspect.Parameter.VAR_POSITIONAL,
|
||||
inspect.Parameter.VAR_KEYWORD,
|
||||
)
|
||||
or parameter.default is not inspect.Parameter.empty
|
||||
)
|
||||
|
||||
|
||||
def test_dtls_session_constructor_keeps_file_memory_and_local_port_inputs():
|
||||
_assert_compatible_signature(
|
||||
DtlsCoapSession,
|
||||
[
|
||||
"host",
|
||||
"port",
|
||||
"cert_path",
|
||||
"key_path",
|
||||
"cert_pem",
|
||||
"key_pem",
|
||||
"on_notification",
|
||||
"mtu",
|
||||
"rate_limit_rps",
|
||||
"local_port",
|
||||
],
|
||||
)
|
||||
auth_parameter = inspect.signature(DtlsCoapSession).parameters["auth"]
|
||||
assert auth_parameter.kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert auth_parameter.default is None
|
||||
|
||||
|
||||
def test_certificate_auth_is_a_public_authentication_provider():
|
||||
provider = CertificateAuth.from_files("/synthetic/cert.pem", "/synthetic/key")
|
||||
assert isinstance(provider, AuthenticationProvider)
|
||||
|
||||
for factory in (CertificateAuth.from_files, CertificateAuth.from_memory):
|
||||
profile_parameter = inspect.signature(factory).parameters["server_profile"]
|
||||
assert profile_parameter.kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert profile_parameter.default is None
|
||||
|
||||
|
||||
def test_samsung_server_profile_is_public_and_explicitly_bound():
|
||||
parameters = inspect.signature(SamsungServerProfile.bound_device).parameters
|
||||
assert list(parameters) == [
|
||||
"expected_certificate_identity",
|
||||
"role",
|
||||
"additional_ca_pem",
|
||||
]
|
||||
assert (
|
||||
parameters["expected_certificate_identity"].default
|
||||
is inspect.Parameter.empty
|
||||
)
|
||||
assert parameters["role"].kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert parameters["role"].default is SamsungServerRole.HOME_APPLIANCE
|
||||
assert parameters["additional_ca_pem"].kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert parameters["additional_ca_pem"].default is None
|
||||
|
||||
|
||||
def test_server_certificate_auth_is_a_public_authentication_provider():
|
||||
profile = SamsungServerProfile.bound_device(
|
||||
"abababab-abab-abab-abab-abababababab",
|
||||
role=SamsungServerRole.VD_DEVICE,
|
||||
)
|
||||
provider = ServerCertificateAuth(server_profile=profile)
|
||||
assert isinstance(provider, AuthenticationProvider)
|
||||
session = DtlsCoapSession("device.example", 5684, auth=provider)
|
||||
assert session.auth is provider
|
||||
assert session.cert_path is None
|
||||
assert session.key_path is None
|
||||
assert session.cert_pem is None
|
||||
assert session.key_pem is None
|
||||
parameters = inspect.signature(ServerCertificateAuth).parameters
|
||||
assert list(parameters) == ["server_profile"]
|
||||
assert parameters["server_profile"].kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert parameters["server_profile"].default is inspect.Parameter.empty
|
||||
|
||||
|
||||
def test_psk_auth_is_a_public_authentication_provider():
|
||||
provider = PskAuth(identity=b"i" * 16, key=b"k" * 16)
|
||||
assert isinstance(provider, AuthenticationProvider)
|
||||
parameters = inspect.signature(PskAuth).parameters
|
||||
assert list(parameters) == ["identity", "key"]
|
||||
assert all(
|
||||
parameter.kind is inspect.Parameter.KEYWORD_ONLY
|
||||
and parameter.default is inspect.Parameter.empty
|
||||
for parameter in parameters.values()
|
||||
)
|
||||
|
||||
|
||||
def test_owner_psk_derivation_keeps_every_security_input_explicit():
|
||||
parameters = inspect.signature(
|
||||
derive_mfg_certificate_owner_psk
|
||||
).parameters
|
||||
assert list(parameters) == [
|
||||
"master_secret",
|
||||
"client_random",
|
||||
"server_random",
|
||||
"owner_uuid",
|
||||
"device_uuid",
|
||||
"cipher_name",
|
||||
"oxm_label",
|
||||
]
|
||||
assert all(
|
||||
parameter.kind is inspect.Parameter.KEYWORD_ONLY
|
||||
and parameter.default is inspect.Parameter.empty
|
||||
for parameter in parameters.values()
|
||||
)
|
||||
|
||||
|
||||
def test_dtls_session_keeps_current_consumer_methods():
|
||||
expected = {
|
||||
"close",
|
||||
"connect",
|
||||
"get",
|
||||
"join",
|
||||
"pace",
|
||||
"ping",
|
||||
"post",
|
||||
"refresh_observes",
|
||||
"start_reader",
|
||||
"subscribe",
|
||||
}
|
||||
assert expected <= set(dir(DtlsCoapSession))
|
||||
assert "abort" not in DtlsCoapSession.__dict__
|
||||
assert "quiesce_for_close" not in DtlsCoapSession.__dict__
|
||||
_assert_compatible_signature(DtlsCoapSession.connect, ["self"])
|
||||
connect_timeout = inspect.signature(DtlsCoapSession.connect).parameters[
|
||||
"timeout"
|
||||
]
|
||||
assert connect_timeout.kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert connect_timeout.default is None
|
||||
connect_cancel = inspect.signature(DtlsCoapSession.connect).parameters[
|
||||
"cancel"
|
||||
]
|
||||
assert connect_cancel.kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert connect_cancel.default is None
|
||||
assert callable(ConnectCancellation().set)
|
||||
_assert_compatible_signature(
|
||||
DtlsCoapSession.get,
|
||||
[
|
||||
"self",
|
||||
"path_segs",
|
||||
"query",
|
||||
"timeout",
|
||||
],
|
||||
)
|
||||
_assert_compatible_signature(
|
||||
DtlsCoapSession.post,
|
||||
[
|
||||
"self",
|
||||
"path_segs",
|
||||
"body_cbor",
|
||||
"timeout",
|
||||
],
|
||||
)
|
||||
_assert_compatible_signature(
|
||||
DtlsCoapSession.subscribe,
|
||||
["self", "path_segs"],
|
||||
)
|
||||
|
||||
|
||||
def test_state_cache_keeps_current_consumer_surface():
|
||||
_assert_compatible_signature(StateCache, ["descriptor"])
|
||||
expected = {
|
||||
"apply_optimistic",
|
||||
"apply_rep",
|
||||
"freshness_s",
|
||||
"get",
|
||||
"index_device_tree",
|
||||
"set_on_change",
|
||||
"snapshot",
|
||||
"stalest",
|
||||
}
|
||||
assert expected <= set(dir(StateCache))
|
||||
|
||||
|
||||
def test_observe_refresh_task_keeps_current_consumer_surface():
|
||||
_assert_compatible_signature(
|
||||
ObserveRefreshTask,
|
||||
[
|
||||
"session",
|
||||
"paths",
|
||||
"interval_s",
|
||||
"logger",
|
||||
],
|
||||
)
|
||||
_assert_compatible_signature(
|
||||
ObserveRefreshTask.run_forever,
|
||||
["self", "stop"],
|
||||
)
|
||||
@@ -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()
|
||||
@@ -0,0 +1,101 @@
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
import pytest
|
||||
|
||||
import setup_cert
|
||||
|
||||
# All of these drive the real `openssl` CLI the way setup_cert does.
|
||||
pytestmark = pytest.mark.skipif(
|
||||
shutil.which("openssl") is None, reason="openssl CLI not available")
|
||||
|
||||
UUID = "04700f20-1111-2222-3333-444455556666"
|
||||
|
||||
|
||||
def _make_ca(dir_path):
|
||||
"""A throwaway self-signed CA standing in for the AC14K_M signer."""
|
||||
cert = dir_path / "ca.pem"
|
||||
key = dir_path / "ca.key"
|
||||
subprocess.run(
|
||||
["openssl", "req", "-x509", "-newkey", "rsa:2048", "-nodes",
|
||||
"-keyout", str(key), "-out", str(cert), "-days", "1",
|
||||
"-subj", "/CN=AC14K_M"],
|
||||
check=True, capture_output=True)
|
||||
return cert, key
|
||||
|
||||
|
||||
def test_mint_cert_produces_sha1_leaf_with_uuid(tmp_path):
|
||||
ca_cert, ca_key = _make_ca(tmp_path)
|
||||
paths = setup_cert.mint_cert(
|
||||
UUID, ca_cert, ca_key, [ca_cert], tmp_path / "out")
|
||||
|
||||
for name in ("key", "leaf", "fullchain"):
|
||||
assert paths[name].exists() and paths[name].stat().st_size > 0
|
||||
|
||||
text = subprocess.run(
|
||||
["openssl", "x509", "-in", str(paths["leaf"]), "-noout", "-text"],
|
||||
check=True, capture_output=True, text=True).stdout
|
||||
assert "sha1WithRSAEncryption" in text # SHA-1 signed leaf
|
||||
assert f"URI:urn:uuid:{UUID}" in text # UUID in the SAN
|
||||
assert "1.3.6.1.4.1.51414" in text # custom OIDs parsed
|
||||
# fullchain is leaf + supplied chain
|
||||
assert paths["fullchain"].read_text().count("BEGIN CERTIFICATE") == 2
|
||||
|
||||
|
||||
def test_mint_cert_surfaces_openssl_error(tmp_path):
|
||||
"""A genuine signing failure raises CommandError carrying openssl's
|
||||
output, instead of a bare non-zero-exit traceback."""
|
||||
ca_cert, _ = _make_ca(tmp_path)
|
||||
with pytest.raises(setup_cert.CommandError) as exc:
|
||||
setup_cert.mint_cert(
|
||||
UUID, ca_cert, tmp_path / "missing.key", [ca_cert],
|
||||
tmp_path / "out")
|
||||
assert "command failed" in str(exc.value)
|
||||
assert len(str(exc.value)) > 40 # includes detail, not just an exit code
|
||||
|
||||
|
||||
def test_mint_cert_retries_when_sha1_signing_blocked(tmp_path, monkeypatch):
|
||||
"""Simulate a Fedora/RHEL crypto policy rejecting SHA-1: the first
|
||||
(plain) signing attempt fails, and the SHA-1-override retry recovers."""
|
||||
ca_cert, ca_key = _make_ca(tmp_path)
|
||||
real_run = setup_cert.run
|
||||
attempts = {"plain": 0}
|
||||
|
||||
def fake_run(cmd, **kw):
|
||||
# Only the plain attempt has no OPENSSL_CONF override in its env.
|
||||
if cmd[:3] == ["openssl", "x509", "-req"] and "env" not in kw:
|
||||
attempts["plain"] += 1
|
||||
raise setup_cert.CommandError(
|
||||
"error: sha1 signature disabled by crypto policy")
|
||||
return real_run(cmd, **kw)
|
||||
|
||||
monkeypatch.setattr(setup_cert, "run", fake_run)
|
||||
paths = setup_cert.mint_cert(
|
||||
UUID, ca_cert, ca_key, [ca_cert], tmp_path / "out")
|
||||
|
||||
assert attempts["plain"] == 1 # the plain path was exercised
|
||||
assert paths["leaf"].exists() # the override retry recovered
|
||||
|
||||
|
||||
def test_sha1_retry_does_not_give_openssl_3_config_to_libressl(monkeypatch):
|
||||
calls = []
|
||||
|
||||
def fake_run(cmd, **kw):
|
||||
calls.append((cmd, kw))
|
||||
if cmd == ["openssl", "version"]:
|
||||
return subprocess.CompletedProcess(cmd, 0, "LibreSSL 3.3.6\n", "")
|
||||
return subprocess.CompletedProcess(cmd, 0, "", "")
|
||||
|
||||
monkeypatch.setattr(setup_cert, "run", fake_run)
|
||||
monkeypatch.setenv("OPENSSL_CONF", "/synthetic/inherited.cnf")
|
||||
|
||||
setup_cert.run_allow_sha1(["openssl", "x509", "-req"])
|
||||
|
||||
assert calls[1][0] == ["openssl", "x509", "-req"]
|
||||
assert "OPENSSL_CONF" not in calls[1][1]["env"]
|
||||
|
||||
|
||||
def test_command_error_includes_stderr():
|
||||
with pytest.raises(setup_cert.CommandError) as exc:
|
||||
setup_cert.run(["openssl", "x509", "-in", "/no/such/file"])
|
||||
assert "command failed" in str(exc.value)
|
||||
@@ -0,0 +1,158 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
|
||||
from tools import check_share_safety
|
||||
|
||||
|
||||
def test_documentation_addresses_and_synthetic_uuid_are_safe():
|
||||
text = (
|
||||
"192.0.2.10 198.51.100.20 203.0.113.30 "
|
||||
"2001:db8::10 11111111-2222-3333-4444-555555555555"
|
||||
)
|
||||
assert check_share_safety.scan_text("fixture.txt", text) == []
|
||||
|
||||
|
||||
def test_dotted_object_identifiers_are_not_ipv4_addresses():
|
||||
text = "extendedKeyUsage = 1.3.6.1.4.1.51414.0.1.2"
|
||||
|
||||
assert check_share_safety.scan_text("fixture.txt", text) == []
|
||||
|
||||
|
||||
def test_public_github_attachment_uuid_is_safe_but_bare_uuid_is_not():
|
||||
value = "cc1dca15-f272-4625-" + "a13c-2dc82283ff95"
|
||||
public_url = f"https://github.com/user-attachments/assets/{value}"
|
||||
|
||||
assert check_share_safety.scan_text("README.md", public_url) == []
|
||||
assert check_share_safety.scan_text("fixture.txt", value) == [
|
||||
check_share_safety.Finding("fixture.txt", 1, "UUID")
|
||||
]
|
||||
|
||||
|
||||
def test_findings_never_echo_matched_content():
|
||||
cases = {
|
||||
"PEM_PRIVATE_KEY": "-----BEGIN " + "PRIVATE KEY-----",
|
||||
"EMAIL_ADDRESS": "person" + "@example.net",
|
||||
"MAC_ADDRESS": "aa:bb:cc:" + "dd:ee:ff",
|
||||
"NON_DOCUMENTATION_IPV4": "10." + "24.8.9",
|
||||
"NON_DOCUMENTATION_IPV6": "fd00" + 2 * chr(58) + "1234",
|
||||
"PRIVATE_DNS": "appliance" + chr(46) + "house" + chr(46) + "local",
|
||||
"HOME_PATH": "/" + "Users/person/private.txt",
|
||||
"CREDENTIAL_URL": "https://user:" + "pass" + chr(64) + "example.net/data",
|
||||
"SECRET_ASSIGNMENT": (
|
||||
"access_token " + chr(61) + " " + chr(34) + "never-print-this" + chr(34)
|
||||
),
|
||||
"SERIAL_ASSIGNMENT": (
|
||||
"serialNumber " + chr(61) + " " + chr(34) + "device-123456" + chr(34)
|
||||
),
|
||||
"REAL_TIMESTAMP": "2026-08-02" + "T12:34:56Z",
|
||||
"QR_PAYLOAD": "qr_" + "payload = value",
|
||||
"UUID": "12345678-1234-4234-9234-" + "123456789abc",
|
||||
}
|
||||
for rule_id, value in cases.items():
|
||||
findings = check_share_safety.scan_text("candidate.txt", value)
|
||||
rendered = "\n".join(finding.render() for finding in findings)
|
||||
assert f"candidate.txt:1:{rule_id}" in rendered
|
||||
assert value not in rendered
|
||||
|
||||
|
||||
def test_binary_and_archive_inputs_are_rejected(tmp_path):
|
||||
binary = tmp_path / "fixture.bin"
|
||||
binary.write_bytes(b"before\x00after")
|
||||
capture = tmp_path / "fixture.pcap"
|
||||
capture.write_text("text-looking content")
|
||||
|
||||
assert check_share_safety.scan_file(binary, "fixture.bin") == [
|
||||
check_share_safety.Finding("fixture.bin", 0, "BINARY_CONTENT")
|
||||
]
|
||||
assert check_share_safety.scan_file(capture, "fixture.pcap") == [
|
||||
check_share_safety.Finding("fixture.pcap", 0, "FORBIDDEN_FILE_TYPE")
|
||||
]
|
||||
|
||||
|
||||
def test_changed_paths_include_staged_unstaged_and_untracked_files(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
def git(*args):
|
||||
return subprocess.run(
|
||||
[
|
||||
"git",
|
||||
"-c",
|
||||
"commit.gpgsign=false",
|
||||
"-c",
|
||||
"user.name=Test",
|
||||
"-c",
|
||||
"user.email=" + "test" + chr(64) + "example.invalid",
|
||||
*args,
|
||||
],
|
||||
cwd=tmp_path,
|
||||
capture_output=True,
|
||||
check=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
git("init", "--quiet")
|
||||
baseline = tmp_path / "baseline.txt"
|
||||
baseline.write_text("before\n")
|
||||
git("add", "baseline.txt")
|
||||
git("commit", "--quiet", "-m", "baseline")
|
||||
base = git("rev-parse", "HEAD").stdout.strip()
|
||||
|
||||
staged = tmp_path / "staged.txt"
|
||||
staged.write_text("staged\n")
|
||||
git("add", "staged.txt")
|
||||
baseline.write_text("after\n")
|
||||
(tmp_path / "untracked.txt").write_text("untracked\n")
|
||||
|
||||
monkeypatch.chdir(tmp_path)
|
||||
assert check_share_safety._changed_paths(base) == [
|
||||
"baseline.txt",
|
||||
"staged.txt",
|
||||
"untracked.txt",
|
||||
]
|
||||
|
||||
|
||||
def test_committed_scan_ignores_unchanged_findings_but_checks_added_lines(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
def git(*args):
|
||||
return subprocess.run(
|
||||
[
|
||||
"git",
|
||||
"-c",
|
||||
"commit.gpgsign=false",
|
||||
"-c",
|
||||
"user.name=Test",
|
||||
"-c",
|
||||
"user.email=" + "test" + chr(64) + "example.invalid",
|
||||
*args,
|
||||
],
|
||||
cwd=tmp_path,
|
||||
capture_output=True,
|
||||
check=True,
|
||||
text=True,
|
||||
)
|
||||
|
||||
candidate = tmp_path / "candidate.txt"
|
||||
private_one = "10." + "24.8.9"
|
||||
private_two = "10." + "24.8.10"
|
||||
git("init", "--quiet")
|
||||
candidate.write_text(f"existing {private_one}\n")
|
||||
git("add", "candidate.txt")
|
||||
git("commit", "--quiet", "-m", "baseline")
|
||||
base = git("rev-parse", "HEAD").stdout.strip()
|
||||
|
||||
candidate.write_text(f"existing {private_one}\nsafe addition\n")
|
||||
git("add", "candidate.txt")
|
||||
git("commit", "--quiet", "-m", "safe change")
|
||||
monkeypatch.chdir(tmp_path)
|
||||
assert check_share_safety.check_changed(base) == []
|
||||
|
||||
candidate.write_text(
|
||||
f"existing {private_one}\nsafe addition\nintroduced {private_two}\n"
|
||||
)
|
||||
git("add", "candidate.txt")
|
||||
git("commit", "--quiet", "-m", "unsafe change")
|
||||
assert check_share_safety.check_changed(base) == [
|
||||
check_share_safety.Finding("candidate.txt", 3, "NON_DOCUMENTATION_IPV4")
|
||||
]
|
||||
@@ -0,0 +1,315 @@
|
||||
"""Deterministic baseline checks for the current OCF worker stop contract."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from mqtt_demo import bridge as bridge_module
|
||||
from mqtt_demo.bridge import PushBridge
|
||||
from smartthings_local.ocf.keepalive import KeepaliveTask
|
||||
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
|
||||
from smartthings_local.ocf.poll_scheduler import PollScheduler, PollTier
|
||||
from smartthings_local.ocf.state_cache import StateCache
|
||||
|
||||
_THREAD_DEADLINE_S = 2.0
|
||||
|
||||
|
||||
class _Session:
|
||||
def ping(self):
|
||||
return None
|
||||
|
||||
def refresh_observes(self, paths):
|
||||
return None
|
||||
|
||||
|
||||
class _Descriptor:
|
||||
def on_observation(self, state, href, rep):
|
||||
return None
|
||||
|
||||
|
||||
class _ObservedEvent(threading.Event):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.waiting = threading.Event()
|
||||
|
||||
def wait(self, timeout=None):
|
||||
self.waiting.set()
|
||||
return super().wait(timeout)
|
||||
|
||||
|
||||
def _assert_worker_stops(target, name: str):
|
||||
stop = _ObservedEvent()
|
||||
errors: list[str] = []
|
||||
|
||||
def run():
|
||||
try:
|
||||
target(stop)
|
||||
except Exception as error: # noqa: BLE001 # pragma: no cover
|
||||
errors.append(type(error).__name__)
|
||||
|
||||
worker = threading.Thread(target=run, name=name, daemon=True)
|
||||
worker.start()
|
||||
assert stop.waiting.wait(_THREAD_DEADLINE_S), (
|
||||
f"{name} did not enter an interruptible wait"
|
||||
)
|
||||
stop.set()
|
||||
worker.join(_THREAD_DEADLINE_S)
|
||||
assert not worker.is_alive(), f"{name} did not stop"
|
||||
assert errors == [], f"{name} raised {errors[0]}"
|
||||
|
||||
|
||||
def test_keepalive_worker_stops_without_waiting_for_interval():
|
||||
task = KeepaliveTask(_Session(), interval_s=3600.0)
|
||||
_assert_worker_stops(task.run_forever, "test-keepalive")
|
||||
|
||||
|
||||
def test_observe_refresh_worker_stops_without_waiting_for_interval():
|
||||
task = ObserveRefreshTask(_Session(), [], interval_s=3600.0)
|
||||
_assert_worker_stops(task.run_forever, "test-observe-refresh")
|
||||
|
||||
|
||||
def test_poll_scheduler_worker_stops_without_leaking_thread():
|
||||
scheduler = PollScheduler(
|
||||
_Session(),
|
||||
StateCache(_Descriptor()),
|
||||
[PollTier("idle", interval_s=3600.0, paths=())],
|
||||
)
|
||||
_assert_worker_stops(scheduler.run_forever, "test-poll-scheduler")
|
||||
|
||||
|
||||
class _SessionWorker:
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.stop = None
|
||||
self.started = threading.Event()
|
||||
self.exited = threading.Event()
|
||||
self.on_reachable = kwargs.get("on_reachable")
|
||||
self.on_unreachable = kwargs.get("on_unreachable")
|
||||
self.last_success_ts = 0.0
|
||||
|
||||
def run_forever(self, stop):
|
||||
self.stop = stop
|
||||
self.started.set()
|
||||
stop.wait()
|
||||
self.exited.set()
|
||||
|
||||
|
||||
class _JoinedSession:
|
||||
def __init__(self, workers, start_index, error=None):
|
||||
self.workers = workers
|
||||
self.start_index = start_index
|
||||
self.error = error
|
||||
|
||||
def join(self):
|
||||
assert all(
|
||||
worker.started.wait(_THREAD_DEADLINE_S)
|
||||
for worker in self.workers[self.start_index:]
|
||||
)
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
|
||||
class _BlockingJoinedSession(_JoinedSession):
|
||||
def __init__(self, workers):
|
||||
super().__init__(workers, 0)
|
||||
self.joined = threading.Event()
|
||||
self.release = threading.Event()
|
||||
|
||||
def join(self):
|
||||
super().join()
|
||||
self.joined.set()
|
||||
assert self.release.wait(_THREAD_DEADLINE_S)
|
||||
|
||||
|
||||
def _bridge():
|
||||
bridge = object.__new__(PushBridge)
|
||||
bridge.descriptor = SimpleNamespace(
|
||||
observe_paths=(),
|
||||
poll_tiers=[],
|
||||
is_active=lambda _state: False,
|
||||
)
|
||||
bridge.shared = SimpleNamespace(PING_INTERVAL_S=3600.0)
|
||||
bridge.app = SimpleNamespace(klass="test")
|
||||
bridge.log = SimpleNamespace(info=lambda *args: None, warning=lambda *args: None)
|
||||
bridge.cache = SimpleNamespace(links={})
|
||||
bridge.stop = threading.Event()
|
||||
bridge._session_stop_lock = threading.Lock()
|
||||
bridge._session_stop = None
|
||||
bridge.scheduler = None
|
||||
bridge.keepalive = None
|
||||
bridge.observe_refresh = None
|
||||
bridge._seed_from_device0 = lambda _session: None
|
||||
bridge._retag_logger_with_serial = lambda: None
|
||||
bridge.maybe_publish_state = lambda **kwargs: None
|
||||
bridge.set_availability = lambda _online: None
|
||||
return bridge
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("session_count", "join_error"),
|
||||
[
|
||||
pytest.param(2, None, id="reconnect"),
|
||||
pytest.param(1, RuntimeError("reader failed"), id="reader-error"),
|
||||
],
|
||||
)
|
||||
def test_bridge_retires_session_workers_before_returning(
|
||||
monkeypatch, session_count, join_error
|
||||
):
|
||||
workers = []
|
||||
|
||||
def worker_factory(*args, **kwargs):
|
||||
worker = _SessionWorker(*args, **kwargs)
|
||||
workers.append(worker)
|
||||
return worker
|
||||
|
||||
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
|
||||
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
|
||||
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
|
||||
|
||||
bridge = _bridge()
|
||||
|
||||
session_stops = []
|
||||
for session_index in range(session_count):
|
||||
start_index = len(workers)
|
||||
session = _JoinedSession(
|
||||
workers,
|
||||
start_index,
|
||||
error=join_error if session_index == session_count - 1 else None,
|
||||
)
|
||||
if session.error is None:
|
||||
bridge._run_session_inner(session)
|
||||
else:
|
||||
with pytest.raises(RuntimeError, match="reader failed"):
|
||||
bridge._run_session_inner(session)
|
||||
|
||||
session_workers = workers[start_index:]
|
||||
assert len(session_workers) == 3
|
||||
assert len({id(worker.stop) for worker in session_workers}) == 1
|
||||
session_stops.append(session_workers[0].stop)
|
||||
|
||||
assert len(workers) == session_count * 3
|
||||
assert len({id(stop) for stop in session_stops}) == session_count
|
||||
assert all(stop is not bridge.stop for stop in session_stops)
|
||||
assert all(stop.is_set() for stop in session_stops)
|
||||
assert not bridge.stop.is_set()
|
||||
assert all(worker.exited.is_set() for worker in workers)
|
||||
keepalive_workers = tuple(
|
||||
workers[index] for index in range(1, len(workers), 3)
|
||||
)
|
||||
assert all(worker.on_reachable is None for worker in keepalive_workers)
|
||||
assert all(worker.on_unreachable is None for worker in keepalive_workers)
|
||||
assert bridge.scheduler is None
|
||||
assert bridge.keepalive is None
|
||||
assert bridge.observe_refresh is None
|
||||
assert bridge._session_stop is None
|
||||
|
||||
|
||||
def test_bridge_retires_started_worker_when_later_thread_fails_to_start(
|
||||
monkeypatch,
|
||||
):
|
||||
workers = []
|
||||
|
||||
def worker_factory(*args, **kwargs):
|
||||
worker = _SessionWorker(*args, **kwargs)
|
||||
workers.append(worker)
|
||||
return worker
|
||||
|
||||
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
|
||||
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
|
||||
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
|
||||
|
||||
bridge = _bridge()
|
||||
|
||||
original_start = threading.Thread.start
|
||||
start_count = 0
|
||||
|
||||
def fail_second_start(thread):
|
||||
nonlocal start_count
|
||||
start_count += 1
|
||||
if start_count == 2:
|
||||
raise RuntimeError("synthetic thread start failure")
|
||||
original_start(thread)
|
||||
|
||||
monkeypatch.setattr(threading.Thread, "start", fail_second_start)
|
||||
|
||||
with pytest.raises(RuntimeError, match="synthetic thread start failure"):
|
||||
bridge._run_session_inner(SimpleNamespace(join=lambda: None))
|
||||
|
||||
assert len(workers) == 3
|
||||
assert workers[0].started.wait(_THREAD_DEADLINE_S)
|
||||
assert workers[0].exited.wait(_THREAD_DEADLINE_S)
|
||||
assert workers[0].stop is not bridge.stop
|
||||
assert workers[0].stop.is_set()
|
||||
assert workers[1].stop is None
|
||||
assert workers[2].stop is None
|
||||
assert bridge.scheduler is None
|
||||
assert bridge.keepalive is None
|
||||
assert bridge.observe_refresh is None
|
||||
assert bridge._session_stop is None
|
||||
|
||||
|
||||
def test_request_stop_wakes_current_session_workers(monkeypatch):
|
||||
workers = []
|
||||
|
||||
def worker_factory(*args, **kwargs):
|
||||
worker = _SessionWorker(*args, **kwargs)
|
||||
workers.append(worker)
|
||||
return worker
|
||||
|
||||
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
|
||||
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
|
||||
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
|
||||
|
||||
bridge = _bridge()
|
||||
session = _BlockingJoinedSession(workers)
|
||||
session_thread = threading.Thread(
|
||||
target=bridge._run_session_inner,
|
||||
args=(session,),
|
||||
daemon=True,
|
||||
)
|
||||
session_thread.start()
|
||||
|
||||
try:
|
||||
assert session.joined.wait(_THREAD_DEADLINE_S)
|
||||
session_stop = bridge._session_stop
|
||||
assert session_stop is not None
|
||||
assert not session_stop.is_set()
|
||||
|
||||
bridge.request_stop()
|
||||
|
||||
assert bridge.stop.is_set()
|
||||
assert session_stop.is_set()
|
||||
assert all(
|
||||
worker.exited.wait(_THREAD_DEADLINE_S) for worker in workers
|
||||
)
|
||||
assert session_thread.is_alive()
|
||||
finally:
|
||||
session.release.set()
|
||||
session_thread.join(_THREAD_DEADLINE_S)
|
||||
|
||||
assert not session_thread.is_alive()
|
||||
assert bridge._session_stop is None
|
||||
|
||||
|
||||
def test_session_workers_observe_stop_requested_before_handoff(monkeypatch):
|
||||
workers = []
|
||||
|
||||
def worker_factory(*args, **kwargs):
|
||||
worker = _SessionWorker(*args, **kwargs)
|
||||
workers.append(worker)
|
||||
return worker
|
||||
|
||||
monkeypatch.setattr(bridge_module, "PollScheduler", worker_factory)
|
||||
monkeypatch.setattr(bridge_module, "KeepaliveTask", worker_factory)
|
||||
monkeypatch.setattr(bridge_module, "ObserveRefreshTask", worker_factory)
|
||||
|
||||
bridge = _bridge()
|
||||
bridge.request_stop()
|
||||
bridge._run_session_inner(_JoinedSession(workers, 0))
|
||||
|
||||
assert len(workers) == 3
|
||||
assert all(worker.stop.is_set() for worker in workers)
|
||||
assert all(worker.exited.wait(_THREAD_DEADLINE_S) for worker in workers)
|
||||
assert bridge._session_stop is None
|
||||
Executable
+130
@@ -0,0 +1,130 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Verify that SmartThings-Local wheel and sdist contents are intentional."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import subprocess
|
||||
import tarfile
|
||||
import zipfile
|
||||
from pathlib import Path, PurePosixPath
|
||||
|
||||
|
||||
class DistributionError(RuntimeError):
|
||||
"""An artifact contains a missing, unexpected, or unsafe member."""
|
||||
|
||||
|
||||
def _tracked_files() -> set[str]:
|
||||
proc = subprocess.run(
|
||||
["git", "ls-files", "-z", "--", "smartthings_local", "tests"],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
)
|
||||
return {value.decode("utf-8") for value in proc.stdout.split(b"\0") if value}
|
||||
|
||||
|
||||
def _safe_member(name: str) -> bool:
|
||||
path = PurePosixPath(name)
|
||||
return bool(name) and not path.is_absolute() and ".." not in path.parts
|
||||
|
||||
|
||||
def _expected_package_files() -> set[str]:
|
||||
tracked = {
|
||||
path for path in _tracked_files() if path.startswith("smartthings_local/")
|
||||
}
|
||||
tracked.add("smartthings_local/_version.py")
|
||||
return tracked
|
||||
|
||||
|
||||
def check_wheel(path: Path) -> None:
|
||||
with zipfile.ZipFile(path) as archive:
|
||||
names = set(archive.namelist())
|
||||
if not names or any(not _safe_member(name) for name in names):
|
||||
raise DistributionError("wheel has an unsafe member")
|
||||
|
||||
package_files = {name for name in names if name.startswith("smartthings_local/")}
|
||||
if package_files != _expected_package_files():
|
||||
raise DistributionError(
|
||||
"wheel package contents differ from the tracked package"
|
||||
)
|
||||
|
||||
metadata = names - package_files
|
||||
roots = {name.split("/", 1)[0] for name in metadata}
|
||||
if len(roots) != 1:
|
||||
raise DistributionError("wheel must contain one dist-info directory")
|
||||
dist_info = roots.pop()
|
||||
if not dist_info.endswith(".dist-info"):
|
||||
raise DistributionError("wheel metadata directory is invalid")
|
||||
expected_metadata = {
|
||||
f"{dist_info}/METADATA",
|
||||
f"{dist_info}/WHEEL",
|
||||
f"{dist_info}/licenses/LICENSE",
|
||||
f"{dist_info}/licenses/NOTICE",
|
||||
f"{dist_info}/RECORD",
|
||||
}
|
||||
if metadata != expected_metadata:
|
||||
raise DistributionError("wheel metadata contents are unexpected")
|
||||
|
||||
|
||||
def check_sdist(path: Path) -> None:
|
||||
with tarfile.open(path, mode="r:gz") as archive:
|
||||
members = archive.getmembers()
|
||||
if not members or any(
|
||||
not member.isfile() or member.issym() or member.islnk() for member in members
|
||||
):
|
||||
raise DistributionError("sdist must contain regular files only")
|
||||
names = {member.name for member in members}
|
||||
if any(not _safe_member(name) for name in names):
|
||||
raise DistributionError("sdist has an unsafe member")
|
||||
|
||||
roots = {name.split("/", 1)[0] for name in names}
|
||||
if len(roots) != 1:
|
||||
raise DistributionError("sdist must contain one top-level directory")
|
||||
root = roots.pop()
|
||||
relative = {name[len(root) + 1 :] for name in names if name.startswith(f"{root}/")}
|
||||
required = _tracked_files() | {
|
||||
"LICENSE",
|
||||
"NOTICE",
|
||||
"PKG-INFO",
|
||||
"README.md",
|
||||
"pyproject.toml",
|
||||
"smartthings_local/_version.py",
|
||||
}
|
||||
# hatchling bundles the VCS ignore files it finds, but which ones ship
|
||||
# depends on the hatchling version (newer releases drop .hgignore), so
|
||||
# treat them as optional rather than exact members.
|
||||
optional = {".gitignore", ".hgignore"}
|
||||
if not required <= relative <= required | optional:
|
||||
raise DistributionError("sdist contents differ from the intended source set")
|
||||
|
||||
|
||||
def check_directory(directory: Path) -> None:
|
||||
wheels = sorted(directory.glob("*.whl"))
|
||||
sdists = sorted(directory.glob("*.tar.gz"))
|
||||
if len(wheels) != 1 or len(sdists) != 1:
|
||||
raise DistributionError("expected exactly one wheel and one sdist")
|
||||
check_wheel(wheels[0])
|
||||
check_sdist(sdists[0])
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("directory", type=Path)
|
||||
args = parser.parse_args()
|
||||
try:
|
||||
check_directory(args.directory)
|
||||
except (
|
||||
DistributionError,
|
||||
OSError,
|
||||
subprocess.SubprocessError,
|
||||
tarfile.TarError,
|
||||
zipfile.BadZipFile,
|
||||
):
|
||||
print("distribution check failed")
|
||||
return 1
|
||||
print("distribution contents verified")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Executable
+308
@@ -0,0 +1,308 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Check introduced public content for common private-data and secret shapes.
|
||||
|
||||
Findings contain only path, line, and rule ID. Matched content is never
|
||||
printed because it may itself be sensitive.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import ipaddress
|
||||
import re
|
||||
import subprocess
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
MAX_TEXT_BYTES = 2 * 1024 * 1024
|
||||
DOCUMENTATION_IPV4 = tuple(
|
||||
ipaddress.ip_network(value)
|
||||
for value in ("192.0.2.0/24", "198.51.100.0/24", "203.0.113.0/24")
|
||||
)
|
||||
DOCUMENTATION_IPV6 = ipaddress.ip_network("2001:db8::/32")
|
||||
SAFE_UUIDS = {
|
||||
"00000000-0000-0000-0000-000000000000",
|
||||
"11111111-2222-3333-4444-555555555555",
|
||||
}
|
||||
FORBIDDEN_SUFFIXES = {
|
||||
".7z",
|
||||
".apk",
|
||||
".cap",
|
||||
".db",
|
||||
".der",
|
||||
".gz",
|
||||
".jks",
|
||||
".key",
|
||||
".p12",
|
||||
".pcap",
|
||||
".pcapng",
|
||||
".pfx",
|
||||
".sqlite",
|
||||
".sqlite3",
|
||||
".tar",
|
||||
".tgz",
|
||||
".zip",
|
||||
}
|
||||
|
||||
PATTERNS = (
|
||||
(
|
||||
"PEM_PRIVATE_KEY",
|
||||
re.compile(r"-----BEGIN (?:RSA |EC |OPENSSH |DSA )?PRIVATE KEY-----"),
|
||||
),
|
||||
(
|
||||
"EMAIL_ADDRESS",
|
||||
re.compile(r"[A-Za-z0-9._%+-]+@(?:[A-Za-z0-9-]+\.)+[A-Za-z]{2,}"),
|
||||
),
|
||||
(
|
||||
"MAC_ADDRESS",
|
||||
re.compile(r"(?i)(?<![0-9a-f])(?:[0-9a-f]{2}[:-]){5}[0-9a-f]{2}(?![0-9a-f])"),
|
||||
),
|
||||
(
|
||||
"PRIVATE_DNS",
|
||||
re.compile(r"(?i)\b(?:[a-z0-9-]+\.)+(?:corp|home|internal|lan|local)\b"),
|
||||
),
|
||||
("HOME_PATH", re.compile(r"(?<![A-Za-z0-9._-])/(?:Users|home)/[^\s'\"`]+")),
|
||||
(
|
||||
"CREDENTIAL_URL",
|
||||
re.compile(
|
||||
r"(?i)\bhttps?://(?:[^\s/@:]+:[^\s/@]+@|[^\s?#]+[?&](?:access_token|api_key|password|refresh_token|token)=)"
|
||||
),
|
||||
),
|
||||
(
|
||||
"SECRET_ASSIGNMENT",
|
||||
re.compile(
|
||||
r"(?i)\b(?:access[_-]?token|api[_-]?key|bearer|owner[_-]?psk|password|passwd|private[_-]?key|psk|refresh[_-]?token|secret)\b\s*(?::|=)\s*(?:b|br|f|r|rb)?['\"][^'\"]+['\"]"
|
||||
),
|
||||
),
|
||||
(
|
||||
"SERIAL_ASSIGNMENT",
|
||||
re.compile(
|
||||
r"(?i)\b(?:device[_-]?)?serial(?:number|num)?\b\s*(?::|=)\s*['\"][^'\"]+['\"]"
|
||||
),
|
||||
),
|
||||
(
|
||||
"REAL_TIMESTAMP",
|
||||
re.compile(
|
||||
r"\b20[0-9]{2}-[01][0-9]-[0-3][0-9][T ][0-2][0-9]:[0-5][0-9](?::[0-6][0-9](?:\.[0-9]+)?)?(?:Z|[+-][0-2][0-9]:?[0-5][0-9])?\b"
|
||||
),
|
||||
),
|
||||
(
|
||||
"QR_PAYLOAD",
|
||||
re.compile(r"(?i)\b(?:qr[_-]?payload|setup[_-]?payload)\b\s*(?::|=)"),
|
||||
),
|
||||
)
|
||||
UUID_PATTERN = r"(?<![0-9a-f])[0-9a-f]{8}-[0-9a-f]{4}-[1-5][0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}(?![0-9a-f])"
|
||||
UUID_RE = re.compile(UUID_PATTERN, re.IGNORECASE)
|
||||
PUBLIC_GITHUB_ATTACHMENT_RE = re.compile(
|
||||
rf"https://github\.com/user-attachments/assets/(?P<uuid>{UUID_PATTERN})",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
IPV4_RE = re.compile(
|
||||
r"(?<![0-9.])(?:25[0-5]|2[0-4][0-9]|1?[0-9]{1,2})(?:\.(?:25[0-5]|2[0-4][0-9]|1?[0-9]{1,2})){3}(?![0-9.])"
|
||||
)
|
||||
IPV6_RE = re.compile(
|
||||
r"(?i)(?<![0-9a-f:])(?:\[)?(?:[0-9a-f]{0,4}:){2,7}[0-9a-f]{0,4}(?:%[A-Za-z0-9_.-]+)?(?:\])?(?![0-9a-f:])"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, order=True)
|
||||
class Finding:
|
||||
path: str
|
||||
line: int
|
||||
rule_id: str
|
||||
|
||||
def render(self) -> str:
|
||||
return f"{self.path}:{self.line}:{self.rule_id}"
|
||||
|
||||
|
||||
def _safe_ipv4(value: str) -> bool:
|
||||
address = ipaddress.ip_address(value)
|
||||
return (
|
||||
address.is_loopback
|
||||
or address.is_unspecified
|
||||
or any(address in network for network in DOCUMENTATION_IPV4)
|
||||
)
|
||||
|
||||
|
||||
def _safe_ipv6(value: str) -> bool:
|
||||
address = ipaddress.ip_address(value.strip("[]").split("%", 1)[0])
|
||||
return (
|
||||
address.is_loopback or address.is_unspecified or address in DOCUMENTATION_IPV6
|
||||
)
|
||||
|
||||
|
||||
def scan_text(path: str, text: str) -> list[Finding]:
|
||||
findings: set[Finding] = set()
|
||||
for line_number, line in enumerate(text.splitlines(), start=1):
|
||||
public_attachment_uuids = {
|
||||
match.group("uuid").lower()
|
||||
for match in PUBLIC_GITHUB_ATTACHMENT_RE.finditer(line)
|
||||
}
|
||||
for rule_id, pattern in PATTERNS:
|
||||
if pattern.search(line):
|
||||
findings.add(Finding(path, line_number, rule_id))
|
||||
for match in UUID_RE.finditer(line):
|
||||
value = match.group(0).lower()
|
||||
if value not in SAFE_UUIDS and value not in public_attachment_uuids:
|
||||
findings.add(Finding(path, line_number, "UUID"))
|
||||
for match in IPV4_RE.finditer(line):
|
||||
if not _safe_ipv4(match.group(0)):
|
||||
findings.add(Finding(path, line_number, "NON_DOCUMENTATION_IPV4"))
|
||||
for match in IPV6_RE.finditer(line):
|
||||
try:
|
||||
safe = _safe_ipv6(match.group(0))
|
||||
except ValueError:
|
||||
continue
|
||||
if not safe:
|
||||
findings.add(Finding(path, line_number, "NON_DOCUMENTATION_IPV6"))
|
||||
return sorted(findings)
|
||||
|
||||
|
||||
def scan_file(path: Path, display_path: str) -> list[Finding]:
|
||||
if path.is_symlink():
|
||||
return [Finding(display_path, 0, "SYMLINK")]
|
||||
if path.suffix.casefold() in FORBIDDEN_SUFFIXES:
|
||||
return [Finding(display_path, 0, "FORBIDDEN_FILE_TYPE")]
|
||||
data = path.read_bytes()
|
||||
if len(data) > MAX_TEXT_BYTES:
|
||||
return [Finding(display_path, 0, "FILE_TOO_LARGE")]
|
||||
if b"\x00" in data:
|
||||
return [Finding(display_path, 0, "BINARY_CONTENT")]
|
||||
try:
|
||||
text = data.decode("utf-8", errors="strict")
|
||||
except UnicodeDecodeError:
|
||||
return [Finding(display_path, 0, "NON_UTF8_CONTENT")]
|
||||
return scan_text(display_path, text)
|
||||
|
||||
|
||||
def _changed_paths(base: str) -> list[str]:
|
||||
commands = (
|
||||
[
|
||||
"git",
|
||||
"diff",
|
||||
"--name-only",
|
||||
"--diff-filter=ACMR",
|
||||
"-z",
|
||||
f"{base}..HEAD",
|
||||
"--",
|
||||
],
|
||||
["git", "diff", "--cached", "--name-only", "--diff-filter=ACMR", "-z", "--"],
|
||||
["git", "diff", "--name-only", "--diff-filter=ACMR", "-z", "--"],
|
||||
)
|
||||
changed = [
|
||||
subprocess.run(command, capture_output=True, check=True) for command in commands
|
||||
]
|
||||
untracked = subprocess.run(
|
||||
["git", "ls-files", "--others", "--exclude-standard", "-z"],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
)
|
||||
return sorted(
|
||||
{
|
||||
value.decode("utf-8")
|
||||
for output in (*(result.stdout for result in changed), untracked.stdout)
|
||||
for value in output.split(b"\0")
|
||||
if value
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _local_changed_paths() -> set[str]:
|
||||
commands = (
|
||||
["git", "diff", "--cached", "--name-only", "--diff-filter=ACMR", "-z", "--"],
|
||||
["git", "diff", "--name-only", "--diff-filter=ACMR", "-z", "--"],
|
||||
["git", "ls-files", "--others", "--exclude-standard", "-z"],
|
||||
)
|
||||
outputs = (
|
||||
subprocess.run(command, capture_output=True, check=True).stdout
|
||||
for command in commands
|
||||
)
|
||||
return {
|
||||
value.decode("utf-8")
|
||||
for output in outputs
|
||||
for value in output.split(b"\0")
|
||||
if value
|
||||
}
|
||||
|
||||
|
||||
HUNK_RE = re.compile(r"^@@ -\d+(?:,\d+)? \+(?P<start>\d+)(?:,(?P<count>\d+))? @@")
|
||||
|
||||
|
||||
def _introduced_lines(base: str, path: str) -> set[int]:
|
||||
result = subprocess.run(
|
||||
[
|
||||
"git",
|
||||
"diff",
|
||||
"--no-color",
|
||||
"--no-ext-diff",
|
||||
"--unified=0",
|
||||
"--diff-filter=ACMR",
|
||||
f"{base}..HEAD",
|
||||
"--",
|
||||
path,
|
||||
],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
text=True,
|
||||
)
|
||||
lines: set[int] = set()
|
||||
for value in result.stdout.splitlines():
|
||||
match = HUNK_RE.match(value)
|
||||
if match is None:
|
||||
continue
|
||||
start = int(match.group("start"))
|
||||
count = int(match.group("count") or 1)
|
||||
lines.update(range(start, start + count))
|
||||
return lines
|
||||
|
||||
|
||||
def check_changed(base: str) -> list[Finding]:
|
||||
"""Scan introduced committed lines and all local-only file content."""
|
||||
local_paths = _local_changed_paths()
|
||||
findings: list[Finding] = []
|
||||
for path in _changed_paths(base):
|
||||
path_findings = scan_file(Path(path), path)
|
||||
if path in local_paths:
|
||||
findings.extend(path_findings)
|
||||
continue
|
||||
introduced = _introduced_lines(base, path)
|
||||
findings.extend(
|
||||
finding
|
||||
for finding in path_findings
|
||||
if finding.line == 0 or finding.line in introduced
|
||||
)
|
||||
return sorted(set(findings))
|
||||
|
||||
|
||||
def check_paths(paths: Iterable[str]) -> list[Finding]:
|
||||
findings: list[Finding] = []
|
||||
for value in paths:
|
||||
findings.extend(scan_file(Path(value), value))
|
||||
return sorted(set(findings))
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("paths", nargs="*")
|
||||
parser.add_argument("--changed-since")
|
||||
args = parser.parse_args()
|
||||
try:
|
||||
paths = _changed_paths(args.changed_since) if args.changed_since else args.paths
|
||||
if not paths:
|
||||
raise ValueError("no paths selected")
|
||||
findings = (
|
||||
check_changed(args.changed_since)
|
||||
if args.changed_since
|
||||
else check_paths(paths)
|
||||
)
|
||||
except (OSError, UnicodeDecodeError, subprocess.SubprocessError, ValueError):
|
||||
print("share-safety scan failed closed:SCAN_ERROR")
|
||||
return 2
|
||||
for finding in findings:
|
||||
print(finding.render())
|
||||
return 1 if findings else 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user