Merge pull request #24 from Moballo-LLC/codex/py-03-connected-endpoints

feat(protocol): resolve and connect UDP endpoints
This commit is contained in:
Quite Yellow
2026-08-07 15:37:36 +01:00
committed by GitHub
5 changed files with 663 additions and 15 deletions
+30
View File
@@ -79,6 +79,36 @@ failure is chained for debugging, the cause is replaced with a fixed redacted
marker; raw backend text is not copied into the public error or its formatted
traceback.
### Resolved UDP endpoints
Sessions resolve a host to a first-class `ResolvedUdpEndpoint` and use a
connected UDP socket for the DTLS transport. Connecting the datagram socket
pins it to the exact resolved peer, so unrelated datagrams from another host
using the same port are discarded by the operating system. IPv4, IPv6, and
scoped IPv6 tuples are preserved without putting the address or scope in the
endpoint's `repr`.
Address family and fixed source-port behavior are explicit and optional:
```python
import socket
sess = DtlsCoapSession(
"device.example",
49154,
cert_pem=cert_pem,
key_pem=key_pem,
family=socket.AF_INET6,
local_port=56830,
)
sess.connect()
assert sess.endpoint.family == socket.AF_INET6
```
The resolver retains candidate order and the socket setup tries the next
candidate after a family, bind, or connect failure. Resolution and socket
setup failures raise the redacted `EndpointError` documented above.
For a full worked integration, the higher-level `smartthings_local.ocf` layer (`StateCache`, `PollScheduler`, `KeepaliveTask`, `ObserveRefreshTask`) coordinates tiered polling and OBSERVE on top of a session. The MQTT bridge demo below wires all of it together.
### What the demo bridge gives you
+39 -14
View File
@@ -29,6 +29,7 @@ from OpenSSL import SSL
from ..errors import (
BlockwiseError,
EndpointError,
SessionClosedError,
SessionError,
SessionTimeoutError,
@@ -41,6 +42,7 @@ from .coap import (
encode_options, parse_coap, build_coap, block_value, fmt_code,
split_dtls as _split_dtls,
)
from .endpoint import open_connected_udp_socket
import logging
logger = logging.getLogger(__name__)
@@ -120,7 +122,7 @@ class DtlsCoapSession:
cert_pem=None, key_pem=None,
on_notification=None, mtu=1200,
rate_limit_rps: float = _DEFAULT_RATE_LIMIT_RPS,
local_port=None):
local_port=None, family=socket.AF_UNSPEC):
if (cert_path is not None or key_path is not None) and \
(cert_pem is not None or key_pem is not None):
raise ValueError(
@@ -152,10 +154,12 @@ class DtlsCoapSession:
# the new handshake and discard the old association. Verified
# accepted by RT-OCF (oven, 2026-07-26).
self.local_port = local_port
self.family = family
self.sock = None
self.conn = None
self.dest = None
self.endpoint = None
self._send_lock = threading.Lock()
# Randomize MID and token counter starting points so reconnects
@@ -210,15 +214,14 @@ class DtlsCoapSession:
conn.set_connect_state()
conn.set_ciphertext_mtu(self.mtu)
sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
if self.local_port is not None:
# Fixed source port → same 5-tuple on reconnect, so the device
# evicts any orphaned association per RFC 6347 §4.2.8 instead
# of serving a second one alongside it. See __init__.
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
sock.bind(('', self.local_port))
sock.settimeout(2.0)
dest = (self.host, self.port)
sock, endpoint = open_connected_udp_socket(
self.host,
self.port,
family=self.family,
local_port=self.local_port,
timeout=2.0,
)
dest = endpoint.sockaddr
t0 = time.time()
backend_failed = False
@@ -232,19 +235,32 @@ class DtlsCoapSession:
sock.close()
backend_failed = True
break
send_failed = False
try:
o = conn.bio_read(65535)
if o:
for r in _split_dtls(o):
sock.sendto(r, dest)
if sock.send(r) != len(r):
raise OSError('incomplete UDP send')
except SSL.WantReadError:
pass
except OSError:
sock.close()
send_failed = True
if send_failed:
raise EndpointError() from OSError('UDP send failed')
receive_failed = False
try:
d, _ = sock.recvfrom(65535)
d = sock.recv(65535)
if d:
conn.bio_write(d)
except socket.timeout:
pass
except OSError:
sock.close()
receive_failed = True
if receive_failed:
raise EndpointError() from OSError('UDP receive failed')
time.sleep(0.05)
else:
sock.close()
@@ -255,6 +271,7 @@ class DtlsCoapSession:
self.sock = sock
self.conn = conn
self.dest = dest
self.endpoint = endpoint
self._stop.clear()
def start_reader(self):
@@ -318,6 +335,8 @@ class DtlsCoapSession:
self._observe_tokens.clear()
self.sock = None
self.conn = None
self.dest = None
self.endpoint = None
# ---- send / receive plumbing -------------------------------------
@@ -350,6 +369,7 @@ class DtlsCoapSession:
with self._send_lock:
if self.conn is None:
raise SessionClosedError()
send_failed = False
try:
self.conn.send(datagram)
self._last_send_ts = time.monotonic()
@@ -358,9 +378,14 @@ class DtlsCoapSession:
if not o:
break
for r in _split_dtls(o):
self.sock.sendto(r, self.dest)
if self.sock.send(r) != len(r):
raise OSError('incomplete UDP send')
except SSL.WantReadError:
pass
except OSError:
send_failed = True
if send_failed:
raise EndpointError() from OSError('UDP send failed')
def _reader_loop(self):
"""Pump UDP socket → DTLS BIO → CoAP parser. Demuxes to pending
@@ -371,7 +396,7 @@ class DtlsCoapSession:
try:
while not self._stop.is_set():
try:
d, _ = sock.recvfrom(65535)
d = sock.recv(65535)
except socket.timeout:
continue
except (OSError, ValueError):
+164
View File
@@ -0,0 +1,164 @@
"""Deterministic UDP endpoint resolution and connected socket setup."""
import math
import socket
from dataclasses import dataclass
from ..errors import EndpointError
__all__ = [
'ResolvedUdpEndpoint',
'open_connected_udp_socket',
'resolve_udp_endpoint',
'resolve_udp_endpoints',
]
_SUPPORTED_FAMILIES = (socket.AF_UNSPEC, socket.AF_INET, socket.AF_INET6)
def _validate_port(port, *, allow_zero=False):
lower_bound = 0 if allow_zero else 1
if isinstance(port, bool) or not isinstance(port, int):
raise TypeError('port must be an integer')
if not lower_bound <= port <= 65535:
raise ValueError(f'port must be between {lower_bound} and 65535')
def _validate_family(family):
if isinstance(family, bool) or not isinstance(family, int):
raise TypeError('family must be an address-family integer')
if family not in _SUPPORTED_FAMILIES:
raise ValueError('family must be AF_UNSPEC, AF_INET, or AF_INET6')
@dataclass(frozen=True, slots=True, repr=False)
class ResolvedUdpEndpoint:
"""One concrete IPv4 or IPv6 UDP destination.
``sockaddr`` is the exact tuple returned by ``getaddrinfo`` and therefore
retains IPv6 flow and scope IDs. The custom representation omits the
address, port, and scope so exception and diagnostic output does not leak
a remote endpoint accidentally.
"""
family: int
sockaddr: tuple
def __post_init__(self):
if self.family not in (socket.AF_INET, socket.AF_INET6):
raise ValueError('resolved family must be AF_INET or AF_INET6')
expected_length = 2 if self.family == socket.AF_INET else 4
if not isinstance(self.sockaddr, tuple) or \
len(self.sockaddr) != expected_length:
raise ValueError('sockaddr does not match its address family')
_validate_port(self.sockaddr[1])
@property
def host(self):
return self.sockaddr[0]
@property
def port(self):
return self.sockaddr[1]
@property
def scope_id(self):
return self.sockaddr[3] if self.family == socket.AF_INET6 else 0
@property
def family_name(self):
return socket.AddressFamily(self.family).name
def bind_address(self, local_port):
"""Return the wildcard bind tuple matching this endpoint's family."""
_validate_port(local_port, allow_zero=True)
if self.family == socket.AF_INET6:
return ('::', local_port, 0, 0)
return ('', local_port)
def __repr__(self):
return f'ResolvedUdpEndpoint(family={self.family_name})'
def resolve_udp_endpoints(host, port, *, family=socket.AF_UNSPEC):
"""Resolve all unique IPv4/IPv6 UDP candidates in resolver order."""
if not isinstance(host, (str, bytes)):
raise TypeError('host must be a string or bytes value')
if not host:
raise ValueError('host must be a non-empty string or bytes value')
_validate_port(port)
_validate_family(family)
infos = None
try:
infos = socket.getaddrinfo(
host, port, family, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
except (OSError, UnicodeError):
pass
if infos is None:
raise EndpointError() from OSError('UDP endpoint resolution failed')
endpoints = []
seen = set()
for resolved_family, socktype, protocol, _canonname, sockaddr in infos:
if resolved_family not in (socket.AF_INET, socket.AF_INET6):
continue
if socktype not in (0, socket.SOCK_DGRAM):
continue
if protocol not in (0, socket.IPPROTO_UDP):
continue
expected_length = 2 if resolved_family == socket.AF_INET else 4
if not isinstance(sockaddr, tuple) or len(sockaddr) != expected_length:
continue
key = (resolved_family, sockaddr)
if key in seen:
continue
seen.add(key)
endpoints.append(ResolvedUdpEndpoint(resolved_family, sockaddr))
if not endpoints:
raise EndpointError()
return tuple(endpoints)
def resolve_udp_endpoint(host, port, *, family=socket.AF_UNSPEC):
"""Resolve the first usable UDP candidate."""
return resolve_udp_endpoints(host, port, family=family)[0]
def open_connected_udp_socket(
host, port, *, family=socket.AF_UNSPEC, local_port=None, timeout=None):
"""Create, optionally bind, and connect a UDP socket.
Candidates are tried in resolver order. A connected UDP socket accepts
datagrams only from its exact remote peer and lets the caller use
``send``/``recv`` instead of passing an address on every operation.
"""
if local_port is not None:
_validate_port(local_port, allow_zero=True)
if timeout is not None:
if isinstance(timeout, bool) or not isinstance(timeout, (int, float)):
raise TypeError('timeout must be a number or None')
if not math.isfinite(timeout) or timeout < 0:
raise ValueError('timeout must be a non-negative number or None')
endpoints = resolve_udp_endpoints(host, port, family=family)
for endpoint in endpoints:
sock = None
try:
sock = socket.socket(
endpoint.family, socket.SOCK_DGRAM, socket.IPPROTO_UDP)
if local_port is not None:
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
sock.bind(endpoint.bind_address(local_port))
sock.connect(endpoint.sockaddr)
sock.settimeout(timeout)
return sock, endpoint
except OSError:
if sock is not None:
try:
sock.close()
except OSError:
pass
raise EndpointError() from OSError('UDP socket setup failed')
+423
View File
@@ -0,0 +1,423 @@
import socket
import traceback
import pytest
from OpenSSL import SSL
from smartthings_local.errors import EndpointError
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.endpoint import (
ResolvedUdpEndpoint,
open_connected_udp_socket,
resolve_udp_endpoint,
resolve_udp_endpoints,
)
def _addrinfo(family, sockaddr, *, socktype=socket.SOCK_DGRAM,
protocol=socket.IPPROTO_UDP):
return family, socktype, protocol, '', sockaddr
class FakeSocket:
def __init__(self, family, *, fail_bind=False, fail_connect=False):
self.family = family
self.fail_bind = fail_bind
self.fail_connect = fail_connect
self.options = []
self.bound = None
self.peer = None
self.timeout = None
self.closed = False
self.sent = []
self.inbound = []
def setsockopt(self, *args):
self.options.append(args)
def bind(self, address):
if self.fail_bind:
raise OSError('synthetic bind failure')
self.bound = address
def connect(self, address):
if self.fail_connect:
raise OSError('synthetic connect failure')
self.peer = address
def settimeout(self, timeout):
self.timeout = timeout
def send(self, data):
self.sent.append(data)
return len(data)
def recv(self, _size):
return self.inbound.pop(0)
def close(self):
self.closed = True
def _socket_factory(monkeypatch, failures=()):
created = []
remaining = list(failures)
def factory(family, _socktype, _protocol):
failure = remaining.pop(0) if remaining else None
sock = FakeSocket(
family,
fail_bind=failure == 'bind',
fail_connect=failure == 'connect',
)
created.append(sock)
return sock
monkeypatch.setattr(
'smartthings_local.protocol.endpoint.socket.socket', factory)
return created
def test_resolve_ipv4_endpoint_and_redacted_repr(monkeypatch):
monkeypatch.setattr(
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
lambda *args: [_addrinfo(socket.AF_INET, ('192.0.2.10', 5684))],
)
endpoint = resolve_udp_endpoint('device.example', 5684)
assert endpoint.family == socket.AF_INET
assert endpoint.host == '192.0.2.10'
assert endpoint.port == 5684
assert endpoint.scope_id == 0
assert repr(endpoint) == 'ResolvedUdpEndpoint(family=AF_INET)'
assert '192.0.2.10' not in repr(endpoint)
assert '5684' not in repr(endpoint)
def test_resolve_scoped_ipv6_preserves_flow_and_scope(monkeypatch):
sockaddr = ('2001:db8::1', 5684, 3, 7)
calls = []
def resolve(*args):
calls.append(args)
return [_addrinfo(socket.AF_INET6, sockaddr)]
monkeypatch.setattr(
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
resolve,
)
endpoint = resolve_udp_endpoint(
'device.example', 5684, family=socket.AF_INET6)
assert endpoint.sockaddr == sockaddr
assert endpoint.scope_id == 7
assert endpoint.bind_address(55000) == ('::', 55000, 0, 0)
assert '2001:db8::1' not in repr(endpoint)
assert '7' not in repr(endpoint)
assert calls == [(
'device.example', 5684, socket.AF_INET6,
socket.SOCK_DGRAM, socket.IPPROTO_UDP)]
def test_resolver_order_is_stable_and_duplicates_are_removed(monkeypatch):
first = _addrinfo(socket.AF_INET6, ('2001:db8::10', 5684, 0, 0))
second = _addrinfo(socket.AF_INET, ('198.51.100.20', 5684))
ignored = _addrinfo(
socket.AF_INET, ('203.0.113.30', 5684), socktype=socket.SOCK_STREAM)
monkeypatch.setattr(
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
lambda *args: [first, first, ignored, second],
)
endpoints = resolve_udp_endpoints('device.example', 5684)
assert [endpoint.sockaddr for endpoint in endpoints] == [
first[-1], second[-1]]
def test_resolver_failure_raises_redacted_endpoint_error(monkeypatch):
def fail(*args):
raise socket.gaierror('credential-value at device.example')
monkeypatch.setattr(
'smartthings_local.protocol.endpoint.socket.getaddrinfo', fail)
remote_host = 'device.example'
with pytest.raises(EndpointError) as exc:
resolve_udp_endpoint(remote_host, 5684)
formatted = ''.join(traceback.format_exception(exc.value))
assert isinstance(exc.value, OSError)
assert exc.value.__context__ is None
assert 'UDP endpoint resolution failed' in formatted
assert 'credential-value' not in formatted
assert 'device.example' not in formatted
def test_socket_setup_tries_next_candidate_after_bind_failure(monkeypatch):
candidates = [
_addrinfo(socket.AF_INET6, ('2001:db8::10', 5684, 0, 0)),
_addrinfo(socket.AF_INET, ('192.0.2.10', 5684)),
]
monkeypatch.setattr(
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
lambda *args: candidates,
)
created = _socket_factory(monkeypatch, failures=('bind', None))
sock, endpoint = open_connected_udp_socket(
'device.example', 5684, local_port=55000, timeout=1.5)
assert created[0].closed
assert sock is created[1]
assert endpoint.family == socket.AF_INET
assert sock.bound == ('', 55000)
assert sock.peer == ('192.0.2.10', 5684)
assert sock.timeout == 1.5
assert (socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) in sock.options
def test_socket_setup_failure_is_redacted_and_closes_candidates(monkeypatch):
monkeypatch.setattr(
'smartthings_local.protocol.endpoint.socket.getaddrinfo',
lambda *args: [
_addrinfo(socket.AF_INET, ('192.0.2.10', 5684)),
_addrinfo(socket.AF_INET, ('198.51.100.20', 5684)),
],
)
created = _socket_factory(monkeypatch, failures=('connect', 'connect'))
with pytest.raises(EndpointError) as exc:
open_connected_udp_socket('device.example', 5684)
formatted = ''.join(traceback.format_exception(exc.value))
assert all(sock.closed for sock in created)
assert exc.value.__context__ is None
assert 'UDP socket setup failed' in formatted
assert 'synthetic connect failure' not in formatted
assert '192.0.2.10' not in formatted
assert '198.51.100.20' not in formatted
def test_same_port_different_hosts_remain_distinct_with_source_reuse(
monkeypatch):
addresses = {
'first.example': '192.0.2.10',
'second.example': '198.51.100.20',
}
def resolve(host, port, *_args):
return [_addrinfo(socket.AF_INET, (addresses[host], port))]
monkeypatch.setattr(
'smartthings_local.protocol.endpoint.socket.getaddrinfo', resolve)
created = _socket_factory(monkeypatch)
first, _ = open_connected_udp_socket(
'first.example', 5684, local_port=55000)
second, _ = open_connected_udp_socket(
'second.example', 5684, local_port=55000)
assert first.peer == ('192.0.2.10', 5684)
assert second.peer == ('198.51.100.20', 5684)
assert first.peer != second.peer
assert [sock.bound for sock in created] == [('', 55000), ('', 55000)]
def test_connected_udp_socket_filters_datagrams_from_another_peer():
expected_peer = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
other_peer = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
client = None
try:
expected_peer.bind(('127.0.0.1', 0))
other_peer.bind(('127.0.0.1', 0))
expected_peer.settimeout(0.5)
client, endpoint = open_connected_udp_socket(
'127.0.0.1',
expected_peer.getsockname()[1],
family=socket.AF_INET,
timeout=0.5,
)
assert endpoint.sockaddr == expected_peer.getsockname()
assert client.send(b'client datagram') == len(b'client datagram')
payload, source = expected_peer.recvfrom(64)
assert payload == b'client datagram'
assert source == client.getsockname()
other_peer.sendto(b'unrelated', client.getsockname())
expected_peer.sendto(b'expected', client.getsockname())
assert client.recv(64) == b'expected'
finally:
if client is not None:
client.close()
expected_peer.close()
other_peer.close()
@pytest.mark.parametrize(
('host', 'port', 'family'),
(
('', 5684, socket.AF_UNSPEC),
('device.example', 0, socket.AF_UNSPEC),
('device.example', 65536, socket.AF_UNSPEC),
('device.example', 5684, 9999),
),
)
def test_invalid_endpoint_inputs_fail_before_resolution(host, port, family):
with pytest.raises(ValueError):
resolve_udp_endpoint(host, port, family=family)
@pytest.mark.parametrize(
('host', 'port', 'family'),
(
(None, 5684, socket.AF_UNSPEC),
('device.example', '5684', socket.AF_UNSPEC),
('device.example', 5684, 'AF_INET'),
),
)
def test_endpoint_input_types_are_explicit(host, port, family):
with pytest.raises(TypeError):
resolve_udp_endpoint(host, port, family=family)
@pytest.mark.parametrize('timeout', (-1, float('nan'), float('inf')))
def test_socket_timeout_must_be_finite_and_non_negative(timeout):
with pytest.raises(ValueError):
open_connected_udp_socket('device.example', 5684, timeout=timeout)
def test_session_uses_connected_socket_send_and_recv(monkeypatch):
endpoint = ResolvedUdpEndpoint(
socket.AF_INET6, ('2001:db8::10', 5684, 0, 0))
sock = FakeSocket(socket.AF_INET6)
sock.peer = endpoint.sockaddr
sock.inbound.append(b'synthetic server flight')
outbound_record = (
b'\x16\xfe\xfd' + b'\x00' * 8 + b'\x00\x01' + b'x')
class FakeContext:
def load_verify_locations(self, *args):
pass
def set_verify(self, *args):
pass
def set_cipher_list(self, *args):
pass
def use_certificate_chain_file(self, *args):
pass
def use_privatekey_file(self, *args):
pass
def check_privatekey(self):
pass
class FakeConnection:
def __init__(self):
self.handshake_calls = 0
self.bio_reads = 0
self.bio_writes = []
def set_connect_state(self):
pass
def set_ciphertext_mtu(self, *args):
pass
def do_handshake(self):
self.handshake_calls += 1
if self.handshake_calls == 1:
raise SSL.WantReadError()
def bio_read(self, _size):
self.bio_reads += 1
if self.bio_reads == 1:
return outbound_record
raise SSL.WantReadError()
def bio_write(self, data):
self.bio_writes.append(data)
connection = FakeConnection()
open_calls = []
def open_socket(*args, **kwargs):
open_calls.append((args, kwargs))
return sock, endpoint
monkeypatch.setattr(dtls_session.SSL, 'Context', lambda *args: FakeContext())
monkeypatch.setattr(
dtls_session.SSL, 'Connection', lambda *args: connection)
monkeypatch.setattr(
dtls_session,
'open_connected_udp_socket',
open_socket,
)
monkeypatch.setattr(dtls_session.time, 'sleep', lambda _delay: None)
session = dtls_session.DtlsCoapSession(
'device.example', 5684,
cert_path='/synthetic/client.pem',
key_path='/synthetic/client.key',
family=socket.AF_INET6,
)
session.connect()
assert sock.sent == [outbound_record]
assert connection.bio_writes == [b'synthetic server flight']
assert session.endpoint is endpoint
assert session.dest == endpoint.sockaddr
assert open_calls == [(('device.example', 5684), {
'family': socket.AF_INET6,
'local_port': None,
'timeout': 2.0,
})]
session.close()
assert session.endpoint is None
assert session.dest is None
def test_session_send_failure_has_no_raw_exception_context():
outbound_record = (
b'\x16\xfe\xfd' + b'\x00' * 8 + b'\x00\x01' + b'x')
class FakeConnection:
def __init__(self):
self.bio_reads = 0
def send(self, _data):
pass
def bio_read(self, _size):
self.bio_reads += 1
if self.bio_reads == 1:
return outbound_record
raise SSL.WantReadError()
class FailingSocket:
def send(self, _data):
raise OSError('credential-value at device.example')
session = dtls_session.DtlsCoapSession(
'device.example', 5684,
cert_path='/synthetic/client.pem',
key_path='/synthetic/client.key',
)
session.conn = FakeConnection()
session.sock = FailingSocket()
with pytest.raises(EndpointError) as exc:
session._send_dgram(b'payload')
formatted = ''.join(traceback.format_exception(exc.value))
assert exc.value.__context__ is None
assert 'UDP send failed' in formatted
assert 'credential-value' not in formatted
assert 'device.example' not in formatted
+7 -1
View File
@@ -18,6 +18,7 @@ from smartthings_local.errors import (
)
from smartthings_local.protocol import dtls_session
from smartthings_local.protocol.dtls_session import DtlsCoapSession
from smartthings_local.protocol.endpoint import ResolvedUdpEndpoint
ERROR_TYPES = (
EndpointError,
@@ -124,8 +125,13 @@ def test_handshake_error_is_classified_without_backend_text(monkeypatch):
monkeypatch.setattr(dtls_session.SSL, 'Context', lambda *args: FakeContext())
monkeypatch.setattr(
dtls_session.SSL, 'Connection', lambda *args: FakeConnection())
endpoint = ResolvedUdpEndpoint(
dtls_session.socket.AF_INET, ('192.0.2.10', 5684))
monkeypatch.setattr(
dtls_session.socket, 'socket', lambda *args: fake_socket)
dtls_session,
'open_connected_udp_socket',
lambda *args, **kwargs: (fake_socket, endpoint),
)
session = DtlsCoapSession(
'device.example', 5684,