diff --git a/README.md b/README.md index 33c9001..52c2f41 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/smartthings_local/protocol/auth.py b/smartthings_local/protocol/auth.py index 8291e09..0e99712 100644 --- a/smartthings_local/protocol/auth.py +++ b/smartthings_local/protocol/auth.py @@ -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"] diff --git a/tests/test_psk_auth.py b/tests/test_psk_auth.py new file mode 100644 index 0000000..c4e957a --- /dev/null +++ b/tests/test_psk_auth.py @@ -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() diff --git a/tests/test_public_api_contract.py b/tests/test_public_api_contract.py index 5125819..f40e5b4 100644 --- a/tests/test_public_api_contract.py +++ b/tests/test_public_api_contract.py @@ -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",