4 Commits
Author SHA1 Message Date
Quite Yellow b0d51abcc8 Merge pull request #38 from QuiteYellow/fix/reader-death-visible
fix(dtls): make reader-thread death visible and fail fast
2026-08-14 17:38:41 +01:00
Jack Nagy 7a74a955f3 fix(dtls): make reader-thread death visible and fail fast
The reader loop exited silently on any socket error, leaving conn/sock
set so the session still looked open. Every later get()/post()/ping()
then waited out its full request timeout on a session nobody was
reading, raising SessionTimeoutError on repeat, forever.

This started biting in v0.1.3 (d677c72), which moved to connected UDP
sockets: a connected socket surfaces ICMP errors on recv, so one
ECONNREFUSED from a rebooting appliance now killed the reader.

- Advisory ICMP errnos (ECONNREFUSED/EHOSTUNREACH/...) no longer kill
  the reader; the next datagram usually works.
- Real reader exits log at WARNING; close()-driven exits stay quiet.
- A _reader_running Event lets get/post/ping/subscribe/refresh_observes
  fail fast via _check_live() with SessionClosedError instead of waiting
  out a timeout. Callers that never start a reader are unaffected.

Refs QuiteYellow/SmartThings-Local#37
2026-08-14 17:33:13 +01:00
Quite Yellow d4aebebca4 Merge pull request #32 from Moballo-LLC/codex/py-06-psk-auth
feat(protocol): add PSK authentication provider
2026-08-09 14:21:54 +01:00
Jason Morcos 2a6fc627f2 feat(protocol): add PSK authentication provider 2026-08-08 16:06:53 -07:00
6 changed files with 727 additions and 15 deletions
+15
View File
@@ -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
+82 -3
View File
@@ -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"]
+60 -11
View File
@@ -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
+173
View File
@@ -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()
+380
View File
@@ -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()
+17 -1
View File
@@ -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",