Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b0d51abcc8 | ||
|
|
7a74a955f3 | ||
|
|
d4aebebca4 | ||
|
|
2a6fc627f2 |
@@ -61,6 +61,21 @@ For compatibility, the existing `cert_path` / `key_path` and `cert_pem` /
|
||||
They are routed through `CertificateAuth` internally. Do not combine `auth`
|
||||
with those legacy arguments.
|
||||
|
||||
An existing OCF PSK credential can be supplied through `PskAuth`:
|
||||
|
||||
```python
|
||||
from smartthings_local.protocol.auth import PskAuth
|
||||
|
||||
auth = PskAuth(identity=psk_identity, key=psk_key)
|
||||
sess = DtlsCoapSession("192.0.2.100", 49154, auth=auth)
|
||||
```
|
||||
|
||||
The identity must be the raw 16-byte OCF UUID and cannot contain a NUL byte;
|
||||
the key must be exactly 16 or 32 bytes. `PskAuth` selects only
|
||||
`ECDHE-PSK-AES128-CBC-SHA256` and does not acquire, derive, provision, rotate,
|
||||
or persist credentials. Ownership transfer and credential discovery are
|
||||
outside this package.
|
||||
|
||||
### Classified errors
|
||||
|
||||
Runtime transport failures use the public types in
|
||||
|
||||
@@ -7,10 +7,15 @@ from os import PathLike
|
||||
from pathlib import Path
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
from OpenSSL import SSL, crypto
|
||||
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,
|
||||
@@ -45,7 +50,7 @@ class AuthenticationProvider(Protocol):
|
||||
"""Configure authentication for a newly created DTLS context."""
|
||||
|
||||
def configure_context(self, context: SSL.Context) -> None:
|
||||
"""Configure trust, verification, ciphers, and client credentials."""
|
||||
"""Configure a context while this provider remains session-owned."""
|
||||
|
||||
|
||||
class CertificateAuth:
|
||||
@@ -164,4 +169,78 @@ class CertificateAuth:
|
||||
context.check_privatekey()
|
||||
|
||||
|
||||
__all__ = ["AuthenticationProvider", "CertificateAuth"]
|
||||
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,6 +18,7 @@ Reader thread owns the UDP socket. Callers issue get()/post() and block
|
||||
on a per-token Event the reader signals. OBSERVE notifications are
|
||||
delivered via the on_notification callback.
|
||||
"""
|
||||
import errno
|
||||
import os
|
||||
import socket
|
||||
import threading
|
||||
@@ -72,6 +73,21 @@ _BLOCK_ACK_TIMEOUT = 4.0
|
||||
# once the ceiling is measured empirically.
|
||||
_DEFAULT_RATE_LIMIT_RPS = 5.0
|
||||
|
||||
# ICMP errors a connected UDP socket surfaces on the next recv. On these
|
||||
# appliances they show up while the device is rebooting, while it holds an
|
||||
# orphaned association, or across a router blip, and the next datagram
|
||||
# usually works. UDP delivery was never guaranteed, so treat them as
|
||||
# advisory and keep reading. Unconnected sockets never see any of this,
|
||||
# which is why the reader survived them before the connected-socket change
|
||||
# in d677c72 (v0.1.3).
|
||||
_ADVISORY_ERRNOS = frozenset(
|
||||
value for value in (
|
||||
getattr(errno, name, None)
|
||||
for name in ('ECONNREFUSED', 'EHOSTUNREACH', 'ENETUNREACH',
|
||||
'EHOSTDOWN', 'ENETDOWN')
|
||||
) if value is not None
|
||||
)
|
||||
|
||||
class DtlsCoapSession:
|
||||
"""Single sustained DTLS-CoAP session.
|
||||
|
||||
@@ -168,6 +184,10 @@ class DtlsCoapSession:
|
||||
|
||||
self._stop = threading.Event()
|
||||
self._reader_thread = None
|
||||
# Set while the reader owns the socket. Cleared when it exits for
|
||||
# any reason, so callers fail fast through _check_live() instead of
|
||||
# waiting out a request timeout against a session nobody is reading.
|
||||
self._reader_running = threading.Event()
|
||||
self._last_send_ts = 0.0
|
||||
|
||||
def pace(self) -> None:
|
||||
@@ -253,11 +273,27 @@ class DtlsCoapSession:
|
||||
"""Spawn the reader thread. Must be called after connect()."""
|
||||
if self.sock is None:
|
||||
raise RuntimeError("connect() before start_reader()")
|
||||
self._reader_running.set()
|
||||
t = threading.Thread(target=self._reader_loop,
|
||||
daemon=True, name='dtls-reader')
|
||||
t.start()
|
||||
self._reader_thread = t
|
||||
|
||||
def _check_live(self):
|
||||
"""Raise if the session cannot carry a request. A dead reader is
|
||||
as fatal as a closed connection: the socket may still accept
|
||||
sends, but no response will ever be dispatched, so waiting out the
|
||||
request timeout only delays the inevitable SessionClosedError.
|
||||
|
||||
Callers that never start a reader (config-flow style) keep the old
|
||||
behaviour — only the conn check applies while _reader_thread is
|
||||
None."""
|
||||
if self.conn is None:
|
||||
raise SessionClosedError()
|
||||
if self._reader_thread is not None and \
|
||||
not self._reader_running.is_set():
|
||||
raise SessionClosedError()
|
||||
|
||||
def join(self):
|
||||
"""Block until the reader thread exits (i.e. socket dies)."""
|
||||
if self._reader_thread is not None:
|
||||
@@ -374,7 +410,23 @@ class DtlsCoapSession:
|
||||
d = sock.recv(65535)
|
||||
except socket.timeout:
|
||||
continue
|
||||
except (OSError, ValueError):
|
||||
except OSError as e:
|
||||
if self._stop.is_set():
|
||||
return # close() got here first
|
||||
if e.errno in _ADVISORY_ERRNOS:
|
||||
logger.debug("reader: advisory %s from %s, continuing",
|
||||
errno.errorcode.get(e.errno, e.errno),
|
||||
self.host)
|
||||
continue
|
||||
logger.warning("reader exiting: socket error %s from %s",
|
||||
errno.errorcode.get(e.errno, e.errno),
|
||||
self.host)
|
||||
return
|
||||
except ValueError:
|
||||
# recv on a socket closed underneath the reader.
|
||||
if not self._stop.is_set():
|
||||
logger.warning("reader exiting: socket closed "
|
||||
"underneath it")
|
||||
return
|
||||
if not d:
|
||||
continue
|
||||
@@ -419,6 +471,8 @@ class DtlsCoapSession:
|
||||
if exit_reader:
|
||||
return
|
||||
finally:
|
||||
# Reader no longer owns the socket — callers must fail fast.
|
||||
self._reader_running.clear()
|
||||
# Make sure pending waiters don't hang if the reader dies.
|
||||
for tok, (ev, container) in list(self._pending.items()):
|
||||
container.setdefault('err', SessionClosedError())
|
||||
@@ -490,8 +544,7 @@ class DtlsCoapSession:
|
||||
response — Samsung's server keys per-transfer state on the
|
||||
token, and dropping a fresh token on block 1+ silently drops
|
||||
the request."""
|
||||
if self.conn is None:
|
||||
raise SessionClosedError()
|
||||
self._check_live()
|
||||
tok = self._next_tok()
|
||||
blob = b''
|
||||
num = 0
|
||||
@@ -566,8 +619,7 @@ class DtlsCoapSession:
|
||||
def post(self, path_segs, body_cbor, timeout=8.0):
|
||||
"""Single-frame POST with a CBOR-encoded body. Returns
|
||||
(code, payload_bytes). body_cbor must already be encoded."""
|
||||
if self.conn is None:
|
||||
raise SessionClosedError()
|
||||
self._check_live()
|
||||
tok = self._next_tok()
|
||||
mid = self._next_mid()
|
||||
opts = [(URI_PATH, s.encode()) for s in path_segs]
|
||||
@@ -600,8 +652,7 @@ class DtlsCoapSession:
|
||||
Real half-open-session detection lives in PollScheduler's
|
||||
`last_success_ts`, surfaced through KeepaliveTask's
|
||||
`liveness_fn`."""
|
||||
if self.conn is None:
|
||||
raise SessionClosedError()
|
||||
self._check_live()
|
||||
mid = self._next_mid()
|
||||
self._send_dgram(build_coap(TYPE_CON, 0, mid, b'', []))
|
||||
return mid
|
||||
@@ -618,8 +669,7 @@ class DtlsCoapSession:
|
||||
tokens via subscribe. Brief race window where a notify on the
|
||||
old token gets dropped as 'stale' — acceptable for a 6h-scale
|
||||
safety net."""
|
||||
if self.conn is None:
|
||||
raise SessionClosedError()
|
||||
self._check_live()
|
||||
for tok, href in list(self._observe_tokens.items()):
|
||||
segs = [s for s in href.split('/') if s]
|
||||
try:
|
||||
@@ -642,8 +692,7 @@ class DtlsCoapSession:
|
||||
|
||||
Returns the token used (in case the caller wants to deregister
|
||||
later)."""
|
||||
if self.conn is None:
|
||||
raise SessionClosedError()
|
||||
self._check_live()
|
||||
tok = self._next_observe_tok()
|
||||
href = '/' + '/'.join(path_segs)
|
||||
# Register the token BEFORE sending — otherwise the device
|
||||
|
||||
@@ -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,7 +6,11 @@ import inspect
|
||||
|
||||
from smartthings_local.ocf.observe_refresh import ObserveRefreshTask
|
||||
from smartthings_local.ocf.state_cache import StateCache
|
||||
from smartthings_local.protocol.auth import AuthenticationProvider, CertificateAuth
|
||||
from smartthings_local.protocol.auth import (
|
||||
AuthenticationProvider,
|
||||
CertificateAuth,
|
||||
PskAuth,
|
||||
)
|
||||
from smartthings_local.protocol.dtls_session import DtlsCoapSession
|
||||
|
||||
|
||||
@@ -51,6 +55,18 @@ def test_certificate_auth_is_a_public_authentication_provider():
|
||||
assert isinstance(provider, AuthenticationProvider)
|
||||
|
||||
|
||||
def test_psk_auth_is_a_public_authentication_provider():
|
||||
provider = PskAuth(identity=b"i" * 16, key=b"k" * 16)
|
||||
assert isinstance(provider, AuthenticationProvider)
|
||||
parameters = inspect.signature(PskAuth).parameters
|
||||
assert list(parameters) == ["identity", "key"]
|
||||
assert all(
|
||||
parameter.kind is inspect.Parameter.KEYWORD_ONLY
|
||||
and parameter.default is inspect.Parameter.empty
|
||||
for parameter in parameters.values()
|
||||
)
|
||||
|
||||
|
||||
def test_dtls_session_keeps_current_consumer_methods():
|
||||
expected = {
|
||||
"close",
|
||||
|
||||
Reference in New Issue
Block a user