Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7a74a955f3 | ||
|
|
d4aebebca4 | ||
|
|
7f0f1e5531 | ||
|
|
2a6fc627f2 | ||
|
|
8fb37ca2ed | ||
|
|
74d78543e1 | ||
|
|
da8f917541 | ||
|
|
a63bc0bb01 |
@@ -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.
|
||||||
@@ -27,12 +27,16 @@ session directly:
|
|||||||
|
|
||||||
```python
|
```python
|
||||||
import cbor2
|
import cbor2
|
||||||
|
from smartthings_local.protocol.auth import CertificateAuth
|
||||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
||||||
|
|
||||||
|
auth = CertificateAuth.from_files(
|
||||||
|
"certs/client_fullchain.pem",
|
||||||
|
"certs/client.key",
|
||||||
|
)
|
||||||
sess = DtlsCoapSession(
|
sess = DtlsCoapSession(
|
||||||
"192.0.2.100", 49154,
|
"192.0.2.100", 49154,
|
||||||
cert_path="certs/client_fullchain.pem",
|
auth=auth,
|
||||||
key_path="certs/client.key",
|
|
||||||
)
|
)
|
||||||
sess.connect()
|
sess.connect()
|
||||||
sess.start_reader()
|
sess.start_reader()
|
||||||
@@ -44,12 +48,34 @@ sess.subscribe(["operational", "state", "vs", "0"], # OBSERVE
|
|||||||
sess.close()
|
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
|
```python
|
||||||
sess = DtlsCoapSession("192.0.2.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
|
### Classified errors
|
||||||
|
|
||||||
Runtime transport failures use the public types in
|
Runtime transport failures use the public types in
|
||||||
@@ -479,6 +505,7 @@ smartthings_local/ The installable library — `pip install sm
|
|||||||
__init__.py
|
__init__.py
|
||||||
protocol/ DTLS-CoAP transport (reusable by any consumer, not just MQTT)
|
protocol/ DTLS-CoAP transport (reusable by any consumer, not just MQTT)
|
||||||
__init__.py
|
__init__.py
|
||||||
|
auth.py Immutable DTLS authentication providers
|
||||||
coap.py CoAP wire protocol: message encode/decode, token handling
|
coap.py CoAP wire protocol: message encode/decode, token handling
|
||||||
dtls_session.py DTLS session: handshake, client-cert auth (file or in-memory PEM), Block2, liveness
|
dtls_session.py DTLS session: handshake, client-cert auth (file or in-memory PEM), Block2, liveness
|
||||||
dtls_probe.py Stateless DTLS liveness + opt-in stateful diagnostic
|
dtls_probe.py Stateless DTLS liveness + opt-in stateful diagnostic
|
||||||
@@ -560,3 +587,13 @@ If reconnects become persistent (e.g. >10 in a minute) something's wrong: check
|
|||||||
## Contributing
|
## 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`.
|
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.
|
||||||
|
|||||||
+2
-1
@@ -9,7 +9,7 @@ description = "Local CoAP-over-DTLS (OCF) transport + polling layer for Samsung
|
|||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
requires-python = ">=3.11"
|
requires-python = ">=3.11"
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
license-files = ["LICENSE"]
|
license-files = ["LICENSE", "NOTICE"]
|
||||||
authors = [{ name = "Jack Nagy" }]
|
authors = [{ name = "Jack Nagy" }]
|
||||||
keywords = ["smartthings", "samsung", "ocf", "coap", "dtls", "home-assistant", "iot"]
|
keywords = ["smartthings", "samsung", "ocf", "coap", "dtls", "home-assistant", "iot"]
|
||||||
classifiers = [
|
classifiers = [
|
||||||
@@ -60,5 +60,6 @@ include = [
|
|||||||
"tests",
|
"tests",
|
||||||
"README.md",
|
"README.md",
|
||||||
"LICENSE",
|
"LICENSE",
|
||||||
|
"NOTICE",
|
||||||
"pyproject.toml",
|
"pyproject.toml",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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"]
|
||||||
@@ -18,12 +18,11 @@ Reader thread owns the UDP socket. Callers issue get()/post() and block
|
|||||||
on a per-token Event the reader signals. OBSERVE notifications are
|
on a per-token Event the reader signals. OBSERVE notifications are
|
||||||
delivered via the on_notification callback.
|
delivered via the on_notification callback.
|
||||||
"""
|
"""
|
||||||
|
import errno
|
||||||
import os
|
import os
|
||||||
import re as _re
|
|
||||||
import socket
|
import socket
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from OpenSSL import SSL
|
from OpenSSL import SSL
|
||||||
|
|
||||||
@@ -42,15 +41,18 @@ from .coap import (
|
|||||||
encode_options, parse_coap, build_coap, block_value, fmt_code,
|
encode_options, parse_coap, build_coap, block_value, fmt_code,
|
||||||
split_dtls as _split_dtls,
|
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
|
from .endpoint import open_connected_udp_socket
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_OCF_ROOT_CA = str(Path(__file__).parent / 'ocf_root_ca.pem')
|
|
||||||
_DTLS_CIPHERS = b'ECDHE-ECDSA-AES128-GCM-SHA256:@SECLEVEL=0'
|
|
||||||
|
|
||||||
|
|
||||||
# Diagnostic logging — when DEBUG_BRIDGE=1 in env, the bridge dumps
|
# Diagnostic logging — when DEBUG_BRIDGE=1 in env, the bridge dumps
|
||||||
# every received CoAP frame, every /operational/state/vs/0 + /oven/vs/0
|
# 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
|
# + /power/vs/0 + /mode/vs/0-options rep change, the full link tree at
|
||||||
@@ -71,32 +73,21 @@ _BLOCK_ACK_TIMEOUT = 4.0
|
|||||||
# once the ceiling is measured empirically.
|
# once the ceiling is measured empirically.
|
||||||
_DEFAULT_RATE_LIMIT_RPS = 5.0
|
_DEFAULT_RATE_LIMIT_RPS = 5.0
|
||||||
|
|
||||||
_PEM_CERT_RE = _re.compile(
|
# ICMP errors a connected UDP socket surfaces on the next recv. On these
|
||||||
rb'-----BEGIN CERTIFICATE-----.*?-----END CERTIFICATE-----',
|
# appliances they show up while the device is rebooting, while it holds an
|
||||||
_re.DOTALL,
|
# 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:
|
class DtlsCoapSession:
|
||||||
"""Single sustained DTLS-CoAP session.
|
"""Single sustained DTLS-CoAP session.
|
||||||
|
|
||||||
@@ -109,10 +100,9 @@ class DtlsCoapSession:
|
|||||||
code, _ = sess.post(['mode','vs','0'], cbor)
|
code, _ = sess.post(['mode','vs','0'], cbor)
|
||||||
sess.close()
|
sess.close()
|
||||||
|
|
||||||
Cert material comes from either a file pair (cert_path, key_path) or
|
Authentication comes from an immutable provider. For compatibility,
|
||||||
an in-memory PEM pair (cert_pem, key_pem) — exactly one pair required.
|
cert_path/key_path and cert_pem/key_pem create a CertificateAuth provider
|
||||||
The in-memory path exists for callers (e.g. an HA config flow) that
|
internally — exactly one legacy pair is required when auth is omitted.
|
||||||
mint a client cert at runtime and never write it to disk.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
HANDSHAKE_TIMEOUT_S = 12.0
|
HANDSHAKE_TIMEOUT_S = 12.0
|
||||||
@@ -123,17 +113,24 @@ class DtlsCoapSession:
|
|||||||
cert_pem=None, key_pem=None,
|
cert_pem=None, key_pem=None,
|
||||||
on_notification=None, mtu=1200,
|
on_notification=None, mtu=1200,
|
||||||
rate_limit_rps: float = _DEFAULT_RATE_LIMIT_RPS,
|
rate_limit_rps: float = _DEFAULT_RATE_LIMIT_RPS,
|
||||||
local_port=None, family=socket.AF_UNSPEC):
|
local_port=None, family=socket.AF_UNSPEC,
|
||||||
if (cert_path is not None or key_path is not None) and \
|
auth: AuthenticationProvider | None = None):
|
||||||
(cert_pem is not None or key_pem is not 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(
|
raise ValueError(
|
||||||
"pass either cert_path/key_path or cert_pem/key_pem, not both")
|
"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:
|
if cert_pem is None or key_pem is None:
|
||||||
raise ValueError("cert_pem and key_pem must be passed together")
|
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(
|
raise ValueError(
|
||||||
"must pass either cert_path/key_path or cert_pem/key_pem")
|
"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.host = host
|
||||||
self.port = port
|
self.port = port
|
||||||
@@ -141,6 +138,12 @@ class DtlsCoapSession:
|
|||||||
self.key_path = str(key_path) if key_path is not None else None
|
self.key_path = str(key_path) if key_path is not None else None
|
||||||
self.cert_pem = cert_pem
|
self.cert_pem = cert_pem
|
||||||
self.key_pem = key_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.on_notification = on_notification # fn(href, payload_bytes)
|
||||||
self.mtu = mtu
|
self.mtu = mtu
|
||||||
self._min_req_interval = 1.0 / rate_limit_rps
|
self._min_req_interval = 1.0 / rate_limit_rps
|
||||||
@@ -181,6 +184,10 @@ class DtlsCoapSession:
|
|||||||
|
|
||||||
self._stop = threading.Event()
|
self._stop = threading.Event()
|
||||||
self._reader_thread = None
|
self._reader_thread = None
|
||||||
|
# Set while the reader owns the socket. Cleared when it exits for
|
||||||
|
# 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
|
self._last_send_ts = 0.0
|
||||||
|
|
||||||
def pace(self) -> None:
|
def pace(self) -> None:
|
||||||
@@ -196,20 +203,7 @@ class DtlsCoapSession:
|
|||||||
"""DTLS handshake. Blocks up to HANDSHAKE_TIMEOUT_S. Raises
|
"""DTLS handshake. Blocks up to HANDSHAKE_TIMEOUT_S. Raises
|
||||||
ConnectionError / TimeoutError on failure."""
|
ConnectionError / TimeoutError on failure."""
|
||||||
ctx = SSL.Context(SSL.DTLS_METHOD)
|
ctx = SSL.Context(SSL.DTLS_METHOD)
|
||||||
|
self.auth.configure_context(ctx)
|
||||||
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(_DTLS_CIPHERS)
|
|
||||||
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()
|
|
||||||
|
|
||||||
conn = SSL.Connection(ctx, None)
|
conn = SSL.Connection(ctx, None)
|
||||||
conn.set_connect_state()
|
conn.set_connect_state()
|
||||||
@@ -279,11 +273,27 @@ class DtlsCoapSession:
|
|||||||
"""Spawn the reader thread. Must be called after connect()."""
|
"""Spawn the reader thread. Must be called after connect()."""
|
||||||
if self.sock is None:
|
if self.sock is None:
|
||||||
raise RuntimeError("connect() before start_reader()")
|
raise RuntimeError("connect() before start_reader()")
|
||||||
|
self._reader_running.set()
|
||||||
t = threading.Thread(target=self._reader_loop,
|
t = threading.Thread(target=self._reader_loop,
|
||||||
daemon=True, name='dtls-reader')
|
daemon=True, name='dtls-reader')
|
||||||
t.start()
|
t.start()
|
||||||
self._reader_thread = t
|
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):
|
def join(self):
|
||||||
"""Block until the reader thread exits (i.e. socket dies)."""
|
"""Block until the reader thread exits (i.e. socket dies)."""
|
||||||
if self._reader_thread is not None:
|
if self._reader_thread is not None:
|
||||||
@@ -400,7 +410,23 @@ class DtlsCoapSession:
|
|||||||
d = sock.recv(65535)
|
d = sock.recv(65535)
|
||||||
except socket.timeout:
|
except socket.timeout:
|
||||||
continue
|
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
|
return
|
||||||
if not d:
|
if not d:
|
||||||
continue
|
continue
|
||||||
@@ -445,6 +471,8 @@ class DtlsCoapSession:
|
|||||||
if exit_reader:
|
if exit_reader:
|
||||||
return
|
return
|
||||||
finally:
|
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.
|
# Make sure pending waiters don't hang if the reader dies.
|
||||||
for tok, (ev, container) in list(self._pending.items()):
|
for tok, (ev, container) in list(self._pending.items()):
|
||||||
container.setdefault('err', SessionClosedError())
|
container.setdefault('err', SessionClosedError())
|
||||||
@@ -516,8 +544,7 @@ class DtlsCoapSession:
|
|||||||
response — Samsung's server keys per-transfer state on the
|
response — Samsung's server keys per-transfer state on the
|
||||||
token, and dropping a fresh token on block 1+ silently drops
|
token, and dropping a fresh token on block 1+ silently drops
|
||||||
the request."""
|
the request."""
|
||||||
if self.conn is None:
|
self._check_live()
|
||||||
raise SessionClosedError()
|
|
||||||
tok = self._next_tok()
|
tok = self._next_tok()
|
||||||
blob = b''
|
blob = b''
|
||||||
num = 0
|
num = 0
|
||||||
@@ -592,8 +619,7 @@ class DtlsCoapSession:
|
|||||||
def post(self, path_segs, body_cbor, timeout=8.0):
|
def post(self, path_segs, body_cbor, timeout=8.0):
|
||||||
"""Single-frame POST with a CBOR-encoded body. Returns
|
"""Single-frame POST with a CBOR-encoded body. Returns
|
||||||
(code, payload_bytes). body_cbor must already be encoded."""
|
(code, payload_bytes). body_cbor must already be encoded."""
|
||||||
if self.conn is None:
|
self._check_live()
|
||||||
raise SessionClosedError()
|
|
||||||
tok = self._next_tok()
|
tok = self._next_tok()
|
||||||
mid = self._next_mid()
|
mid = self._next_mid()
|
||||||
opts = [(URI_PATH, s.encode()) for s in path_segs]
|
opts = [(URI_PATH, s.encode()) for s in path_segs]
|
||||||
@@ -626,8 +652,7 @@ class DtlsCoapSession:
|
|||||||
Real half-open-session detection lives in PollScheduler's
|
Real half-open-session detection lives in PollScheduler's
|
||||||
`last_success_ts`, surfaced through KeepaliveTask's
|
`last_success_ts`, surfaced through KeepaliveTask's
|
||||||
`liveness_fn`."""
|
`liveness_fn`."""
|
||||||
if self.conn is None:
|
self._check_live()
|
||||||
raise SessionClosedError()
|
|
||||||
mid = self._next_mid()
|
mid = self._next_mid()
|
||||||
self._send_dgram(build_coap(TYPE_CON, 0, mid, b'', []))
|
self._send_dgram(build_coap(TYPE_CON, 0, mid, b'', []))
|
||||||
return mid
|
return mid
|
||||||
@@ -644,8 +669,7 @@ class DtlsCoapSession:
|
|||||||
tokens via subscribe. Brief race window where a notify on the
|
tokens via subscribe. Brief race window where a notify on the
|
||||||
old token gets dropped as 'stale' — acceptable for a 6h-scale
|
old token gets dropped as 'stale' — acceptable for a 6h-scale
|
||||||
safety net."""
|
safety net."""
|
||||||
if self.conn is None:
|
self._check_live()
|
||||||
raise SessionClosedError()
|
|
||||||
for tok, href in list(self._observe_tokens.items()):
|
for tok, href in list(self._observe_tokens.items()):
|
||||||
segs = [s for s in href.split('/') if s]
|
segs = [s for s in href.split('/') if s]
|
||||||
try:
|
try:
|
||||||
@@ -668,8 +692,7 @@ class DtlsCoapSession:
|
|||||||
|
|
||||||
Returns the token used (in case the caller wants to deregister
|
Returns the token used (in case the caller wants to deregister
|
||||||
later)."""
|
later)."""
|
||||||
if self.conn is None:
|
self._check_live()
|
||||||
raise SessionClosedError()
|
|
||||||
tok = self._next_observe_tok()
|
tok = self._next_observe_tok()
|
||||||
href = '/' + '/'.join(path_segs)
|
href = '/' + '/'.join(path_segs)
|
||||||
# Register the token BEFORE sending — otherwise the device
|
# Register the token BEFORE sending — otherwise the device
|
||||||
|
|||||||
@@ -1,31 +1,69 @@
|
|||||||
|
import gc
|
||||||
|
import traceback
|
||||||
|
import weakref
|
||||||
|
from dataclasses import asdict
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
|
||||||
import pytest
|
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 OpenSSL import SSL, crypto
|
||||||
|
|
||||||
|
from smartthings_local.protocol.auth import CertificateAuth
|
||||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession, _load_pem_chain
|
from smartthings_local.protocol.dtls_session import DtlsCoapSession, _load_pem_chain
|
||||||
|
|
||||||
|
|
||||||
def _make_self_signed_pem_pair():
|
def _make_generated_pem_chain():
|
||||||
"""A throwaway self-signed cert + key, just to exercise PEM loading —
|
"""Create a throwaway leaf + root chain unrelated to Samsung devices."""
|
||||||
not meant to resemble a real Samsung client cert."""
|
now = datetime.now(timezone.utc)
|
||||||
key = crypto.PKey()
|
root_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||||
key.generate_key(crypto.TYPE_RSA, 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()
|
leaf_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||||
cert.get_subject().CN = "test"
|
leaf_name = x509.Name(
|
||||||
cert.set_serial_number(1)
|
[x509.NameAttribute(NameOID.COMMON_NAME, "Synthetic test client")]
|
||||||
cert.gmtime_adj_notBefore(0)
|
)
|
||||||
cert.gmtime_adj_notAfter(3600)
|
leaf_cert = (
|
||||||
cert.set_issuer(cert.get_subject())
|
x509.CertificateBuilder()
|
||||||
cert.set_pubkey(key)
|
.subject_name(leaf_name)
|
||||||
cert.sign(key, "sha256")
|
.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()
|
cert_pem = (
|
||||||
key_pem = crypto.dump_privatekey(crypto.FILETYPE_PEM, key).decode()
|
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
|
return cert_pem, key_pem
|
||||||
|
|
||||||
|
|
||||||
def test_load_pem_chain_loads_cert_and_key_in_memory():
|
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)
|
ctx = SSL.Context(SSL.DTLS_METHOD)
|
||||||
_load_pem_chain(ctx, cert_pem, key_pem)
|
_load_pem_chain(ctx, cert_pem, key_pem)
|
||||||
ctx.check_privatekey() # raises if cert/key don't match
|
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():
|
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):
|
with pytest.raises(ValueError):
|
||||||
DtlsCoapSession("host", 1234) # neither pair given
|
DtlsCoapSession("host", 1234) # neither pair given
|
||||||
@@ -49,10 +87,226 @@ def test_session_requires_exactly_one_cert_source():
|
|||||||
with pytest.raises(ValueError):
|
with pytest.raises(ValueError):
|
||||||
DtlsCoapSession("host", 1234, cert_pem=cert_pem) # key_pem missing
|
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():
|
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)
|
sess = DtlsCoapSession("host", 1234, cert_pem=cert_pem, key_pem=key_pem)
|
||||||
assert sess.cert_path is None
|
assert sess.cert_path is None
|
||||||
assert sess.key_path is None
|
assert sess.key_path is None
|
||||||
assert sess.cert_pem == cert_pem
|
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,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()
|
||||||
@@ -6,6 +6,11 @@ import inspect
|
|||||||
|
|
||||||
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
|
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
|
||||||
from smartthings_local.ocf.state_cache import StateCache
|
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
|
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
||||||
|
|
||||||
|
|
||||||
@@ -40,6 +45,26 @@ def test_dtls_session_constructor_keeps_file_memory_and_local_port_inputs():
|
|||||||
"local_port",
|
"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():
|
def test_dtls_session_keeps_current_consumer_methods():
|
||||||
|
|||||||
@@ -59,6 +59,7 @@ def check_wheel(path: Path) -> None:
|
|||||||
f"{dist_info}/METADATA",
|
f"{dist_info}/METADATA",
|
||||||
f"{dist_info}/WHEEL",
|
f"{dist_info}/WHEEL",
|
||||||
f"{dist_info}/licenses/LICENSE",
|
f"{dist_info}/licenses/LICENSE",
|
||||||
|
f"{dist_info}/licenses/NOTICE",
|
||||||
f"{dist_info}/RECORD",
|
f"{dist_info}/RECORD",
|
||||||
}
|
}
|
||||||
if metadata != expected_metadata:
|
if metadata != expected_metadata:
|
||||||
@@ -83,6 +84,7 @@ def check_sdist(path: Path) -> None:
|
|||||||
relative = {name[len(root) + 1 :] for name in names if name.startswith(f"{root}/")}
|
relative = {name[len(root) + 1 :] for name in names if name.startswith(f"{root}/")}
|
||||||
required = _tracked_files() | {
|
required = _tracked_files() | {
|
||||||
"LICENSE",
|
"LICENSE",
|
||||||
|
"NOTICE",
|
||||||
"PKG-INFO",
|
"PKG-INFO",
|
||||||
"README.md",
|
"README.md",
|
||||||
"pyproject.toml",
|
"pyproject.toml",
|
||||||
|
|||||||
Reference in New Issue
Block a user