netfishing-discovery-server/netfishing_discovery/server.py

507 lines
19 KiB
Python
Raw Normal View History

2026-08-10 10:19:57 -04:00
"""Small HTTP/JSON API for NETfishing public-room discovery."""
from __future__ import annotations
import argparse
from dataclasses import dataclass
import ipaddress
import json
import logging
import os
import signal
2026-08-10 19:49:47 -04:00
import socketserver
2026-08-10 10:19:57 -04:00
import threading
from http import HTTPStatus
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from socketserver import BaseServer
from typing import Any
from urllib.parse import parse_qs, urlsplit
from . import __version__
from .registry import (
DEFAULT_MAX_ROOMS,
DEFAULT_MAX_ROOMS_PER_ADDRESS,
LeaseAuthorizationError,
RoomLimitError,
RoomNotFoundError,
RoomRegistry,
ValidationError,
)
LOGGER = logging.getLogger("netfishing.discovery")
DEFAULT_MAX_BODY_BYTES = 8 * 1024
2026-08-10 19:49:47 -04:00
TRAVERSAL_PACKET_PREFIX = b"NETFISHING_TRAVERSAL_V1 "
MAX_TRAVERSAL_PACKET_BYTES = 2048
2026-08-10 10:19:57 -04:00
@dataclass(frozen=True, slots=True)
class ServerConfig:
bind_host: str = "127.0.0.1"
bind_port: int = 7770
room_ttl_seconds: float = 45.0
max_rooms: int = DEFAULT_MAX_ROOMS
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, ...] = ()
2026-08-10 19:49:47 -04:00
traversal_bind_host: str = "0.0.0.0"
traversal_port: int = 7771
traversal_public_host: str = "127.0.0.1"
2026-08-11 10:37:03 -04:00
build_revision: str = "unknown"
2026-08-10 10:19:57 -04:00
@classmethod
def from_environment(cls) -> "ServerConfig":
proxy_cidrs: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = []
raw_cidrs = os.getenv("NETFISHING_DISCOVERY_TRUSTED_PROXY_CIDRS", "")
for value in raw_cidrs.split(","):
value = value.strip()
if value:
proxy_cidrs.append(ipaddress.ip_network(value, strict=False))
return cls(
bind_host=os.getenv("NETFISHING_DISCOVERY_HOST", "127.0.0.1"),
bind_port=int(os.getenv("NETFISHING_DISCOVERY_PORT", "7770")),
room_ttl_seconds=float(os.getenv("NETFISHING_DISCOVERY_ROOM_TTL", "45")),
max_rooms=int(
os.getenv("NETFISHING_DISCOVERY_MAX_ROOMS", str(DEFAULT_MAX_ROOMS))
),
max_rooms_per_address=int(
os.getenv(
"NETFISHING_DISCOVERY_MAX_ROOMS_PER_ADDRESS",
str(DEFAULT_MAX_ROOMS_PER_ADDRESS),
)
),
max_body_bytes=int(
os.getenv("NETFISHING_DISCOVERY_MAX_BODY_BYTES", str(DEFAULT_MAX_BODY_BYTES))
),
trusted_proxy_cidrs=tuple(proxy_cidrs),
2026-08-10 19:49:47 -04:00
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"
),
2026-08-11 10:37:03 -04:00
build_revision=(
os.getenv("NETFISHING_DISCOVERY_BUILD_REVISION", "unknown").strip()
or "unknown"
)[:64],
2026-08-10 10:19:57 -04:00
)
class DiscoveryHTTPServer(ThreadingHTTPServer):
daemon_threads = True
allow_reuse_address = True
def __init__(self, config: ServerConfig, registry: RoomRegistry | None = None) -> None:
self.config = config
self.registry = registry or RoomRegistry(
config.room_ttl_seconds,
config.max_rooms,
config.max_rooms_per_address,
)
super().__init__((config.bind_host, config.bind_port), DiscoveryRequestHandler)
class DiscoveryRequestHandler(BaseHTTPRequestHandler):
server: DiscoveryHTTPServer
protocol_version = "HTTP/1.1"
def do_GET(self) -> None:
route = urlsplit(self.path)
if route.path == "/health":
self._send_json(
HTTPStatus.OK,
{
"status": "ok",
"service": "netfishing-discovery-server",
"version": __version__,
2026-08-11 10:37:03 -04:00
"build_revision": self.server.config.build_revision,
2026-08-10 10:19:57 -04:00
"active_rooms": self.server.registry.room_count(),
},
)
return
2026-08-10 19:49:47 -04:00
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
2026-08-10 10:19:57 -04:00
if route.path == "/v1/rooms":
try:
filters = parse_qs(route.query, keep_blank_values=False)
game_version = self._single_query_value(filters, "game_version")
raw_protocol = self._single_query_value(filters, "protocol_version")
protocol_version = int(raw_protocol) if raw_protocol is not None else None
if protocol_version is not None and protocol_version <= 0:
raise ValueError
except ValueError:
self._send_error(HTTPStatus.BAD_REQUEST, "invalid_query", "invalid room filter")
return
rooms = self.server.registry.list_rooms(
game_version=game_version,
protocol_version=protocol_version,
)
self._send_json(
HTTPStatus.OK,
{
"rooms": [room.public_dict() for room in rooms],
"ttl_seconds": self.server.registry.ttl_seconds,
},
)
return
self._send_error(HTTPStatus.NOT_FOUND, "not_found", "route not found")
def do_POST(self) -> None:
2026-08-10 19:49:47 -04:00
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":
2026-08-10 10:19:57 -04:00
self._send_error(HTTPStatus.NOT_FOUND, "not_found", "route not found")
return
payload = self._read_json_object()
if payload is None:
return
2026-08-11 20:31:02 -04:00
if self._reject_incompatible_game_version(payload):
return
2026-08-10 10:19:57 -04:00
try:
2026-08-10 19:49:47 -04:00
room, token, verification_token = self.server.registry.create(
self._client_address(), payload
)
2026-08-10 10:19:57 -04:00
except ValidationError as error:
self._send_error(HTTPStatus.BAD_REQUEST, "invalid_room", str(error))
return
except RoomLimitError as error:
self._send_error(HTTPStatus.TOO_MANY_REQUESTS, "room_limit", str(error))
return
self._send_json(
HTTPStatus.CREATED,
2026-08-10 19:49:47 -04:00
{
"room": room.public_dict(),
"lease_token": token,
"traversal": self._traversal_details(verification_token),
},
2026-08-10 10:19:57 -04:00
)
def do_PUT(self) -> None:
room_id = self._room_id_from_path()
if room_id is None:
self._send_error(HTTPStatus.NOT_FOUND, "not_found", "route not found")
return
payload = self._read_json_object()
if payload is None:
return
2026-08-11 20:31:02 -04:00
if self._reject_incompatible_game_version(payload):
return
2026-08-10 10:19:57 -04:00
try:
room = self.server.registry.update(
room_id,
self._bearer_token(),
self._client_address(),
payload,
)
except ValidationError as error:
self._send_error(HTTPStatus.BAD_REQUEST, "invalid_room", str(error))
return
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, {"room": room.public_dict()})
def do_DELETE(self) -> None:
room_id = self._room_id_from_path()
if room_id is None:
self._send_error(HTTPStatus.NOT_FOUND, "not_found", "route not found")
return
try:
self.server.registry.delete(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_response(HTTPStatus.NO_CONTENT)
self.send_header("Content-Length", "0")
self.end_headers()
def log_message(self, format_string: str, *args: object) -> None:
LOGGER.info("%s - %s", self.client_address[0], format_string % args)
def _read_json_object(self) -> dict[str, Any] | None:
content_type = self.headers.get("Content-Type", "").split(";", 1)[0].strip().lower()
if content_type != "application/json":
self._send_error(
HTTPStatus.UNSUPPORTED_MEDIA_TYPE,
"unsupported_media_type",
"Content-Type must be application/json",
)
return None
try:
content_length = int(self.headers.get("Content-Length", ""))
except ValueError:
content_length = -1
if content_length < 0:
self._send_error(
HTTPStatus.LENGTH_REQUIRED,
"length_required",
"Content-Length required",
)
return None
if content_length > self.server.config.max_body_bytes:
self._send_error(
HTTPStatus.REQUEST_ENTITY_TOO_LARGE,
"body_too_large",
"request too large",
)
return None
try:
payload = json.loads(self.rfile.read(content_length))
except (json.JSONDecodeError, UnicodeDecodeError):
self._send_error(HTTPStatus.BAD_REQUEST, "invalid_json", "invalid JSON body")
return None
if not isinstance(payload, dict):
self._send_error(HTTPStatus.BAD_REQUEST, "invalid_json", "JSON body must be an object")
return None
return payload
def _client_address(self) -> str:
peer_text = self.client_address[0]
try:
peer = ipaddress.ip_address(peer_text)
except ValueError:
return peer_text
if not any(peer in network for network in self.server.config.trusted_proxy_cidrs):
return peer_text
forwarded = self.headers.get("X-Forwarded-For", "").split(",", 1)[0].strip()
try:
return str(ipaddress.ip_address(forwarded))
except ValueError:
return peer_text
2026-08-11 20:31:02 -04:00
def _reject_incompatible_game_version(self, payload: dict[str, Any]) -> bool:
game_version = payload.get("game_version")
if not isinstance(game_version, str) or game_version.strip() == __version__:
return False
self._send_error(
HTTPStatus.CONFLICT,
"game_version_mismatch",
(
f"This discovery server only lists NETfishing {__version__} rooms. "
"Your room will not be listed until the game and discovery server "
"use the same version."
),
{"required_game_version": __version__},
)
return True
2026-08-10 10:19:57 -04:00
def _room_id_from_path(self) -> str | None:
path = urlsplit(self.path).path
prefix = "/v1/rooms/"
room_id = path[len(prefix) :] if path.startswith(prefix) else ""
if not room_id or "/" in room_id:
return None
return room_id
2026-08-10 19:49:47 -04:00
@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
2026-08-10 10:19:57 -04:00
def _bearer_token(self) -> str:
value = self.headers.get("Authorization", "")
scheme, separator, token = value.partition(" ")
if separator and scheme.lower() == "bearer":
return token.strip()
return ""
@staticmethod
def _single_query_value(filters: dict[str, list[str]], key: str) -> str | None:
values = filters.get(key)
if not values:
return None
if len(values) != 1:
raise ValueError
return values[0]
2026-08-11 20:31:02 -04:00
def _send_error(
self,
status: HTTPStatus,
code: str,
message: str,
details: dict[str, Any] | None = None,
) -> None:
error: dict[str, Any] = {"code": code, "message": message}
if details is not None:
error.update(details)
self._send_json(status, {"error": error})
2026-08-10 10:19:57 -04:00
def _send_json(self, status: HTTPStatus, payload: dict[str, Any]) -> None:
body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
self.send_response(status)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.send_header("Cache-Control", "no-store")
self.send_header("X-Content-Type-Options", "nosniff")
self.end_headers()
self.wfile.write(body)
def _install_signal_handlers(server: BaseServer) -> None:
def stop_server(_signum: int, _frame: object) -> None:
LOGGER.info("shutdown requested")
# BaseServer.shutdown() must run on a different thread from serve_forever().
threading.Thread(target=server.shutdown, name="discovery-shutdown", daemon=True).start()
signal.signal(signal.SIGINT, stop_server)
signal.signal(signal.SIGTERM, stop_server)
2026-08-11 10:37:03 -04:00
class TraversalUDPServer(socketserver.UDPServer):
# Each packet performs one bounded JSON decode and one locked registry
# operation. A serial loop avoids creating an attacker-controlled thread
# for every untrusted datagram.
2026-08-10 19:49:47 -04:00
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
2026-08-10 10:19:57 -04:00
def _argument_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--host", help="bind address (overrides environment)")
parser.add_argument("--port", type=int, help="bind port (overrides environment)")
parser.add_argument("--log-level", default=os.getenv("NETFISHING_DISCOVERY_LOG_LEVEL", "INFO"))
return parser
def main() -> None:
arguments = _argument_parser().parse_args()
config = ServerConfig.from_environment()
if arguments.host is not None or arguments.port is not None:
config = ServerConfig(
bind_host=arguments.host or config.bind_host,
bind_port=arguments.port or config.bind_port,
room_ttl_seconds=config.room_ttl_seconds,
max_rooms=config.max_rooms,
max_rooms_per_address=config.max_rooms_per_address,
max_body_bytes=config.max_body_bytes,
trusted_proxy_cidrs=config.trusted_proxy_cidrs,
2026-08-10 19:49:47 -04:00
traversal_bind_host=config.traversal_bind_host,
traversal_port=config.traversal_port,
traversal_public_host=config.traversal_public_host,
2026-08-11 10:37:03 -04:00
build_revision=config.build_revision,
2026-08-10 10:19:57 -04:00
)
logging.basicConfig(
level=getattr(logging, arguments.log_level.upper(), logging.INFO),
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
)
server = DiscoveryHTTPServer(config)
2026-08-10 19:49:47 -04:00
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()
2026-08-10 10:19:57 -04:00
_install_signal_handlers(server)
LOGGER.info(
"NETfishing discovery server %s listening on %s:%d (room TTL %.1fs)",
__version__,
config.bind_host,
server.server_address[1],
config.room_ttl_seconds,
)
2026-08-10 19:49:47 -04:00
LOGGER.info(
"NETfishing traversal rendezvous listening on %s:%d/udp",
config.traversal_bind_host,
traversal_server.server_address[1],
)
2026-08-10 10:19:57 -04:00
try:
server.serve_forever(poll_interval=0.25)
finally:
2026-08-10 19:49:47 -04:00
traversal_server.shutdown()
traversal_server.server_close()
traversal_thread.join(timeout=2.0)
2026-08-10 10:19:57 -04:00
server.server_close()
LOGGER.info("server stopped")
if __name__ == "__main__":
main()