Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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,95 @@ 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:
|
||||
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
|
||||
sess = DtlsCoapSession("192.168.1.100", 49154, cert_pem=cert_pem, key_pem=key_pem)
|
||||
auth = CertificateAuth.from_memory(cert_pem, key_pem)
|
||||
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
|
||||
```
|
||||
|
||||
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.
|
||||
|
||||
### 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 +155,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 and need separate authentication work.
|
||||
|
||||
---
|
||||
|
||||
@@ -75,23 +164,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 may still require an unsupported authentication profile.
|
||||
- **`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 +194,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 +224,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 +246,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 +271,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 +303,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 +439,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 +505,10 @@ 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
|
||||
ocf_root_ca.pem Samsung OCF root CA, bundled for handshake verification
|
||||
ocf/ OCF resource + state layer (reusable)
|
||||
__init__.py
|
||||
@@ -426,7 +535,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)
|
||||
.github/workflows/publish.yml Build + PyPI Trusted Publishing on `v*` tags
|
||||
```
|
||||
|
||||
@@ -478,3 +587,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.
|
||||
+50
-52
@@ -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,14 +72,15 @@ 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
|
||||
@@ -263,35 +266,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 +295,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()
|
||||
|
||||
+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,246 @@
|
||||
"""Immutable authentication providers for DTLS sessions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from os import PathLike
|
||||
from pathlib import Path
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
from OpenSSL import SSL, _util, crypto
|
||||
|
||||
_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"
|
||||
_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 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",
|
||||
)
|
||||
|
||||
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,
|
||||
) -> 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"
|
||||
)
|
||||
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)
|
||||
|
||||
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],
|
||||
) -> CertificateAuth:
|
||||
"""Create a provider backed by certificate-chain and key files."""
|
||||
return cls(
|
||||
certificate_path=certificate_path,
|
||||
private_key_path=private_key_path,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_memory(
|
||||
cls,
|
||||
certificate_pem: str,
|
||||
private_key_pem: str,
|
||||
) -> 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,
|
||||
)
|
||||
|
||||
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."""
|
||||
context.load_verify_locations(_OCF_ROOT_CA)
|
||||
context.set_verify(SSL.VERIFY_PEER, _verify_peer)
|
||||
# @SECLEVEL=0 permits SHA-1 in Samsung's server cert chain (AC14K_M
|
||||
# intermediate is SHA-1 signed). This is the only channel that reaches
|
||||
# the OpenSSL instance cryptography bundles; ctypes and cffi bindings
|
||||
# do not expose SSL_CTX_set_security_level on this build.
|
||||
context.set_cipher_list(_DTLS_CIPHERS)
|
||||
if self._certificate_pem is not None:
|
||||
_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) # noqa: SLF001
|
||||
|
||||
|
||||
__all__ = ["AuthenticationProvider", "CertificateAuth", "PskAuth"]
|
||||
@@ -7,6 +7,8 @@ independently.
|
||||
"""
|
||||
import struct
|
||||
|
||||
from ..errors import MalformedMessageError
|
||||
|
||||
# CoAP option numbers (RFC 7252 + 7641 + 7959)
|
||||
URI_PATH = 11
|
||||
URI_QUERY = 15
|
||||
@@ -81,7 +83,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 +91,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
|
||||
|
||||
@@ -23,18 +23,25 @@ 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 .coap import split_dtls
|
||||
from .dtls_session import _OCF_ROOT_CA, _load_pem_chain
|
||||
from .dtls_session import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
|
||||
from .endpoint import open_connected_udp_socket
|
||||
|
||||
# DTLS record content types (RFC 6347 §4.1)
|
||||
_CT_CHANGE_CIPHER_SPEC = 20
|
||||
@@ -90,6 +97,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 +496,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 +577,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,40 +589,53 @@ 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)
|
||||
try:
|
||||
sock, _endpoint = open_connected_udp_socket(
|
||||
host,
|
||||
port,
|
||||
family=family,
|
||||
timeout=min(0.5, timeout),
|
||||
)
|
||||
except OSError:
|
||||
result.error = ProbeError()
|
||||
return result
|
||||
|
||||
t0 = time.time()
|
||||
started = time.monotonic()
|
||||
deadline = started + timeout
|
||||
seen = set()
|
||||
retransmits = 0
|
||||
try:
|
||||
while time.time() - t0 < timeout:
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
conn.do_handshake()
|
||||
result.outcome = COMPLETED
|
||||
if result.rtt_s is None:
|
||||
result.rtt_s = time.time() - t0
|
||||
result.rtt_s = time.monotonic() - started
|
||||
break
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
except SSL.Error as e:
|
||||
except SSL.Error:
|
||||
# A fatal Alert lands here; the alert record was already
|
||||
# captured below, so classification still works.
|
||||
result.error = str(e)
|
||||
result.error = ProbeError()
|
||||
break
|
||||
|
||||
try:
|
||||
o = conn.bio_read(65535)
|
||||
if o:
|
||||
for r in split_dtls(o):
|
||||
sock.sendto(r, dest)
|
||||
if sock.send(r) != len(r):
|
||||
raise OSError('short UDP send')
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
break
|
||||
sock.settimeout(min(0.5, remaining))
|
||||
try:
|
||||
d, _ = sock.recvfrom(65535)
|
||||
except socket.timeout:
|
||||
d = sock.recv(65535)
|
||||
except TimeoutError:
|
||||
# No answer to the last flight. Service OpenSSL's DTLS
|
||||
# retransmit timer: once it has counted down to 0,
|
||||
# handle_timeout() re-queues the previous flight into the
|
||||
@@ -250,9 +655,8 @@ def probe(host, port, *, cert_pem=None, key_pem=None,
|
||||
continue
|
||||
|
||||
if result.rtt_s is None:
|
||||
result.rtt_s = time.time() - t0
|
||||
result.rtt_s = time.monotonic() - started
|
||||
result.datagrams.append(d)
|
||||
server_flight = False
|
||||
for ct, detail in classify_datagram(d):
|
||||
if ct == _CT_HANDSHAKE:
|
||||
if detail not in seen:
|
||||
@@ -260,23 +664,14 @@ def probe(host, port, *, cert_pem=None, key_pem=None,
|
||||
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)
|
||||
except OSError:
|
||||
result.error = ProbeError()
|
||||
finally:
|
||||
sock.close()
|
||||
|
||||
@@ -288,14 +683,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 +695,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)]
|
||||
|
||||
|
||||
@@ -18,15 +18,21 @@ Reader thread owns the UDP socket. Callers issue get()/post() and block
|
||||
on a per-token Event the reader signals. OBSERVE notifications are
|
||||
delivered via the on_notification callback.
|
||||
"""
|
||||
import errno
|
||||
import os
|
||||
import re as _re
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from OpenSSL import SSL
|
||||
|
||||
from ..errors import (
|
||||
BlockwiseError,
|
||||
EndpointError,
|
||||
SessionClosedError,
|
||||
SessionError,
|
||||
SessionTimeoutError,
|
||||
)
|
||||
from .coap import (
|
||||
URI_PATH, URI_QUERY, OBSERVE, CONTENT_FORMAT, ACCEPT, BLOCK2, SIZE2,
|
||||
TYPE_CON, TYPE_NON, TYPE_ACK, TYPE_RST,
|
||||
@@ -35,13 +41,18 @@ from .coap import (
|
||||
encode_options, parse_coap, build_coap, block_value, fmt_code,
|
||||
split_dtls as _split_dtls,
|
||||
)
|
||||
from .auth import (
|
||||
AuthenticationProvider,
|
||||
CertificateAuth,
|
||||
_DTLS_CIPHERS,
|
||||
_OCF_ROOT_CA,
|
||||
_load_pem_chain,
|
||||
)
|
||||
from .endpoint import open_connected_udp_socket
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_OCF_ROOT_CA = str(Path(__file__).parent / 'ocf_root_ca.pem')
|
||||
|
||||
|
||||
# Diagnostic logging — when DEBUG_BRIDGE=1 in env, the bridge dumps
|
||||
# every received CoAP frame, every /operational/state/vs/0 + /oven/vs/0
|
||||
# + /power/vs/0 + /mode/vs/0-options rep change, the full link tree at
|
||||
@@ -62,32 +73,21 @@ _BLOCK_ACK_TIMEOUT = 4.0
|
||||
# once the ceiling is measured empirically.
|
||||
_DEFAULT_RATE_LIMIT_RPS = 5.0
|
||||
|
||||
_PEM_CERT_RE = _re.compile(
|
||||
rb'-----BEGIN CERTIFICATE-----.*?-----END CERTIFICATE-----',
|
||||
_re.DOTALL,
|
||||
# ICMP errors a connected UDP socket surfaces on the next recv. On these
|
||||
# appliances they show up while the device is rebooting, while it holds an
|
||||
# orphaned association, or across a router blip, and the next datagram
|
||||
# usually works. UDP delivery was never guaranteed, so treat them as
|
||||
# advisory and keep reading. Unconnected sockets never see any of this,
|
||||
# which is why the reader survived them before the connected-socket change
|
||||
# in d677c72 (v0.1.3).
|
||||
_ADVISORY_ERRNOS = frozenset(
|
||||
value for value in (
|
||||
getattr(errno, name, None)
|
||||
for name in ('ECONNREFUSED', 'EHOSTUNREACH', 'ENETUNREACH',
|
||||
'EHOSTDOWN', 'ENETDOWN')
|
||||
) if value is not None
|
||||
)
|
||||
|
||||
|
||||
def _load_pem_chain(ctx: SSL.Context, cert_pem: str, key_pem: str) -> None:
|
||||
"""Load a PEM cert chain and private key into an SSL context in memory.
|
||||
|
||||
Parses all certificate blocks from cert_pem: the first is the leaf
|
||||
(use_certificate), the rest are intermediates (add_extra_chain_cert).
|
||||
No temp files are written.
|
||||
"""
|
||||
from OpenSSL import crypto
|
||||
certs = _PEM_CERT_RE.findall(cert_pem.encode())
|
||||
if not certs:
|
||||
raise ValueError("No certificates found in cert_pem")
|
||||
ctx.use_certificate(crypto.load_certificate(crypto.FILETYPE_PEM, certs[0]))
|
||||
for extra in certs[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()
|
||||
|
||||
|
||||
class DtlsCoapSession:
|
||||
"""Single sustained DTLS-CoAP session.
|
||||
|
||||
@@ -100,10 +100,9 @@ class DtlsCoapSession:
|
||||
code, _ = sess.post(['mode','vs','0'], cbor)
|
||||
sess.close()
|
||||
|
||||
Cert material comes from either a file pair (cert_path, key_path) or
|
||||
an in-memory PEM pair (cert_pem, key_pem) — exactly one pair required.
|
||||
The in-memory path exists for callers (e.g. an HA config flow) that
|
||||
mint a client cert at runtime and never write it to disk.
|
||||
Authentication comes from an immutable provider. For compatibility,
|
||||
cert_path/key_path and cert_pem/key_pem create a CertificateAuth provider
|
||||
internally — exactly one legacy pair is required when auth is omitted.
|
||||
"""
|
||||
|
||||
HANDSHAKE_TIMEOUT_S = 12.0
|
||||
@@ -114,17 +113,24 @@ class DtlsCoapSession:
|
||||
cert_pem=None, key_pem=None,
|
||||
on_notification=None, mtu=1200,
|
||||
rate_limit_rps: float = _DEFAULT_RATE_LIMIT_RPS,
|
||||
local_port=None):
|
||||
if (cert_path is not None or key_path is not None) and \
|
||||
(cert_pem is not None or key_pem is not None):
|
||||
local_port=None, family=socket.AF_UNSPEC,
|
||||
auth: AuthenticationProvider | None = None):
|
||||
file_supplied = cert_path is not None or key_path is not None
|
||||
memory_supplied = cert_pem is not None or key_pem is not None
|
||||
if auth is not None and (file_supplied or memory_supplied):
|
||||
raise ValueError(
|
||||
"pass auth or legacy certificate arguments, not both")
|
||||
if auth is None and file_supplied and memory_supplied:
|
||||
raise ValueError(
|
||||
"pass either cert_path/key_path or cert_pem/key_pem, not both")
|
||||
if cert_pem is not None or key_pem is not None:
|
||||
if auth is None and memory_supplied:
|
||||
if cert_pem is None or key_pem is None:
|
||||
raise ValueError("cert_pem and key_pem must be passed together")
|
||||
elif cert_path is None or key_path is None:
|
||||
elif auth is None and (cert_path is None or key_path is None):
|
||||
raise ValueError(
|
||||
"must pass either cert_path/key_path or cert_pem/key_pem")
|
||||
if auth is not None and not isinstance(auth, AuthenticationProvider):
|
||||
raise TypeError("auth must implement AuthenticationProvider")
|
||||
|
||||
self.host = host
|
||||
self.port = port
|
||||
@@ -132,6 +138,12 @@ class DtlsCoapSession:
|
||||
self.key_path = str(key_path) if key_path is not None else None
|
||||
self.cert_pem = cert_pem
|
||||
self.key_pem = key_pem
|
||||
if auth is None:
|
||||
if cert_pem is not None:
|
||||
auth = CertificateAuth.from_memory(cert_pem, key_pem)
|
||||
else:
|
||||
auth = CertificateAuth.from_files(self.cert_path, self.key_path)
|
||||
self.auth = auth
|
||||
self.on_notification = on_notification # fn(href, payload_bytes)
|
||||
self.mtu = mtu
|
||||
self._min_req_interval = 1.0 / rate_limit_rps
|
||||
@@ -146,10 +158,12 @@ class DtlsCoapSession:
|
||||
# the new handshake and discard the old association. Verified
|
||||
# accepted by RT-OCF (oven, 2026-07-26).
|
||||
self.local_port = local_port
|
||||
self.family = family
|
||||
|
||||
self.sock = None
|
||||
self.conn = None
|
||||
self.dest = None
|
||||
self.endpoint = None
|
||||
|
||||
self._send_lock = threading.Lock()
|
||||
# Randomize MID and token counter starting points so reconnects
|
||||
@@ -170,6 +184,10 @@ class DtlsCoapSession:
|
||||
|
||||
self._stop = threading.Event()
|
||||
self._reader_thread = None
|
||||
# Set while the reader owns the socket. Cleared when it exits for
|
||||
# any reason, so callers fail fast through _check_live() instead of
|
||||
# waiting out a request timeout against a session nobody is reading.
|
||||
self._reader_running = threading.Event()
|
||||
self._last_send_ts = 0.0
|
||||
|
||||
def pace(self) -> None:
|
||||
@@ -185,78 +203,97 @@ class DtlsCoapSession:
|
||||
"""DTLS handshake. Blocks up to HANDSHAKE_TIMEOUT_S. Raises
|
||||
ConnectionError / TimeoutError on failure."""
|
||||
ctx = SSL.Context(SSL.DTLS_METHOD)
|
||||
|
||||
ctx.load_verify_locations(_OCF_ROOT_CA)
|
||||
ctx.set_verify(SSL.VERIFY_PEER, lambda conn, cert, err, depth, ok: ok)
|
||||
# @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.
|
||||
ctx.set_cipher_list(b'ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0')
|
||||
if self.cert_pem is not None:
|
||||
_load_pem_chain(ctx, self.cert_pem, self.key_pem)
|
||||
else:
|
||||
ctx.use_certificate_chain_file(self.cert_path)
|
||||
ctx.use_privatekey_file(self.key_path)
|
||||
ctx.check_privatekey()
|
||||
self.auth.configure_context(ctx)
|
||||
|
||||
conn = SSL.Connection(ctx, None)
|
||||
conn.set_connect_state()
|
||||
conn.set_ciphertext_mtu(self.mtu)
|
||||
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
if self.local_port is not None:
|
||||
# Fixed source port → same 5-tuple on reconnect, so the device
|
||||
# evicts any orphaned association per RFC 6347 §4.2.8 instead
|
||||
# of serving a second one alongside it. See __init__.
|
||||
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
sock.bind(('', self.local_port))
|
||||
sock.settimeout(2.0)
|
||||
dest = (self.host, self.port)
|
||||
sock, endpoint = open_connected_udp_socket(
|
||||
self.host,
|
||||
self.port,
|
||||
family=self.family,
|
||||
local_port=self.local_port,
|
||||
timeout=2.0,
|
||||
)
|
||||
dest = endpoint.sockaddr
|
||||
|
||||
t0 = time.time()
|
||||
backend_failed = False
|
||||
while time.time() - t0 < self.HANDSHAKE_TIMEOUT_S:
|
||||
try:
|
||||
conn.do_handshake()
|
||||
break
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
except SSL.Error as e:
|
||||
except SSL.Error:
|
||||
sock.close()
|
||||
raise ConnectionError(f"DTLS handshake error: {e}") from e
|
||||
backend_failed = True
|
||||
break
|
||||
send_failed = False
|
||||
try:
|
||||
o = conn.bio_read(65535)
|
||||
if o:
|
||||
for r in _split_dtls(o):
|
||||
sock.sendto(r, dest)
|
||||
if sock.send(r) != len(r):
|
||||
raise OSError('incomplete UDP send')
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
except OSError:
|
||||
sock.close()
|
||||
send_failed = True
|
||||
if send_failed:
|
||||
raise EndpointError() from OSError('UDP send failed')
|
||||
receive_failed = False
|
||||
try:
|
||||
d, _ = sock.recvfrom(65535)
|
||||
d = sock.recv(65535)
|
||||
if d:
|
||||
conn.bio_write(d)
|
||||
except socket.timeout:
|
||||
pass
|
||||
except OSError:
|
||||
sock.close()
|
||||
receive_failed = True
|
||||
if receive_failed:
|
||||
raise EndpointError() from OSError('UDP receive failed')
|
||||
time.sleep(0.05)
|
||||
else:
|
||||
sock.close()
|
||||
raise TimeoutError(
|
||||
f"DTLS handshake timeout to {self.host}:{self.port}")
|
||||
raise SessionTimeoutError()
|
||||
if backend_failed:
|
||||
raise SessionError() from ConnectionError('DTLS backend failed')
|
||||
|
||||
self.sock = sock
|
||||
self.conn = conn
|
||||
self.dest = dest
|
||||
self.endpoint = endpoint
|
||||
self._stop.clear()
|
||||
|
||||
def start_reader(self):
|
||||
"""Spawn the reader thread. Must be called after connect()."""
|
||||
if self.sock is None:
|
||||
raise RuntimeError("connect() before start_reader()")
|
||||
self._reader_running.set()
|
||||
t = threading.Thread(target=self._reader_loop,
|
||||
daemon=True, name='dtls-reader')
|
||||
t.start()
|
||||
self._reader_thread = t
|
||||
|
||||
def _check_live(self):
|
||||
"""Raise if the session cannot carry a request. A dead reader is
|
||||
as fatal as a closed connection: the socket may still accept
|
||||
sends, but no response will ever be dispatched, so waiting out the
|
||||
request timeout only delays the inevitable SessionClosedError.
|
||||
|
||||
Callers that never start a reader (config-flow style) keep the old
|
||||
behaviour — only the conn check applies while _reader_thread is
|
||||
None."""
|
||||
if self.conn is None:
|
||||
raise SessionClosedError()
|
||||
if self._reader_thread is not None and \
|
||||
not self._reader_running.is_set():
|
||||
raise SessionClosedError()
|
||||
|
||||
def join(self):
|
||||
"""Block until the reader thread exits (i.e. socket dies)."""
|
||||
if self._reader_thread is not None:
|
||||
@@ -303,12 +340,14 @@ class DtlsCoapSession:
|
||||
except Exception:
|
||||
pass
|
||||
for tok, (ev, container) in list(self._pending.items()):
|
||||
container.setdefault('err', 'socket closed')
|
||||
container.setdefault('err', SessionClosedError())
|
||||
ev.set()
|
||||
self._pending.clear()
|
||||
self._observe_tokens.clear()
|
||||
self.sock = None
|
||||
self.conn = None
|
||||
self.dest = None
|
||||
self.endpoint = None
|
||||
|
||||
# ---- send / receive plumbing -------------------------------------
|
||||
|
||||
@@ -340,7 +379,8 @@ class DtlsCoapSession:
|
||||
BIO-drain so two writers can't interleave records."""
|
||||
with self._send_lock:
|
||||
if self.conn is None:
|
||||
raise ConnectionError("DTLS session closed")
|
||||
raise SessionClosedError()
|
||||
send_failed = False
|
||||
try:
|
||||
self.conn.send(datagram)
|
||||
self._last_send_ts = time.monotonic()
|
||||
@@ -349,9 +389,14 @@ class DtlsCoapSession:
|
||||
if not o:
|
||||
break
|
||||
for r in _split_dtls(o):
|
||||
self.sock.sendto(r, self.dest)
|
||||
if self.sock.send(r) != len(r):
|
||||
raise OSError('incomplete UDP send')
|
||||
except SSL.WantReadError:
|
||||
pass
|
||||
except OSError:
|
||||
send_failed = True
|
||||
if send_failed:
|
||||
raise EndpointError() from OSError('UDP send failed')
|
||||
|
||||
def _reader_loop(self):
|
||||
"""Pump UDP socket → DTLS BIO → CoAP parser. Demuxes to pending
|
||||
@@ -362,10 +407,26 @@ class DtlsCoapSession:
|
||||
try:
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
d, _ = sock.recvfrom(65535)
|
||||
d = sock.recv(65535)
|
||||
except socket.timeout:
|
||||
continue
|
||||
except (OSError, ValueError):
|
||||
except OSError as e:
|
||||
if self._stop.is_set():
|
||||
return # close() got here first
|
||||
if e.errno in _ADVISORY_ERRNOS:
|
||||
logger.debug("reader: advisory %s from %s, continuing",
|
||||
errno.errorcode.get(e.errno, e.errno),
|
||||
self.host)
|
||||
continue
|
||||
logger.warning("reader exiting: socket error %s from %s",
|
||||
errno.errorcode.get(e.errno, e.errno),
|
||||
self.host)
|
||||
return
|
||||
except ValueError:
|
||||
# recv on a socket closed underneath the reader.
|
||||
if not self._stop.is_set():
|
||||
logger.warning("reader exiting: socket closed "
|
||||
"underneath it")
|
||||
return
|
||||
if not d:
|
||||
continue
|
||||
@@ -410,9 +471,11 @@ class DtlsCoapSession:
|
||||
if exit_reader:
|
||||
return
|
||||
finally:
|
||||
# Reader no longer owns the socket — callers must fail fast.
|
||||
self._reader_running.clear()
|
||||
# Make sure pending waiters don't hang if the reader dies.
|
||||
for tok, (ev, container) in list(self._pending.items()):
|
||||
container.setdefault('err', 'reader exited')
|
||||
container.setdefault('err', SessionClosedError())
|
||||
ev.set()
|
||||
|
||||
def _dispatch_coap(self, datagram):
|
||||
@@ -481,8 +544,7 @@ class DtlsCoapSession:
|
||||
response — Samsung's server keys per-transfer state on the
|
||||
token, and dropping a fresh token on block 1+ silently drops
|
||||
the request."""
|
||||
if self.conn is None:
|
||||
raise ConnectionError("DTLS session closed")
|
||||
self._check_live()
|
||||
tok = self._next_tok()
|
||||
blob = b''
|
||||
num = 0
|
||||
@@ -518,8 +580,7 @@ class DtlsCoapSession:
|
||||
"GET %s /%s block %d: timed out after %d attempt(s)",
|
||||
self.host, '/'.join(path_segs), num, attempt + 1,
|
||||
)
|
||||
raise TimeoutError(
|
||||
f"GET /{'/'.join(path_segs)} block {num} timeout")
|
||||
raise SessionTimeoutError()
|
||||
logger.debug(
|
||||
"GET %s /%s block %d: attempt %d/%d timeout, retrying",
|
||||
self.host, '/'.join(path_segs), num,
|
||||
@@ -528,7 +589,7 @@ class DtlsCoapSession:
|
||||
finally:
|
||||
self._pending.pop(tok, None)
|
||||
if 'err' in container:
|
||||
raise ConnectionError(container['err'])
|
||||
raise container['err']
|
||||
|
||||
code = container['code']
|
||||
payload = container['payload']
|
||||
@@ -552,16 +613,13 @@ class DtlsCoapSession:
|
||||
break
|
||||
num += 1
|
||||
if num > self.MAX_BLOCKS:
|
||||
raise ConnectionError(
|
||||
f"GET /{'/'.join(path_segs)}: >{self.MAX_BLOCKS} "
|
||||
f"blocks, aborting")
|
||||
raise BlockwiseError()
|
||||
return last_code, blob
|
||||
|
||||
def post(self, path_segs, body_cbor, timeout=8.0):
|
||||
"""Single-frame POST with a CBOR-encoded body. Returns
|
||||
(code, payload_bytes). body_cbor must already be encoded."""
|
||||
if self.conn is None:
|
||||
raise ConnectionError("DTLS session closed")
|
||||
self._check_live()
|
||||
tok = self._next_tok()
|
||||
mid = self._next_mid()
|
||||
opts = [(URI_PATH, s.encode()) for s in path_segs]
|
||||
@@ -575,10 +633,9 @@ class DtlsCoapSession:
|
||||
try:
|
||||
self._send_dgram(datagram)
|
||||
if not ev.wait(timeout):
|
||||
raise TimeoutError(
|
||||
f"POST /{'/'.join(path_segs)} timeout")
|
||||
raise SessionTimeoutError()
|
||||
if 'err' in container:
|
||||
raise ConnectionError(container['err'])
|
||||
raise container['err']
|
||||
return container['code'], container['payload']
|
||||
finally:
|
||||
self._pending.pop(tok, None)
|
||||
@@ -595,8 +652,7 @@ class DtlsCoapSession:
|
||||
Real half-open-session detection lives in PollScheduler's
|
||||
`last_success_ts`, surfaced through KeepaliveTask's
|
||||
`liveness_fn`."""
|
||||
if self.conn is None:
|
||||
raise ConnectionError("DTLS session closed")
|
||||
self._check_live()
|
||||
mid = self._next_mid()
|
||||
self._send_dgram(build_coap(TYPE_CON, 0, mid, b'', []))
|
||||
return mid
|
||||
@@ -613,8 +669,7 @@ class DtlsCoapSession:
|
||||
tokens via subscribe. Brief race window where a notify on the
|
||||
old token gets dropped as 'stale' — acceptable for a 6h-scale
|
||||
safety net."""
|
||||
if self.conn is None:
|
||||
raise ConnectionError("DTLS session closed")
|
||||
self._check_live()
|
||||
for tok, href in list(self._observe_tokens.items()):
|
||||
segs = [s for s in href.split('/') if s]
|
||||
try:
|
||||
@@ -637,8 +692,7 @@ class DtlsCoapSession:
|
||||
|
||||
Returns the token used (in case the caller wants to deregister
|
||||
later)."""
|
||||
if self.conn is None:
|
||||
raise ConnectionError("DTLS session closed")
|
||||
self._check_live()
|
||||
tok = self._next_observe_tok()
|
||||
href = '/' + '/'.join(path_segs)
|
||||
# Register the token BEFORE sending — otherwise the device
|
||||
|
||||
@@ -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')
|
||||
@@ -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)
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
import pytest
|
||||
|
||||
from smartthings_local.errors import MalformedMessageError
|
||||
from smartthings_local.protocol.coap import (
|
||||
build_coap, parse_coap, encode_options, block_value, fmt_code,
|
||||
TYPE_CON, METHOD_GET, URI_PATH, ACCEPT, CF_CBOR, BLOCK2,
|
||||
@@ -43,3 +46,13 @@ def test_block_value_promotes_to_two_bytes_when_num_is_large():
|
||||
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': 2.0,
|
||||
})]
|
||||
|
||||
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,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,136 @@
|
||||
"""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,
|
||||
)
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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_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_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,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,75 @@
|
||||
"""Deterministic baseline checks for the current OCF worker stop contract."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
|
||||
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")
|
||||
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