313 lines
10 KiB
Python
313 lines
10 KiB
Python
import gc
|
|
import traceback
|
|
import weakref
|
|
from dataclasses import asdict
|
|
from datetime import datetime, timedelta, timezone
|
|
|
|
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 smartthings_local.protocol.auth import CertificateAuth
|
|
from smartthings_local.protocol.dtls_session import DtlsCoapSession, _load_pem_chain
|
|
|
|
|
|
def _make_generated_pem_chain():
|
|
"""Create a throwaway leaf + root chain unrelated to Samsung devices."""
|
|
now = datetime.now(timezone.utc)
|
|
root_key = rsa.generate_private_key(public_exponent=65537, key_size=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())
|
|
)
|
|
|
|
leaf_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
|
leaf_name = x509.Name(
|
|
[x509.NameAttribute(NameOID.COMMON_NAME, "Synthetic test client")]
|
|
)
|
|
leaf_cert = (
|
|
x509.CertificateBuilder()
|
|
.subject_name(leaf_name)
|
|
.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 = (
|
|
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
|
|
|
|
|
|
def test_load_pem_chain_loads_cert_and_key_in_memory():
|
|
cert_pem, key_pem = _make_generated_pem_chain()
|
|
ctx = SSL.Context(SSL.DTLS_METHOD)
|
|
_load_pem_chain(ctx, cert_pem, key_pem)
|
|
ctx.check_privatekey() # raises if cert/key don't match
|
|
|
|
|
|
def test_load_pem_chain_rejects_cert_pem_with_no_certificates():
|
|
with pytest.raises(ValueError):
|
|
_load_pem_chain(SSL.Context(SSL.DTLS_METHOD), "not a cert", "not a key")
|
|
|
|
|
|
def test_session_requires_exactly_one_cert_source():
|
|
cert_pem, key_pem = _make_generated_pem_chain()
|
|
|
|
with pytest.raises(ValueError):
|
|
DtlsCoapSession("host", 1234) # neither pair given
|
|
|
|
with pytest.raises(ValueError):
|
|
DtlsCoapSession("host", 1234, cert_path="/a", key_path="/b",
|
|
cert_pem=cert_pem, key_pem=key_pem) # both given
|
|
|
|
with pytest.raises(ValueError):
|
|
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():
|
|
cert_pem, key_pem = _make_generated_pem_chain()
|
|
sess = DtlsCoapSession("host", 1234, cert_pem=cert_pem, key_pem=key_pem)
|
|
assert sess.cert_path is None
|
|
assert sess.key_path is None
|
|
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
|