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

@ -2,7 +2,7 @@
from __future__ import annotations
from dataclasses import dataclass
from dataclasses import dataclass, replace
import hmac
import secrets
import threading
@ -18,6 +18,8 @@ VERSION_MAX_LENGTH = 32
MAX_ROOM_CAPACITY = 128
DEFAULT_MAX_ROOMS = 4096
DEFAULT_MAX_ROOMS_PER_ADDRESS = 32
JOIN_ATTEMPT_TTL_SECONDS = 12.0
MAX_JOIN_ATTEMPTS_PER_ROOM = 64
class RegistryError(Exception):
@ -53,6 +55,7 @@ class RoomAdvertisement:
created_at: float
updated_at: float
expires_at: float
verified: bool = False
def public_dict(self) -> dict[str, Any]:
return {
@ -67,6 +70,7 @@ class RoomAdvertisement:
"created_at": _iso_utc(self.created_at),
"updated_at": _iso_utc(self.updated_at),
"expires_at": _iso_utc(self.expires_at),
"verified": self.verified,
}
@ -74,6 +78,15 @@ class RoomAdvertisement:
class _RoomLease:
advertisement: RoomAdvertisement
token: str
verification_token: str
@dataclass(slots=True)
class _JoinAttempt:
room_id: str
expires_at: float
address: str = ""
port: int = 0
def _iso_utc(timestamp: float) -> str:
@ -122,16 +135,20 @@ class RoomRegistry:
self._max_rooms_per_address = int(max_rooms_per_address)
self._clock = clock
self._rooms: dict[str, _RoomLease] = {}
self._join_attempts: dict[str, _JoinAttempt] = {}
self._lock = threading.RLock()
@property
def ttl_seconds(self) -> float:
return self._ttl_seconds
def create(self, address: str, payload: Mapping[str, Any]) -> tuple[RoomAdvertisement, str]:
def create(
self, address: str, payload: Mapping[str, Any]
) -> tuple[RoomAdvertisement, str, str]:
now = self._clock()
room_id = str(uuid.uuid4())
token = secrets.token_urlsafe(32)
verification_token = secrets.token_urlsafe(32)
advertisement = self._build_advertisement(room_id, address, payload, now, now)
with self._lock:
self._purge_locked(now)
@ -144,8 +161,10 @@ class RoomRegistry:
)
if address_room_count >= self._max_rooms_per_address:
raise RoomLimitError("too many active rooms from this address")
self._rooms[room_id] = _RoomLease(advertisement, token)
return advertisement, token
self._rooms[room_id] = _RoomLease(
advertisement, token, verification_token
)
return advertisement, token, verification_token
def update(
self,
@ -160,10 +179,14 @@ class RoomRegistry:
lease = self._authorized_lease_locked(room_id, token)
advertisement = self._build_advertisement(
room_id,
address,
lease.advertisement.address if lease.advertisement.verified else address,
payload,
lease.advertisement.created_at,
now,
endpoint_port=(
lease.advertisement.port if lease.advertisement.verified else None
),
verified=lease.advertisement.verified,
)
lease.advertisement = advertisement
return advertisement
@ -184,7 +207,11 @@ class RoomRegistry:
now = self._clock()
with self._lock:
self._purge_locked(now)
rooms = [lease.advertisement for lease in self._rooms.values()]
rooms = [
lease.advertisement
for lease in self._rooms.values()
if lease.advertisement.verified
]
if game_version is not None:
rooms = [room for room in rooms if room.game_version == game_version]
if protocol_version is not None:
@ -197,6 +224,93 @@ class RoomRegistry:
self._purge_locked(now)
return len(self._rooms)
def verify_endpoint(
self,
room_id: str,
verification_token: str,
address: str,
port: int,
) -> RoomAdvertisement:
if not address or port < 1 or port > 65_535:
raise ValidationError("invalid observed UDP endpoint")
now = self._clock()
with self._lock:
self._purge_locked(now)
lease = self._rooms.get(room_id)
if lease is None:
raise RoomNotFoundError("room does not exist or its lease expired")
if not hmac.compare_digest(
lease.verification_token, verification_token
):
raise LeaseAuthorizationError("invalid endpoint verification token")
lease.advertisement = replace(
lease.advertisement,
address=address,
port=port,
updated_at=now,
expires_at=now + self._ttl_seconds,
verified=True,
)
return lease.advertisement
def create_join_attempt(self, room_id: str) -> str:
now = self._clock()
with self._lock:
self._purge_locked(now)
lease = self._rooms.get(room_id)
if lease is None or not lease.advertisement.verified:
raise RoomNotFoundError("room is not available")
active_count = sum(
1
for attempt in self._join_attempts.values()
if attempt.room_id == room_id
)
if active_count >= MAX_JOIN_ATTEMPTS_PER_ROOM:
raise RoomLimitError("too many pending joins for this room")
token = secrets.token_urlsafe(32)
self._join_attempts[token] = _JoinAttempt(
room_id=room_id,
expires_at=now + JOIN_ATTEMPT_TTL_SECONDS,
)
return token
def register_join_endpoint(
self, token: str, address: str, port: int
) -> None:
if not address or port < 1 or port > 65_535:
raise ValidationError("invalid observed join endpoint")
now = self._clock()
with self._lock:
self._purge_locked(now)
attempt = self._join_attempts.get(token)
if attempt is None:
raise RoomNotFoundError("join attempt does not exist or expired")
attempt.address = address
attempt.port = port
def consume_join_endpoints(
self, room_id: str, lease_token: str
) -> list[dict[str, Any]]:
now = self._clock()
with self._lock:
self._purge_locked(now)
self._authorized_lease_locked(room_id, lease_token)
consumed_tokens = [
token
for token, attempt in self._join_attempts.items()
if attempt.room_id == room_id and attempt.address and attempt.port > 0
]
endpoints = [
{
"address": self._join_attempts[token].address,
"port": self._join_attempts[token].port,
}
for token in consumed_tokens
]
for token in consumed_tokens:
del self._join_attempts[token]
return endpoints
def _build_advertisement(
self,
room_id: str,
@ -204,9 +318,12 @@ class RoomRegistry:
payload: Mapping[str, Any],
created_at: float,
now: float,
endpoint_port: int | None = None,
verified: bool = False,
) -> RoomAdvertisement:
room_name = _clean_text(payload.get("room_name"), "room_name", ROOM_NAME_MAX_LENGTH)
port = _clean_int(payload.get("port"), "port", 1, 65_535)
supplied_port = _clean_int(payload.get("port"), "port", 1, 65_535)
port = endpoint_port if endpoint_port is not None else supplied_port
current_players = _clean_int(
payload.get("current_players"), "current_players", 0, MAX_ROOM_CAPACITY
)
@ -233,6 +350,7 @@ class RoomRegistry:
created_at=created_at,
updated_at=now,
expires_at=now + self._ttl_seconds,
verified=verified,
)
def _authorized_lease_locked(self, room_id: str, token: str) -> _RoomLease:
@ -251,3 +369,10 @@ class RoomRegistry:
]
for room_id in expired:
del self._rooms[room_id]
expired_attempts = [
token
for token, attempt in self._join_attempts.items()
if attempt.expires_at <= now or attempt.room_id not in self._rooms
]
for token in expired_attempts:
del self._join_attempts[token]

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")