netfishing-discovery-server/netfishing_discovery/server.py

506 lines
19 KiB
Python

"""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
import socketserver
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
TRAVERSAL_PACKET_PREFIX = b"NETFISHING_TRAVERSAL_V1 "
MAX_TRAVERSAL_PACKET_BYTES = 2048
@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, ...] = ()
traversal_bind_host: str = "0.0.0.0"
traversal_port: int = 7771
traversal_public_host: str = "127.0.0.1"
build_revision: str = "unknown"
@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),
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"
),
build_revision=(
os.getenv("NETFISHING_DISCOVERY_BUILD_REVISION", "unknown").strip()
or "unknown"
)[:64],
)
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__,
"build_revision": self.server.config.build_revision,
"active_rooms": self.server.registry.room_count(),
},
)
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)
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:
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
if self._reject_incompatible_game_version(payload):
return
try:
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
except RoomLimitError as error:
self._send_error(HTTPStatus.TOO_MANY_REQUESTS, "room_limit", str(error))
return
self._send_json(
HTTPStatus.CREATED,
{
"room": room.public_dict(),
"lease_token": token,
"traversal": self._traversal_details(verification_token),
},
)
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
if self._reject_incompatible_game_version(payload):
return
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
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
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
@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(" ")
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]
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})
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)
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.
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)")
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,
traversal_bind_host=config.traversal_bind_host,
traversal_port=config.traversal_port,
traversal_public_host=config.traversal_public_host,
build_revision=config.build_revision,
)
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)",
__version__,
config.bind_host,
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")
if __name__ == "__main__":
main()