feat: add UDP NAT rendezvous
This commit is contained in:
parent
6a7c7a676d
commit
804e19c51a
7 changed files with 370 additions and 27 deletions
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue