Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 34 additions & 5 deletions smartthings_local/protocol/dtls_probe.py
Original file line number Diff line number Diff line change
Expand Up @@ -750,6 +750,38 @@ def _diagnostic_context(*, auth, cert_pem, key_pem, cert_path, key_path):
return ctx


def _diagnostic_connection(
*, auth, cert_pem, key_pem, cert_path, key_path, mtu):
"""Build the connection one diagnostic run drives.

The factory is asked before the context path, and that order is the
whole point: a built-in provider carrying its own engine keeps
``configure_context`` for direct callers, so ``_validate_diagnostic_auth``
accepts it either way and nothing further down would notice a credential
being handed to the wrong engine.

Nothing is relaxed for a provider that brings its own engine, because
there is nothing to relax: a diagnostic must report the appliance's own
alert rather than a local verdict, and an engine with no X.509 reaches
that by construction instead of by re-asserting accept-any.
"""
factory = getattr(auth, "_create_dtls_connection", None)
if factory is not None:
return factory(mtu=mtu)

ctx = _diagnostic_context(
auth=auth,
cert_pem=cert_pem,
key_pem=key_pem,
cert_path=cert_path,
key_path=key_path,
)
conn = SSL.Connection(ctx, None)
conn.set_connect_state()
conn.set_ciphertext_mtu(mtu)
return conn


def diagnose_dtls_handshake(
host, port, *, auth=None, cert_pem=None, key_pem=None,
cert_path=None, key_path=None,
Expand Down Expand Up @@ -780,18 +812,15 @@ def diagnose_dtls_handshake(
_validate_diagnostic_auth(auth, cert_pem, key_pem, cert_path, key_path)
result = ProbeResult(host, port)

ctx = _diagnostic_context(
conn = _diagnostic_connection(
auth=auth,
cert_pem=cert_pem,
key_pem=key_pem,
cert_path=cert_path,
key_path=key_path,
mtu=mtu,
)

conn = SSL.Connection(ctx, None)
conn.set_connect_state()
conn.set_ciphertext_mtu(mtu)

try:
sock, _endpoint = open_host_filtered_udp_socket(
host,
Expand Down
70 changes: 58 additions & 12 deletions smartthings_local/protocol/dtls_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -731,6 +731,63 @@ def _pace_orderly_close(self) -> None:

# ---- lifecycle ---------------------------------------------------

def _new_dtls_connection(self, cancel):
"""Build the DTLS connection for one handshake attempt.

The single construction site for a session's connection, so an
engine other than pyOpenSSL can be reached without an isinstance
fork spreading through connect() and the reader loop.

A provider may carry a private ``_create_dtls_connection`` factory,
which owns its own engine and never sees an ``SSL.Context``. Found
by ``getattr`` rather than declared on ``AuthenticationProvider``:
that Protocol is ``runtime_checkable`` and ``__init__`` gates on
``isinstance``, so a required method would reject every third-party
provider that has not grown one. A provider that deliberately
supplies the hook is therefore routed whether or not this package
ships it, which does not make it a supported public extension API --
a real one would be designed separately.

The contract the hook has to meet, for anyone experimenting with
another engine:

- Return a fresh connection for this attempt, already in client
state, with ``mtu`` applied.
- Own no socket, and start no network I/O. The caller owns the
socket, the cancellation and deadline handling, and publishing
the session.
- Raise and behave like the memory-BIO subset of
``OpenSSL.SSL.Connection`` that ``_drive_dtls_handshake`` drives:
``WantReadError`` until a handshake completes, ``ZeroReturnError``
on an orderly close, ``Error`` otherwise.

Cancellation is checked before and after the provider's work, never
inside it. A ``configure_context`` or factory call already blocked
on a slow PEM read is not interrupted; what the checks guarantee is
that no socket is created once cancellation is set.
"""
factory = getattr(self.auth, "_create_dtls_connection", None)
if factory is not None:
conn = factory(mtu=self.mtu)
if self._lifecycle_cancel.is_set() or \
(cancel is not None and cancel.is_set()):
raise SessionClosedError()
return conn

ctx = SSL.Context(SSL.DTLS_METHOD)
self.auth.configure_context(ctx)
if self._lifecycle_cancel.is_set() or \
(cancel is not None and cancel.is_set()):
raise SessionClosedError()

conn = SSL.Connection(ctx, None)
conn.set_connect_state()
conn.set_ciphertext_mtu(self.mtu)
if self._lifecycle_cancel.is_set() or \
(cancel is not None and cancel.is_set()):
raise SessionClosedError()
return conn

def connect(
self,
*,
Expand Down Expand Up @@ -778,18 +835,7 @@ def connect(
(cancel is not None and cancel.is_set()):
raise SessionClosedError()
deadline = time.monotonic() + handshake_timeout
ctx = SSL.Context(SSL.DTLS_METHOD)
self.auth.configure_context(ctx)
if self._lifecycle_cancel.is_set() or \
(cancel is not None and cancel.is_set()):
raise SessionClosedError()

conn = SSL.Connection(ctx, None)
conn.set_connect_state()
conn.set_ciphertext_mtu(self.mtu)
if self._lifecycle_cancel.is_set() or \
(cancel is not None and cancel.is_set()):
raise SessionClosedError()
conn = self._new_dtls_connection(cancel)

remaining = deadline - time.monotonic()
if remaining <= 0:
Expand Down
Loading
Loading