7 Commits
Author SHA1 Message Date
Jack Nagy dec84c9ad8 docs(readme): fill gaps left by the auth/handshake stack merges
- add dtls_handshake.py (from #34) to the repo-layout tree
- point the issue #16 / #20 notes at SamsungServerProfile / ServerCertificateAuth instead of calling that path unsupported
- list the new certificate-profile, connect-deadline, and session-interruption test modules
2026-08-15 11:17:42 +01:00
Quite Yellow e9aee1c235 Merge pull request #33 from Moballo-LLC/codex/py-07-certificate-profiles
Add bound Samsung server certificate profile
2026-08-15 10:20:14 +01:00
Quite Yellow 9ef5598813 Merge pull request #35 from Moballo-LLC/codex/py-08b-session-interruption
Add cancellable DTLS connection attempts
2026-08-15 10:20:05 +01:00
Quite Yellow 917b0e47c5 Merge pull request #34 from Moballo-LLC/codex/py-08a-bounded-connect
Bound DTLS handshakes with a monotonic deadline
2026-08-15 10:19:52 +01:00
Jason Morcos 83a5973434 feat(protocol): add bound server certificate profile 2026-08-14 14:55:10 -07:00
Jason Morcos 3f0e437880 feat(protocol): add cancellable session interruption 2026-08-14 14:37:30 -07:00
Jason Morcos a44930f9df feat(protocol): bound DTLS handshake deadline 2026-08-14 14:26:33 -07:00
10 changed files with 2405 additions and 120 deletions
+99 -3
View File
@@ -48,6 +48,31 @@ sess.subscribe(["operational", "state", "vs", "0"], # OBSERVE
sess.close()
```
`connect()` uses a 12-second monotonic DTLS handshake deadline by default. A
caller that needs a shorter bounded attempt can pass a positive finite value
without changing later reader timeouts. OpenSSL's DTLS timer schedules flight
retransmissions within that same deadline:
```python
sess.connect(timeout=4.0)
```
Connection attempts can also use a one-way cancellation signal. The signal is
backed by a socketpair, so setting it wakes the network wait immediately while
OpenSSL retains control of DTLS retransmission timing:
```python
from smartthings_local.protocol.dtls_session import ConnectCancellation
cancel_connect = ConnectCancellation()
# Another thread may call cancel_connect.set().
sess.connect(timeout=8.0, cancel=cancel_connect)
```
Setting the signal stops subscribed connection attempts and closes their
temporary UDP sockets. It does not alter an already established session or add
new session lifecycle methods. Interrupted attempts raise `SessionClosedError`.
If the cert/key are minted at runtime and never written to disk (e.g. inside
an HA config flow), create the provider from memory instead:
@@ -56,6 +81,76 @@ auth = CertificateAuth.from_memory(cert_pem, key_pem)
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
```
Some newer OCF-PKI devices require an exact Samsung DTLS offer and present a
hardware certificate whose subject contains a certificate UUID. That UUID can
be distinct from the runtime OCF device UUID reported by `/oic/d`, so callers
must obtain and verify the certificate identity independently. When the caller
already has an authorized client certificate and a previously verified
hardware-certificate UUID, opt in to both requirements explicitly:
```python
from smartthings_local.protocol.auth import (
CertificateAuth,
SamsungServerProfile,
)
server_profile = SamsungServerProfile.bound_device(
expected_certificate_uuid,
additional_ca_pem=additional_samsung_ca_pem,
)
auth = CertificateAuth.from_memory(
cert_pem,
key_pem,
server_profile=server_profile,
)
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
```
The default profile is restricted to Samsung home-appliance leaves with
`OU=OCF HA Device`. The profile limits the ClientHello to P-256,
`ECDHE-ECDSA-AES128-GCM-SHA256`, and the observed SHA-256/SHA-1 RSA/ECDSA
signature set, disables session tickets, preserves certificate-chain
verification, and requires the exact subject role
`C=KR, O=Samsung Electronics, OU=OCF HA Device` with a common name ending in
the expected certificate UUID. `additional_ca_pem` is optional and accepts
only a bounded PEM CA-certificate chain; it is applied only to this profiled
context. Without a profile, `CertificateAuth` retains its existing verification
behavior.
Samsung VD-family devices can present the same wire profile with the distinct
`OU=OCF VD Device` role. Select that role explicitly; profiles never fall back
between device classes:
```python
from smartthings_local.protocol.auth import (
SamsungServerProfile,
SamsungServerRole,
ServerCertificateAuth,
)
server_profile = SamsungServerProfile.bound_device(
expected_certificate_uuid,
role=SamsungServerRole.VD_DEVICE,
)
auth = ServerCertificateAuth(server_profile=server_profile)
sess = DtlsCoapSession("192.0.2.100", 5684, auth=auth)
```
`ServerCertificateAuth` is for a server-authenticated channel that does not
send a client certificate, such as the initial DTLS carrier used by
manufacturer-certificate OTM. It still verifies the CA chain, exact selected
subject role, and pinned certificate UUID. It does not learn an identity from
the first endpoint it reaches and cannot be combined with client credentials.
This API deliberately does not discover, mint, authorize, provision, rotate,
or persist credentials, and it performs no ownership transfer or OCF security
resource writes. In particular, the server-only provider can authenticate the
initial manufacturer-certificate channel, but it does not implement the OTM
that follows. The already-owned new-PKI case in
[issue #16](https://github.com/QuiteYellow/SmartThings-Local/issues/16) still
requires an authorized client identity before ordinary protected resources
can be used.
For compatibility, the existing `cert_path` / `key_path` and `cert_pem` /
`key_pem` session arguments remain supported without a deprecation warning.
They are routed through `CertificateAuth` internally. Do not combine `auth`
@@ -155,7 +250,7 @@ For a full worked integration, the higher-level `smartthings_local.ocf` layer (`
Each appliance runs an independent bridge built around three coordinated pieces over one persistent DTLS session: a `StateCache` (single source of truth for all reps), a `PollScheduler` (tiered adaptive polling: hot/warm/cold plus a periodic `/device/0` sweep), and a `KeepaliveTask` (CoAP empty-CON ping for DTLS-layer liveness, with consecutive-failure detection for MQTT availability). Tier cadences are descriptor-declared and were calibrated against the empirically-measured per-firmware ceilings: dryer ~14 req/s, oven ~8 req/s. OBSERVE registrations (RFC 7641) are kept as an opportunistic freshness accelerator: when the appliance has internet and emits notifications, the cache absorbs them and the next-poll timer is reset for that resource; when it's air-gapped, polling alone carries the UX with no other code change. Token-stable Block2 (RFC 7959) handles multi-block reads. Writes are optimistically merged into the cache the moment the device 2.04-confirms, with the scheduler deferring that resource's next poll past the fetchback-revert window. Reconnect with exponential backoff on session errors, gated by a stateless DTLS ClientHello pre-flight (`smartthings_local/protocol/dtls_probe.py`) so a silent/rebooting device or wrong port drops into backoff in ~1 RTT instead of eating the full handshake timeout; when `OCF_PORT` is unset the same probe auto-discovers the live port across the OCF band.
On the currently supported firmware families, authentication uses a client cert keyed to the UUID published in Samsung's own wildcard cloud TLS cert. Their factory ACL grants that UUID `perm=31` (full CRUDN) on `href=*`. That certificate path is not universal: the WD53 profile in issue #16 and the washer in issue #20 reject it and need separate authentication work.
On the currently supported firmware families, authentication uses a client cert keyed to the UUID published in Samsung's own wildcard cloud TLS cert. Their factory ACL grants that UUID `perm=31` (full CRUDN) on `href=*`. That certificate path is not universal: the WD53 profile in issue #16 and the washer in issue #20 reject it. For those newer OCF-PKI devices, the `SamsungServerProfile` and `ServerCertificateAuth` providers (see Quick start) pin and verify the device's hardware certificate, but getting an authorized client credential to reach protected resources is still an open problem.
---
@@ -170,7 +265,7 @@ nmap -Pn -sU -p 5683,5684,49152-49160 "$APPLIANCE_IP"
Read the result:
- **`5684/udp` or a 4915x port with a DTLS first-flight response** → an OCF DTLS listener. Standard-port OCF-PKI firmware may still require an unsupported authentication profile.
- **`5684/udp` or a 4915x port with a DTLS first-flight response** → an OCF DTLS listener. Standard-port OCF-PKI firmware needs the Samsung server-certificate profile (`SamsungServerProfile` / `ServerCertificateAuth`, see Quick start), and no working client credential for it exists yet.
- **`5683/udp` responds to public OCF security/resource GETs** → use `/oic/res` to learn the device's advertised secure endpoint; do not assume that endpoint is fixed.
- **Only `8888/tcp` open (token-based HTTPS)** → older firmware (~2018–2022). **Not supported here.**
@@ -509,6 +604,7 @@ smartthings_local/ The installable library — `pip install sm
coap.py CoAP wire protocol: message encode/decode, token handling
dtls_session.py DTLS session: handshake, client-cert auth (file or in-memory PEM), Block2, liveness
dtls_probe.py Stateless DTLS liveness + opt-in stateful diagnostic
dtls_handshake.py Shared memory-BIO handshake driver, bounded by a monotonic deadline (used by session + probe)
ocf_root_ca.pem Samsung OCF root CA, bundled for handshake verification
ocf/ OCF resource + state layer (reusable)
__init__.py
@@ -535,7 +631,7 @@ mqtt_demo/ MQTT bridge demo (consumes smartthings_loca
.env.example Template — copy to .env, fill in
setup_cert.py One-shot cert minting script (live-fetches AC14K_M + UUID)
pyproject.toml Packaging — PyPI dist `smartthings-local`, hatch-vcs versioning
tests/ pytest suite (CoAP wire, state cache, import isolation, cert loading, DTLS probe, bridge port resolution, cert signing)
tests/ pytest suite (CoAP wire, state cache, import isolation, cert loading, DTLS probe, bridge port resolution, cert signing, certificate profiles, connect deadline, session interruption)
.github/workflows/publish.yml Build + PyPI Trusted Publishing on `v*` tags
```
+334 -9
View File
@@ -2,16 +2,33 @@
from __future__ import annotations
import logging
import re
import warnings
from enum import Enum
from os import PathLike
from pathlib import Path
from typing import Protocol, runtime_checkable
from uuid import UUID
from cryptography.x509.oid import ExtensionOID
from OpenSSL import SSL, _util, crypto
logger = logging.getLogger(__name__)
_OCF_ROOT_CA = str(Path(__file__).with_name("ocf_root_ca.pem"))
_DTLS_CIPHERS = b"ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0"
_DTLS_PSK_CIPHERS = b"ECDHE-PSK-AES128-CBC-SHA256:@SECLEVEL=0"
_SAMSUNG_SERVER_CURVES = b"prime256v1"
_SAMSUNG_SERVER_SIGNATURE_ALGORITHMS = (
b"RSA+SHA256:ECDSA+SHA256:RSA+SHA1:ECDSA+SHA1"
)
_SAMSUNG_SERVER_CN_RE = re.compile(
r"\AOCF Device: [^()\r\n]{1,96} "
r"\((?P<device_identity>[0-9a-f]{8}-[0-9a-f]{4}-"
r"[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12})\)\Z",
re.IGNORECASE,
)
_PSK_CLIENT_CALLBACK_CDEF = (
"unsigned int (*)(SSL *, char *, char *, unsigned int, "
"unsigned char *, unsigned int)"
@@ -53,6 +70,299 @@ class AuthenticationProvider(Protocol):
"""Configure a context while this provider remains session-owned."""
class SamsungServerRole(Enum):
"""Known Samsung OCF hardware-certificate subject roles."""
HOME_APPLIANCE = "OCF HA Device"
VD_DEVICE = "OCF VD Device"
class SamsungServerProfile:
"""Opt-in Samsung hardware-certificate verification profile."""
__slots__ = (
"_additional_ca_certificates",
"_expected_certificate_identity",
"_role",
)
def __init__(
self,
*,
expected_certificate_identity: UUID | str,
role: SamsungServerRole = SamsungServerRole.HOME_APPLIANCE,
additional_ca_pem: str | None = None,
) -> None:
if type(expected_certificate_identity) is UUID:
parsed_identity = expected_certificate_identity
elif type(expected_certificate_identity) is str:
try:
parsed_identity = UUID(expected_certificate_identity)
except ValueError:
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
) from None
if expected_certificate_identity != str(parsed_identity):
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
)
else:
raise TypeError(
"expected_certificate_identity must be a UUID or string"
)
if parsed_identity.int == 0:
raise ValueError(
"expected_certificate_identity must be a canonical "
"non-zero UUID"
)
if type(role) is not SamsungServerRole:
raise TypeError("role must be a SamsungServerRole")
certificates: tuple[bytes, ...] = ()
if additional_ca_pem is not None:
if type(additional_ca_pem) is not str:
raise TypeError("additional_ca_pem must be a string")
try:
raw_ca_pem = additional_ca_pem.encode("ascii")
except UnicodeEncodeError:
raise ValueError(
"additional_ca_pem must contain ASCII PEM certificates"
) from None
parsed_certificates = tuple(_PEM_CERT_RE.findall(raw_ca_pem))
if (
not 1 <= len(parsed_certificates) <= 4
or len(raw_ca_pem) > 32 * 1024
or _PEM_CERT_RE.sub(b"", raw_ca_pem).strip()
):
raise ValueError(
"additional_ca_pem must contain one to four PEM certificates"
)
try:
loaded_certificates = [
crypto.load_certificate(crypto.FILETYPE_PEM, certificate)
for certificate in parsed_certificates
]
basic_constraints = [
[
extension
for extension in certificate.to_cryptography().extensions
if extension.oid == ExtensionOID.BASIC_CONSTRAINTS
]
for certificate in loaded_certificates
]
except (crypto.Error, ValueError):
raise ValueError(
"additional_ca_pem contains an invalid certificate"
) from None
if any(
len(constraints) != 1 or not constraints[0].value.ca
for constraints in basic_constraints
):
raise ValueError(
"additional_ca_pem must contain only CA certificates"
)
fingerprints = {
crypto.dump_certificate(crypto.FILETYPE_ASN1, certificate)
for certificate in loaded_certificates
}
if len(fingerprints) != len(loaded_certificates):
raise ValueError(
"additional_ca_pem must not contain duplicate certificates"
)
certificates = parsed_certificates
object.__setattr__(
self,
"_expected_certificate_identity",
parsed_identity,
)
object.__setattr__(self, "_role", role)
object.__setattr__(self, "_additional_ca_certificates", certificates)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("SamsungServerProfile is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("SamsungServerProfile is immutable")
@classmethod
def bound_device(
cls,
expected_certificate_identity: UUID | str,
*,
role: SamsungServerRole = SamsungServerRole.HOME_APPLIANCE,
additional_ca_pem: str | None = None,
) -> SamsungServerProfile:
"""Bind a verified Samsung hardware leaf to its certificate UUID."""
return cls(
expected_certificate_identity=expected_certificate_identity,
role=role,
additional_ca_pem=additional_ca_pem,
)
def __repr__(self) -> str:
"""Return a representation without device or trust-chain details."""
return "SamsungServerProfile()"
def _configure_context(self, context: SSL.Context) -> None:
curve_setter = getattr(_util.lib, "SSL_CTX_set1_curves_list", None)
if curve_setter is not None:
if curve_setter(context._context, _SAMSUNG_SERVER_CURVES) != 1:
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
)
else:
# pyOpenSSL 23.1 does not expose SSL_CTX_set1_curves_list. Its
# public set_tmp_ecdh fallback produces the same single P-256
# supported-groups ClientHello extension; wire-level tests protect
# that compatibility path. Newer pyOpenSSL uses the exact setter.
try:
with warnings.catch_warnings():
warnings.simplefilter("ignore", DeprecationWarning)
curve = crypto.get_elliptic_curve("prime256v1")
context.set_tmp_ecdh(curve)
except (AttributeError, TypeError, ValueError, SSL.Error):
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
) from None
signature_setter = getattr(
_util.lib,
"SSL_CTX_set1_sigalgs_list",
None,
)
if (
signature_setter is None
or signature_setter(
context._context,
_SAMSUNG_SERVER_SIGNATURE_ALGORITHMS,
)
!= 1
):
raise RuntimeError(
"OpenSSL rejected the Samsung server certificate profile"
)
context.set_options(SSL.OP_NO_TICKET)
if self._additional_ca_certificates:
store = context.get_cert_store()
try:
for certificate in self._additional_ca_certificates:
store.add_cert(
crypto.load_certificate(
crypto.FILETYPE_PEM,
certificate,
)
)
except crypto.Error:
raise RuntimeError(
"OpenSSL rejected the Samsung server trust profile"
) from None
def _verify_peer(
self,
_connection,
certificate,
_error,
depth,
ok,
) -> bool:
if not ok or certificate is None or depth < 0:
return False
if depth > 0:
return True
try:
with warnings.catch_warnings():
# pyOpenSSL deprecates this API in favor of cryptography, but
# reparsing Samsung's non-DER factory leaves with cryptography
# rejects certificates that OpenSSL has already verified.
warnings.simplefilter("ignore", DeprecationWarning)
components = [
(name.decode("ascii"), value.decode("ascii"))
for name, value in certificate.get_subject().get_components()
]
except (
AttributeError,
TypeError,
UnicodeDecodeError,
ValueError,
crypto.Error,
):
logger.warning("Unable to parse Samsung server certificate subject")
return False
common_names = [value for name, value in components if name == "CN"]
organizational_units = [
value for name, value in components if name == "OU"
]
organizations = [value for name, value in components if name == "O"]
countries = [value for name, value in components if name == "C"]
# This deliberately pins the complete Samsung subject role. The OCF
# reference implementation reads only the UUID-bearing CN, but
# relaxing C/O/OU here could accept a different certificate cohort.
if (
len(common_names) != 1
or organizational_units != [self._role.value]
or organizations != ["Samsung Electronics"]
or countries != ["KR"]
):
return False
match = _SAMSUNG_SERVER_CN_RE.fullmatch(common_names[0])
return (
match is not None
and UUID(match.group("device_identity"))
== self._expected_certificate_identity
)
def _configure_certificate_server(
context: SSL.Context,
server_profile: SamsungServerProfile | None,
) -> None:
"""Configure certificate-server verification for one DTLS context."""
context.load_verify_locations(_OCF_ROOT_CA)
if server_profile is None:
context.set_verify(SSL.VERIFY_PEER, _verify_peer)
else:
server_profile._configure_context(context)
context.set_verify(
SSL.VERIFY_PEER,
server_profile._verify_peer,
)
# @SECLEVEL=0 permits SHA-1 in Samsung's server cert chain (AC14K_M
# intermediate is SHA-1 signed). This is the only channel that reaches
# the OpenSSL instance cryptography bundles; ctypes and cffi bindings
# do not expose SSL_CTX_set_security_level on this build.
context.set_cipher_list(_DTLS_CIPHERS)
class ServerCertificateAuth:
"""Verify a pinned Samsung server without a client certificate."""
__slots__ = ("_server_profile",)
def __init__(self, *, server_profile: SamsungServerProfile) -> None:
if type(server_profile) is not SamsungServerProfile:
raise TypeError("server_profile must be a SamsungServerProfile")
object.__setattr__(self, "_server_profile", server_profile)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("ServerCertificateAuth is immutable")
def __delattr__(self, _name: str) -> None:
raise AttributeError("ServerCertificateAuth is immutable")
def __repr__(self) -> str:
"""Return a representation without server identity or trust details."""
return "ServerCertificateAuth()"
def configure_context(self, context: SSL.Context) -> None:
"""Verify the selected server profile without loading client material."""
_configure_certificate_server(context, self._server_profile)
class CertificateAuth:
"""Certificate authentication loaded from files or in-memory PEM data.
@@ -65,6 +375,7 @@ class CertificateAuth:
"_certificate_pem",
"_private_key_path",
"_private_key_pem",
"_server_profile",
)
def __init__(
@@ -74,6 +385,7 @@ class CertificateAuth:
private_key_path: str | PathLike[str] | None = None,
certificate_pem: str | None = None,
private_key_pem: str | None = None,
server_profile: SamsungServerProfile | None = None,
) -> None:
file_supplied = (
certificate_path is not None or private_key_path is not None
@@ -101,6 +413,11 @@ class CertificateAuth:
"must pass either certificate_path/private_key_path or "
"certificate_pem/private_key_pem"
)
if (
server_profile is not None
and type(server_profile) is not SamsungServerProfile
):
raise TypeError("server_profile must be a SamsungServerProfile")
object.__setattr__(
self,
"_certificate_path",
@@ -113,6 +430,7 @@ class CertificateAuth:
)
object.__setattr__(self, "_certificate_pem", certificate_pem)
object.__setattr__(self, "_private_key_pem", private_key_pem)
object.__setattr__(self, "_server_profile", server_profile)
def __setattr__(self, _name: str, _value: object) -> None:
raise AttributeError("CertificateAuth is immutable")
@@ -125,11 +443,14 @@ class CertificateAuth:
cls,
certificate_path: str | PathLike[str],
private_key_path: str | PathLike[str],
*,
server_profile: SamsungServerProfile | None = None,
) -> CertificateAuth:
"""Create a provider backed by certificate-chain and key files."""
return cls(
certificate_path=certificate_path,
private_key_path=private_key_path,
server_profile=server_profile,
)
@classmethod
@@ -137,11 +458,14 @@ class CertificateAuth:
cls,
certificate_pem: str,
private_key_pem: str,
*,
server_profile: SamsungServerProfile | None = None,
) -> CertificateAuth:
"""Create a provider backed by an in-memory PEM chain and key."""
return cls(
certificate_pem=certificate_pem,
private_key_pem=private_key_pem,
server_profile=server_profile,
)
def __repr__(self) -> str:
@@ -150,13 +474,7 @@ class 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)
_configure_certificate_server(context, self._server_profile)
if self._certificate_pem is not None:
_load_pem_chain(
context,
@@ -240,7 +558,14 @@ class PskAuth:
"the installed OpenSSL binding does not support DTLS PSK"
)
context.set_cipher_list(_DTLS_PSK_CIPHERS)
setter(context._context, self._callback) # noqa: SLF001
setter(context._context, self._callback)
__all__ = ["AuthenticationProvider", "CertificateAuth", "PskAuth"]
__all__ = [
"AuthenticationProvider",
"CertificateAuth",
"PskAuth",
"SamsungServerProfile",
"SamsungServerRole",
"ServerCertificateAuth",
]
@@ -0,0 +1,97 @@
"""Shared memory-BIO driver for bounded DTLS handshakes."""
from __future__ import annotations
import select
import time
from collections.abc import Callable
from OpenSSL import SSL
from .coap import split_dtls
_HANDSHAKE_POLL_S = 0.5
_MAX_DATAGRAM_SIZE = 65535
class _HandshakeCancelled(Exception):
"""Internal signal that a handshake wake socket became readable."""
def _drive_dtls_handshake(
connection,
sock,
*,
deadline: float,
retries: int | None = None,
on_datagram: Callable[[bytes], None] | None = None,
wake_socket=None,
) -> bool:
"""Drive one memory-BIO DTLS handshake up to a monotonic deadline.
OpenSSL owns the retransmission schedule. ``retries`` optionally limits
how many expired retransmission timers are serviced; a normal session is
bounded only by its deadline, while the diagnostic probe retains its
explicit retry budget.
Return ``True`` only when the handshake completes before the deadline.
TLS and socket failures are left to the caller to classify.
"""
retransmits = 0
while time.monotonic() < deadline:
try:
connection.do_handshake()
return time.monotonic() < deadline
except SSL.WantReadError:
pass
try:
output = connection.bio_read(_MAX_DATAGRAM_SIZE)
except SSL.WantReadError:
output = None
if output:
for record in split_dtls(output):
if sock.send(record) != len(record):
raise OSError("incomplete UDP send")
remaining = deadline - time.monotonic()
if remaining <= 0:
break
wait = min(_HANDSHAKE_POLL_S, remaining)
timed_out = False
if wake_socket is None:
sock.settimeout(wait)
try:
datagram = sock.recv(_MAX_DATAGRAM_SIZE)
except TimeoutError:
timed_out = True
else:
readable, _, _ = select.select(
(sock, wake_socket),
(),
(),
wait,
)
if wake_socket in readable:
raise _HandshakeCancelled()
if sock not in readable:
timed_out = True
else:
datagram = sock.recv(_MAX_DATAGRAM_SIZE)
if timed_out:
timer = connection.DTLSv1_get_timeout()
if timer is not None and timer <= 0:
if retries is not None and retransmits >= retries:
break
connection.DTLSv1_handle_timeout()
retransmits += 1
continue
if not datagram:
continue
if on_datagram is not None:
on_datagram(datagram)
connection.bio_write(datagram)
return False
+33 -65
View File
@@ -39,8 +39,9 @@ from dataclasses import dataclass
from OpenSSL import SSL
from ..errors import ProbeError
from .auth import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
from .coap import split_dtls
from .dtls_session import _DTLS_CIPHERS, _OCF_ROOT_CA, _load_pem_chain
from .dtls_handshake import _drive_dtls_handshake
from .endpoint import open_connected_udp_socket
# DTLS record content types (RFC 6347 §4.1)
@@ -603,73 +604,40 @@ def diagnose_dtls_handshake(
started = time.monotonic()
deadline = started + timeout
seen = set()
retransmits = 0
def record_datagram(datagram):
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
result.datagrams.append(datagram)
for content_type, detail in classify_datagram(datagram):
if content_type == _CT_HANDSHAKE:
if detail not in seen:
seen.add(detail)
result.handshake_msgs.append(detail)
if result.outcome == DEAD:
result.outcome = LIVE
elif content_type == _CT_ALERT and detail is not None:
level, name = detail
result.alert = (level, name)
if level == 2: # fatal
result.outcome = REJECTED
try:
while time.monotonic() < deadline:
try:
conn.do_handshake()
result.outcome = COMPLETED
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
break
except SSL.WantReadError:
pass
except SSL.Error:
# A fatal Alert lands here; the alert record was already
# captured below, so classification still works.
result.error = ProbeError()
break
try:
o = conn.bio_read(65535)
if o:
for r in split_dtls(o):
if sock.send(r) != len(r):
raise OSError('short UDP send')
except SSL.WantReadError:
pass
remaining = deadline - time.monotonic()
if remaining <= 0:
break
sock.settimeout(min(0.5, remaining))
try:
d = sock.recv(65535)
except TimeoutError:
# No answer to the last flight. Service OpenSSL's DTLS
# retransmit timer: once it has counted down to 0,
# handle_timeout() re-queues the previous flight into the
# write BIO for the next iteration to flush. A live server
# answers within a flight or two; a silent/non-DTLS port
# never does, so we give up only after `retries`
# retransmits — one dropped ClientHello no longer reads as
# a false DEAD.
to = conn.DTLSv1_get_timeout()
if to is not None and to <= 0:
if retransmits >= retries:
break
conn.DTLSv1_handle_timeout()
retransmits += 1
continue
if not d:
continue
completed = _drive_dtls_handshake(
conn,
sock,
deadline=deadline,
retries=retries,
on_datagram=record_datagram,
)
if completed:
result.outcome = COMPLETED
if result.rtt_s is None:
result.rtt_s = time.monotonic() - started
result.datagrams.append(d)
for ct, detail in classify_datagram(d):
if ct == _CT_HANDSHAKE:
if detail not in seen:
seen.add(detail)
result.handshake_msgs.append(detail)
if result.outcome == DEAD:
result.outcome = LIVE
elif ct == _CT_ALERT and detail is not None:
level, name = detail
result.alert = (level, name)
if level == 2: # fatal
result.outcome = REJECTED
conn.bio_write(d)
except SSL.Error:
# A fatal Alert lands here; record_datagram() has already classified
# the alert record before it is fed back into OpenSSL.
result.error = ProbeError()
except OSError:
result.error = ProbeError()
finally:
+150 -41
View File
@@ -19,6 +19,7 @@ on a per-token Event the reader signals. OBSERVE notifications are
delivered via the on_notification callback.
"""
import errno
import math
import os
import socket
import threading
@@ -48,6 +49,11 @@ from .auth import (
_OCF_ROOT_CA,
_load_pem_chain,
)
from .dtls_handshake import (
_HANDSHAKE_POLL_S,
_HandshakeCancelled,
_drive_dtls_handshake,
)
from .endpoint import open_connected_udp_socket
import logging
@@ -88,6 +94,75 @@ _ADVISORY_ERRNOS = frozenset(
) if value is not None
)
def _validate_handshake_timeout(timeout, default):
"""Return one finite, positive DTLS handshake timeout."""
value = default if timeout is None else timeout
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise TypeError('timeout must be a number or None')
try:
value = float(value)
except OverflowError:
raise ValueError(
'timeout must be a positive finite number or None') from None
if not math.isfinite(value) or value <= 0:
raise ValueError('timeout must be a positive finite number or None')
return value
class ConnectCancellation:
"""One-way, socket-backed cancellation signal for ``connect()``.
Each active connection attempt receives its own wake socket. ``set()``
makes every subscribed socket readable immediately, without a polling
thread or a session-level abort API.
"""
__slots__ = ("_is_set", "_lock", "_writers")
def __init__(self) -> None:
self._is_set = False
self._lock = threading.Lock()
self._writers: set[socket.socket] = set()
def set(self) -> None:
"""Cancel current and future connection attempts using this signal."""
with self._lock:
if self._is_set:
return
self._is_set = True
for writer in self._writers:
try:
writer.send(b"\0")
except OSError:
pass
def is_set(self) -> bool:
"""Return whether cancellation has been requested."""
with self._lock:
return self._is_set
def _subscribe(self) -> tuple[socket.socket, socket.socket]:
reader, writer = socket.socketpair()
reader.setblocking(False)
with self._lock:
self._writers.add(writer)
if self._is_set:
writer.send(b"\0")
return reader, writer
def _unsubscribe(
self,
reader: socket.socket,
writer: socket.socket,
) -> bool:
with self._lock:
self._writers.discard(writer)
interrupted = self._is_set
reader.close()
writer.close()
return interrupted
class DtlsCoapSession:
"""Single sustained DTLS-CoAP session.
@@ -199,69 +274,103 @@ class DtlsCoapSession:
# ---- lifecycle ---------------------------------------------------
def connect(self):
"""DTLS handshake. Blocks up to HANDSHAKE_TIMEOUT_S. Raises
ConnectionError / TimeoutError on failure."""
def connect(
self,
*,
timeout: float | None = None,
cancel: ConnectCancellation | None = None,
):
"""Perform a cancellable DTLS handshake within a monotonic deadline.
``timeout`` overrides ``HANDSHAKE_TIMEOUT_S`` for this call. OpenSSL
owns DTLS retransmission timing while every receive is capped by the
remaining budget, so wall-clock adjustments cannot change the bound.
A ``ConnectCancellation`` wakes the network wait immediately and does
not alter an already established session.
"""
handshake_timeout = _validate_handshake_timeout(
timeout, self.HANDSHAKE_TIMEOUT_S)
if cancel is not None and not isinstance(cancel, ConnectCancellation):
raise TypeError("cancel must be a ConnectCancellation or None")
if cancel is not None and cancel.is_set():
raise SessionClosedError()
deadline = time.monotonic() + handshake_timeout
ctx = SSL.Context(SSL.DTLS_METHOD)
self.auth.configure_context(ctx)
if cancel is not None and cancel.is_set():
raise SessionClosedError()
conn = SSL.Connection(ctx, None)
conn.set_connect_state()
conn.set_ciphertext_mtu(self.mtu)
if cancel is not None and cancel.is_set():
raise SessionClosedError()
remaining = deadline - time.monotonic()
if remaining <= 0:
raise SessionTimeoutError()
sock, endpoint = open_connected_udp_socket(
self.host,
self.port,
family=self.family,
local_port=self.local_port,
timeout=2.0,
timeout=min(_HANDSHAKE_POLL_S, remaining),
)
dest = endpoint.sockaddr
if cancel is not None and cancel.is_set():
sock.close()
raise SessionClosedError()
wake_subscription = None
subscription_failed = False
if cancel is not None:
try:
wake_subscription = cancel._subscribe()
except OSError:
subscription_failed = True
if subscription_failed:
sock.close()
raise SessionError() from OSError(
"connection cancellation setup failed"
)
t0 = time.time()
backend_failed = False
while time.time() - t0 < self.HANDSHAKE_TIMEOUT_S:
io_failed = False
cancelled = False
interrupted = False
try:
try:
conn.do_handshake()
break
except SSL.WantReadError:
pass
completed = _drive_dtls_handshake(
conn,
sock,
deadline=deadline,
wake_socket=(
wake_subscription[0]
if wake_subscription is not None
else None
),
)
except _HandshakeCancelled:
cancelled = True
except SSL.Error:
sock.close()
backend_failed = True
break
send_failed = False
try:
o = conn.bio_read(65535)
if o:
for r in _split_dtls(o):
if sock.send(r) != len(r):
raise OSError('incomplete UDP send')
except SSL.WantReadError:
pass
except OSError:
sock.close()
send_failed = True
if send_failed:
raise EndpointError() from OSError('UDP send failed')
receive_failed = False
try:
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:
io_failed = True
finally:
if wake_subscription is not None:
interrupted = cancel._unsubscribe(*wake_subscription)
if cancelled or interrupted:
sock.close()
raise SessionClosedError()
if backend_failed:
sock.close()
raise SessionError() from ConnectionError('DTLS backend failed')
if io_failed:
sock.close()
raise EndpointError() from OSError('UDP handshake I/O failed')
if not completed:
sock.close()
raise SessionTimeoutError()
if backend_failed:
raise SessionError() from ConnectionError('DTLS backend failed')
self.sock = sock
self.conn = conn
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -376,7 +376,7 @@ def test_session_uses_connected_socket_send_and_recv(monkeypatch):
assert open_calls == [(('device.example', 5684), {
'family': socket.AF_INET6,
'local_port': None,
'timeout': 2.0,
'timeout': 0.5,
})]
session.close()
+62 -1
View File
@@ -10,8 +10,14 @@ from smartthings_local.protocol.auth import (
AuthenticationProvider,
CertificateAuth,
PskAuth,
SamsungServerProfile,
SamsungServerRole,
ServerCertificateAuth,
)
from smartthings_local.protocol.dtls_session import (
ConnectCancellation,
DtlsCoapSession,
)
from smartthings_local.protocol.dtls_session import DtlsCoapSession
def _assert_compatible_signature(callable_object, expected: list[str]) -> None:
@@ -54,6 +60,47 @@ def test_certificate_auth_is_a_public_authentication_provider():
provider = CertificateAuth.from_files("/synthetic/cert.pem", "/synthetic/key")
assert isinstance(provider, AuthenticationProvider)
for factory in (CertificateAuth.from_files, CertificateAuth.from_memory):
profile_parameter = inspect.signature(factory).parameters["server_profile"]
assert profile_parameter.kind is inspect.Parameter.KEYWORD_ONLY
assert profile_parameter.default is None
def test_samsung_server_profile_is_public_and_explicitly_bound():
parameters = inspect.signature(SamsungServerProfile.bound_device).parameters
assert list(parameters) == [
"expected_certificate_identity",
"role",
"additional_ca_pem",
]
assert (
parameters["expected_certificate_identity"].default
is inspect.Parameter.empty
)
assert parameters["role"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["role"].default is SamsungServerRole.HOME_APPLIANCE
assert parameters["additional_ca_pem"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["additional_ca_pem"].default is None
def test_server_certificate_auth_is_a_public_authentication_provider():
profile = SamsungServerProfile.bound_device(
"abababab-abab-abab-abab-abababababab",
role=SamsungServerRole.VD_DEVICE,
)
provider = ServerCertificateAuth(server_profile=profile)
assert isinstance(provider, AuthenticationProvider)
session = DtlsCoapSession("device.example", 5684, auth=provider)
assert session.auth is provider
assert session.cert_path is None
assert session.key_path is None
assert session.cert_pem is None
assert session.key_pem is None
parameters = inspect.signature(ServerCertificateAuth).parameters
assert list(parameters) == ["server_profile"]
assert parameters["server_profile"].kind is inspect.Parameter.KEYWORD_ONLY
assert parameters["server_profile"].default is inspect.Parameter.empty
def test_psk_auth_is_a_public_authentication_provider():
provider = PskAuth(identity=b"i" * 16, key=b"k" * 16)
@@ -81,6 +128,20 @@ def test_dtls_session_keeps_current_consumer_methods():
"subscribe",
}
assert expected <= set(dir(DtlsCoapSession))
assert "abort" not in DtlsCoapSession.__dict__
assert "quiesce_for_close" not in DtlsCoapSession.__dict__
_assert_compatible_signature(DtlsCoapSession.connect, ["self"])
connect_timeout = inspect.signature(DtlsCoapSession.connect).parameters[
"timeout"
]
assert connect_timeout.kind is inspect.Parameter.KEYWORD_ONLY
assert connect_timeout.default is None
connect_cancel = inspect.signature(DtlsCoapSession.connect).parameters[
"cancel"
]
assert connect_cancel.kind is inspect.Parameter.KEYWORD_ONLY
assert connect_cancel.default is None
assert callable(ConnectCancellation().set)
_assert_compatible_signature(
DtlsCoapSession.get,
[
+344
View File
@@ -0,0 +1,344 @@
"""Deterministic tests for bounded DTLS handshake timing."""
from __future__ import annotations
import socket
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionTimeoutError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.dtls_session import DtlsCoapSession
from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint
class _Clock:
def __init__(self):
self.now = 100.0
def monotonic(self):
return self.now
def advance(self, seconds):
self.now += seconds
class _Auth:
def __init__(self, clock=None, configure_delay=0.0):
self.clock = clock
self.configure_delay = configure_delay
def configure_context(self, _context):
if self.clock is not None:
self.clock.advance(self.configure_delay)
class _Connection:
def __init__(self, outcomes=None, outputs=None, timer=None):
self.outcomes = list(outcomes or ())
self.outputs = list(outputs or ())
self.timer = timer
self.bio_writes = []
self.timeout_calls = 0
def set_connect_state(self):
return None
def set_ciphertext_mtu(self, _mtu):
return None
def do_handshake(self):
outcome = self.outcomes.pop(0) if self.outcomes else "want-read"
if outcome == "want-read":
raise SSL.WantReadError()
if isinstance(outcome, Exception):
raise outcome
def bio_read(self, _size):
if self.outputs:
return self.outputs.pop(0)
raise SSL.WantReadError()
def bio_write(self, data):
self.bio_writes.append(data)
def DTLSv1_get_timeout(self):
return self.timer
def DTLSv1_handle_timeout(self):
self.timeout_calls += 1
class _Socket:
def __init__(self, clock, inbound=()):
self.clock = clock
self.inbound = list(inbound)
self.timeouts = []
self.sent = []
self.closed = False
def settimeout(self, timeout):
self.timeouts.append(timeout)
def send(self, data):
self.sent.append(data)
return len(data)
def recv(self, _size):
if self.inbound:
result = self.inbound.pop(0)
if isinstance(result, Exception):
self.clock.advance(self.timeouts[-1])
raise result
return result
self.clock.advance(self.timeouts[-1])
raise TimeoutError()
def close(self):
self.closed = True
def _session(auth=None):
return DtlsCoapSession(
"device.example",
5684,
auth=auth or _Auth(),
)
def _install_handshake(
monkeypatch,
clock,
*,
outcomes=(),
outputs=(),
inbound=(),
timer=None,
):
connection = _Connection(outcomes, outputs, timer)
sock = _Socket(clock, inbound)
endpoint = ResolvedUdpEndpoint(
socket.AF_INET,
("192.0.2.10", 5684),
)
open_calls = []
def open_socket(*args, **kwargs):
open_calls.append((args, kwargs))
sock.settimeout(kwargs["timeout"])
return sock, endpoint
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(
dtls_session.SSL,
"Connection",
lambda *_args: connection,
)
monkeypatch.setattr(
dtls_session,
"open_connected_udp_socket",
open_socket,
)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
monkeypatch.setattr(dtls_session.time, "sleep", clock.advance)
monkeypatch.setattr(
dtls_session.time,
"time",
lambda: pytest.fail("wall clock must not control handshake deadlines"),
)
return connection, sock, endpoint, open_calls
@pytest.mark.parametrize("timeout", (True, "1", object()))
def test_connect_timeout_type_is_explicit(timeout):
with pytest.raises(TypeError, match="number or None"):
_session().connect(timeout=timeout)
@pytest.mark.parametrize(
"timeout",
(
0,
-1,
float("nan"),
float("inf"),
float("-inf"),
10**1000,
),
)
def test_connect_timeout_must_be_positive_and_finite(timeout):
with pytest.raises(ValueError, match="positive finite"):
_session().connect(timeout=timeout)
def test_connect_timeout_caps_every_blocking_poll(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, open_calls = _install_handshake(
monkeypatch,
clock,
)
session = _session()
with pytest.raises(SessionTimeoutError):
session.connect(timeout=4.75)
assert clock.now == pytest.approx(104.75)
assert sock.closed
assert open_calls == [
(
("device.example", 5684),
{
"family": socket.AF_UNSPEC,
"local_port": None,
"timeout": 0.5,
},
)
]
assert max(sock.timeouts) <= 0.5
assert sock.timeouts[-1] == pytest.approx(0.25)
def test_short_timeout_is_not_rounded_up_to_poll_interval(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, open_calls = _install_handshake(
monkeypatch,
clock,
)
with pytest.raises(SessionTimeoutError):
_session().connect(timeout=0.125)
assert clock.now == pytest.approx(100.125)
assert open_calls[0][1]["timeout"] == pytest.approx(0.125)
assert sock.timeouts == pytest.approx([0.125, 0.125])
def test_default_timeout_uses_session_constant(monkeypatch):
clock = _Clock()
_connection, sock, _endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
)
session = _session()
session.HANDSHAKE_TIMEOUT_S = 0.2
with pytest.raises(SessionTimeoutError):
session.connect()
assert clock.now == pytest.approx(100.2)
assert sock.closed
def test_context_setup_consumes_the_same_deadline(monkeypatch):
clock = _Clock()
socket_opened = False
def open_socket(*_args, **_kwargs):
nonlocal socket_opened
socket_opened = True
raise AssertionError("expired setup must not open a socket")
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(dtls_session.SSL, "Connection", lambda *_args: _Connection())
monkeypatch.setattr(dtls_session, "open_connected_udp_socket", open_socket)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
session = _session(_Auth(clock, configure_delay=0.2))
with pytest.raises(SessionTimeoutError):
session.connect(timeout=0.1)
assert not socket_opened
def test_socket_setup_consumes_the_same_deadline(monkeypatch):
clock = _Clock()
connection = _Connection()
connection.do_handshake = lambda: pytest.fail(
"expired socket setup must not start a handshake"
)
sock = _Socket(clock)
endpoint = ResolvedUdpEndpoint(
socket.AF_INET,
("192.0.2.10", 5684),
)
def open_socket(*_args, **kwargs):
sock.settimeout(kwargs["timeout"])
clock.advance(0.2)
return sock, endpoint
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(
dtls_session.SSL,
"Connection",
lambda *_args: connection,
)
monkeypatch.setattr(dtls_session, "open_connected_udp_socket", open_socket)
monkeypatch.setattr(dtls_session.time, "monotonic", clock.monotonic)
with pytest.raises(SessionTimeoutError):
_session().connect(timeout=0.1)
assert sock.closed
def test_successful_handshake_preserves_connected_session_state(monkeypatch):
clock = _Clock()
connection, sock, endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
outcomes=("want-read", "success"),
inbound=(b"synthetic server flight",),
)
session = _session()
session.connect(timeout=1.0)
assert connection.bio_writes == [b"synthetic server flight"]
assert session.conn is connection
assert session.sock is sock
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert not sock.closed
def test_connect_services_openssl_retransmit_timer(monkeypatch):
clock = _Clock()
outbound = b"\x16\xfe\xfd" + b"\x00" * 8 + b"\x00\x01x"
connection, sock, _endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
outcomes=("want-read", "want-read", "success"),
outputs=(outbound, outbound),
inbound=(TimeoutError(), b"synthetic server flight"),
timer=0.0,
)
_session().connect(timeout=2.0)
assert connection.timeout_calls == 1
assert sock.sent == [outbound, outbound]
assert connection.bio_writes == [b"synthetic server flight"]
@pytest.mark.parametrize("success_delay", (0.1, 0.2))
def test_handshake_success_at_or_after_deadline_is_rejected(
monkeypatch,
success_delay,
):
clock = _Clock()
connection, sock, _endpoint, _open_calls = _install_handshake(
monkeypatch,
clock,
)
def late_success():
clock.advance(success_delay)
connection.do_handshake = late_success
with pytest.raises(SessionTimeoutError):
_session().connect(timeout=0.1)
assert sock.closed
+280
View File
@@ -0,0 +1,280 @@
"""Deterministic tests for connection-attempt cancellation."""
from __future__ import annotations
import select
import socket
import threading
import time
import traceback
import pytest
from OpenSSL import SSL
from smartthings_local.errors import SessionClosedError, SessionError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.dtls_session import (
ConnectCancellation,
DtlsCoapSession,
)
from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint
class _Auth:
def __init__(self, on_configure=None):
self.on_configure = on_configure
def configure_context(self, _context):
if self.on_configure is not None:
self.on_configure()
class _Connection:
def __init__(self, *, started=None, on_success=None, succeed=False):
self.started = started
self.on_success = on_success
self.succeed = succeed
self.bio_writes = []
self.handshake_calls = 0
def set_connect_state(self):
return None
def set_ciphertext_mtu(self, _mtu):
return None
def do_handshake(self):
self.handshake_calls += 1
if self.started is not None:
self.started.set()
if self.succeed:
if self.on_success is not None:
self.on_success()
return
raise SSL.WantReadError()
def bio_read(self, _size):
raise SSL.WantReadError()
def bio_write(self, datagram):
self.bio_writes.append(datagram)
self.succeed = True
def DTLSv1_get_timeout(self):
return None
def shutdown(self):
return None
def _session(auth=None):
return DtlsCoapSession(
"device.example",
5684,
auth=auth or _Auth(),
)
def _install_connection(monkeypatch, connection, data_socket, *, on_open=None):
endpoint = ResolvedUdpEndpoint(
socket.AF_INET,
("192.0.2.10", 5684),
)
def open_socket(*_args, **_kwargs):
if on_open is not None:
on_open()
return data_socket, endpoint
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(
dtls_session.SSL,
"Connection",
lambda *_args: connection,
)
monkeypatch.setattr(
dtls_session,
"open_connected_udp_socket",
open_socket,
)
return endpoint
def _run_connect(session, cancel):
outcome = {}
def worker():
try:
session.connect(timeout=2.0, cancel=cancel)
except Exception as error: # noqa: BLE001 - captured for assertion
outcome["error"] = error
else:
outcome["connected"] = True
thread = threading.Thread(target=worker)
thread.start()
return thread, outcome
@pytest.mark.parametrize("cancel", (True, threading.Event(), object(), "signal"))
def test_connect_cancel_type_is_explicit(cancel):
with pytest.raises(TypeError, match="ConnectCancellation or None"):
_session().connect(cancel=cancel)
def test_pre_cancelled_connect_stops_before_context_setup(monkeypatch):
cancel = ConnectCancellation()
cancel.set()
monkeypatch.setattr(
dtls_session.SSL,
"Context",
lambda *_args: pytest.fail("cancelled connect configured TLS"),
)
with pytest.raises(SessionClosedError):
_session().connect(cancel=cancel)
def test_cancel_during_context_setup_stops_before_socket_setup(monkeypatch):
cancel = ConnectCancellation()
monkeypatch.setattr(dtls_session.SSL, "Context", lambda *_args: object())
monkeypatch.setattr(
dtls_session,
"open_connected_udp_socket",
lambda *_args, **_kwargs: pytest.fail(
"cancelled connect opened a socket"
),
)
with pytest.raises(SessionClosedError):
_session(_Auth(cancel.set)).connect(cancel=cancel)
def test_cancel_during_socket_setup_closes_before_handshake(monkeypatch):
cancel = ConnectCancellation()
connection = _Connection(succeed=True)
data_socket, peer = socket.socketpair()
_install_connection(
monkeypatch,
connection,
data_socket,
on_open=cancel.set,
)
try:
with pytest.raises(SessionClosedError):
_session().connect(cancel=cancel)
assert data_socket.fileno() == -1
assert connection.handshake_calls == 0
finally:
peer.close()
def test_socket_signal_wakes_every_subscribed_waiter():
cancel = ConnectCancellation()
first = cancel._subscribe()
second = cancel._subscribe()
try:
cancel.set()
readable, _, _ = select.select(
(first[0], second[0]),
(),
(),
0,
)
assert set(readable) == {first[0], second[0]}
assert cancel._unsubscribe(*first)
assert cancel._unsubscribe(*second)
finally:
for reader, writer in (first, second):
reader.close()
writer.close()
def test_cancel_wakes_blocked_connect_without_poll_latency(monkeypatch):
cancel = ConnectCancellation()
started = threading.Event()
connection = _Connection(started=started)
data_socket, peer = socket.socketpair()
_install_connection(monkeypatch, connection, data_socket)
session = _session()
thread, outcome = _run_connect(session, cancel)
try:
assert started.wait(1.0)
before = time.monotonic()
cancel.set()
thread.join(1.0)
elapsed = time.monotonic() - before
assert not thread.is_alive()
assert elapsed < 0.25
assert isinstance(outcome.get("error"), SessionClosedError)
assert data_socket.fileno() == -1
assert session.sock is None
assert session.conn is None
assert not cancel._writers
finally:
cancel.set()
thread.join(1.0)
peer.close()
def test_cancel_wins_race_with_reported_handshake_success(monkeypatch):
cancel = ConnectCancellation()
connection = _Connection(on_success=cancel.set, succeed=True)
data_socket, peer = socket.socketpair()
_install_connection(monkeypatch, connection, data_socket)
session = _session()
try:
with pytest.raises(SessionClosedError):
session.connect(cancel=cancel)
assert data_socket.fileno() == -1
assert session.sock is None
assert session.conn is None
finally:
peer.close()
def test_successful_connect_does_not_set_cancel_or_close_session(monkeypatch):
cancel = ConnectCancellation()
connection = _Connection()
data_socket, peer = socket.socketpair()
endpoint = _install_connection(monkeypatch, connection, data_socket)
peer.send(b"synthetic server flight")
session = _session()
try:
session.connect(cancel=cancel)
assert not cancel.is_set()
assert session.sock is data_socket
assert session.conn is connection
assert session.endpoint is endpoint
assert connection.bio_writes == [b"synthetic server flight"]
assert not cancel._writers
finally:
session.close()
peer.close()
def test_cancellation_socket_failure_is_redacted_and_closes_udp(monkeypatch):
class FailingCancellation(ConnectCancellation):
def _subscribe(self):
raise OSError("credential-value at device.example")
connection = _Connection()
data_socket, peer = socket.socketpair()
_install_connection(monkeypatch, connection, data_socket)
try:
with pytest.raises(SessionError) as exc:
_session().connect(cancel=FailingCancellation())
formatted = "".join(traceback.format_exception(exc.value))
assert data_socket.fileno() == -1
assert exc.value.__context__ is None
assert "credential-value" not in formatted
assert "device.example" not in formatted
finally:
peer.close()