feat: add UDP NAT rendezvous

This commit is contained in:
Alexander Sellite 2026-08-10 19:49:47 -04:00
parent 6a7c7a676d
commit 804e19c51a
7 changed files with 370 additions and 27 deletions

View file

@ -9,6 +9,7 @@ import json
import logging
import os
import signal
import socketserver
import threading
from http import HTTPStatus
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
@ -30,6 +31,8 @@ from .registry import (
LOGGER = logging.getLogger("netfishing.discovery")
DEFAULT_MAX_BODY_BYTES = 8 * 1024
TRAVERSAL_PACKET_PREFIX = b"NETFISHING_TRAVERSAL_V1 "
MAX_TRAVERSAL_PACKET_BYTES = 2048
@dataclass(frozen=True, slots=True)
@ -41,6 +44,9 @@ class ServerConfig:
max_rooms_per_address: int = DEFAULT_MAX_ROOMS_PER_ADDRESS
max_body_bytes: int = DEFAULT_MAX_BODY_BYTES
trusted_proxy_cidrs: tuple[ipaddress.IPv4Network | ipaddress.IPv6Network, ...] = ()
traversal_bind_host: str = "0.0.0.0"
traversal_port: int = 7771
traversal_public_host: str = "127.0.0.1"
@classmethod
def from_environment(cls) -> "ServerConfig":
@ -67,6 +73,15 @@ class ServerConfig:
os.getenv("NETFISHING_DISCOVERY_MAX_BODY_BYTES", str(DEFAULT_MAX_BODY_BYTES))
),
trusted_proxy_cidrs=tuple(proxy_cidrs),
traversal_bind_host=os.getenv(
"NETFISHING_DISCOVERY_TRAVERSAL_HOST", "0.0.0.0"
),
traversal_port=int(
os.getenv("NETFISHING_DISCOVERY_TRAVERSAL_PORT", "7771")
),
traversal_public_host=os.getenv(
"NETFISHING_DISCOVERY_TRAVERSAL_PUBLIC_HOST", "127.0.0.1"
),
)
@ -101,6 +116,20 @@ class DiscoveryRequestHandler(BaseHTTPRequestHandler):
},
)
return
room_id = self._room_id_for_suffix(route.path, "/join-attempts")
if room_id is not None:
try:
endpoints = self.server.registry.consume_join_endpoints(
room_id, self._bearer_token()
)
except LeaseAuthorizationError as error:
self._send_error(HTTPStatus.UNAUTHORIZED, "invalid_lease", str(error))
return
except RoomNotFoundError as error:
self._send_error(HTTPStatus.NOT_FOUND, "room_not_found", str(error))
return
self._send_json(HTTPStatus.OK, {"endpoints": endpoints})
return
if route.path == "/v1/rooms":
try:
filters = parse_qs(route.query, keep_blank_values=False)
@ -127,14 +156,32 @@ class DiscoveryRequestHandler(BaseHTTPRequestHandler):
self._send_error(HTTPStatus.NOT_FOUND, "not_found", "route not found")
def do_POST(self) -> None:
if urlsplit(self.path).path != "/v1/rooms":
path = urlsplit(self.path).path
room_id = self._room_id_for_suffix(path, "/join-attempts")
if room_id is not None:
try:
token = self.server.registry.create_join_attempt(room_id)
except RoomNotFoundError as error:
self._send_error(HTTPStatus.NOT_FOUND, "room_not_found", str(error))
return
except RoomLimitError as error:
self._send_error(HTTPStatus.TOO_MANY_REQUESTS, "join_limit", str(error))
return
self._send_json(
HTTPStatus.CREATED,
{"join_token": token, "traversal": self._traversal_details()},
)
return
if path != "/v1/rooms":
self._send_error(HTTPStatus.NOT_FOUND, "not_found", "route not found")
return
payload = self._read_json_object()
if payload is None:
return
try:
room, token = self.server.registry.create(self._client_address(), payload)
room, token, verification_token = self.server.registry.create(
self._client_address(), payload
)
except ValidationError as error:
self._send_error(HTTPStatus.BAD_REQUEST, "invalid_room", str(error))
return
@ -143,7 +190,11 @@ class DiscoveryRequestHandler(BaseHTTPRequestHandler):
return
self._send_json(
HTTPStatus.CREATED,
{"room": room.public_dict(), "lease_token": token},
{
"room": room.public_dict(),
"lease_token": token,
"traversal": self._traversal_details(verification_token),
},
)
def do_PUT(self) -> None:
@ -251,6 +302,23 @@ class DiscoveryRequestHandler(BaseHTTPRequestHandler):
return None
return room_id
@staticmethod
def _room_id_for_suffix(path: str, suffix: str) -> str | None:
prefix = "/v1/rooms/"
if not path.startswith(prefix) or not path.endswith(suffix):
return None
room_id = path[len(prefix) : -len(suffix)]
return room_id if room_id and "/" not in room_id else None
def _traversal_details(self, verification_token: str = "") -> dict[str, Any]:
details: dict[str, Any] = {
"host": self.server.config.traversal_public_host,
"port": self.server.config.traversal_port,
}
if verification_token:
details["verification_token"] = verification_token
return details
def _bearer_token(self) -> str:
value = self.headers.get("Authorization", "")
scheme, separator, token = value.partition(" ")
@ -291,6 +359,51 @@ def _install_signal_handlers(server: BaseServer) -> None:
signal.signal(signal.SIGTERM, stop_server)
class TraversalUDPServer(socketserver.ThreadingUDPServer):
daemon_threads = True
allow_reuse_address = True
def __init__(self, config: ServerConfig, registry: RoomRegistry) -> None:
self.registry = registry
super().__init__(
(config.traversal_bind_host, config.traversal_port),
TraversalRequestHandler,
)
class TraversalRequestHandler(socketserver.BaseRequestHandler):
server: TraversalUDPServer
def handle(self) -> None:
packet = self.request[0]
if (
not isinstance(packet, bytes)
or len(packet) > MAX_TRAVERSAL_PACKET_BYTES
or not packet.startswith(TRAVERSAL_PACKET_PREFIX)
):
return
try:
payload = json.loads(packet[len(TRAVERSAL_PACKET_PREFIX) :])
except (json.JSONDecodeError, UnicodeDecodeError):
return
if not isinstance(payload, dict):
return
address = str(self.client_address[0])
port = int(self.client_address[1])
kind = payload.get("kind")
token = payload.get("token")
room_id = payload.get("room_id")
if not isinstance(token, str) or not isinstance(room_id, str):
return
try:
if kind == "host":
self.server.registry.verify_endpoint(room_id, token, address, port)
elif kind == "join":
self.server.registry.register_join_endpoint(token, address, port)
except (ValidationError, LeaseAuthorizationError, RoomNotFoundError):
return
def _argument_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--host", help="bind address (overrides environment)")
@ -311,12 +424,23 @@ def main() -> None:
max_rooms_per_address=config.max_rooms_per_address,
max_body_bytes=config.max_body_bytes,
trusted_proxy_cidrs=config.trusted_proxy_cidrs,
traversal_bind_host=config.traversal_bind_host,
traversal_port=config.traversal_port,
traversal_public_host=config.traversal_public_host,
)
logging.basicConfig(
level=getattr(logging, arguments.log_level.upper(), logging.INFO),
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
)
server = DiscoveryHTTPServer(config)
traversal_server = TraversalUDPServer(config, server.registry)
traversal_thread = threading.Thread(
target=traversal_server.serve_forever,
kwargs={"poll_interval": 0.25},
name="discovery-traversal",
daemon=True,
)
traversal_thread.start()
_install_signal_handlers(server)
LOGGER.info(
"NETfishing discovery server %s listening on %s:%d (room TTL %.1fs)",
@ -325,9 +449,17 @@ def main() -> None:
server.server_address[1],
config.room_ttl_seconds,
)
LOGGER.info(
"NETfishing traversal rendezvous listening on %s:%d/udp",
config.traversal_bind_host,
traversal_server.server_address[1],
)
try:
server.serve_forever(poll_interval=0.25)
finally:
traversal_server.shutdown()
traversal_server.server_close()
traversal_thread.join(timeout=2.0)
server.server_close()
LOGGER.info("server stopped")