feat(protocol): add resolved connected UDP endpoints
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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')
|
||||
Reference in New Issue
Block a user