"""Small HTTP/JSON API for straywild 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, ) from .social_registry import FriendOfflineError, SocialRegistry LOGGER = logging.getLogger("straywild.discovery") DEFAULT_MAX_BODY_BYTES = 64 * 1024 TRAVERSAL_PACKET_PREFIX = b"straywild_TRAVERSAL_V1 " MAX_TRAVERSAL_PACKET_BYTES = 2048 def _environment(primary: str, legacy: str, default: str) -> str: """Read the Straywild setting first, then its NETfishing predecessor.""" if primary in os.environ: return os.environ[primary] return os.getenv(legacy, default) @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 = _environment( "straywild_DISCOVERY_TRUSTED_PROXY_CIDRS", "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=_environment( "straywild_DISCOVERY_HOST", "NETFISHING_DISCOVERY_HOST", "127.0.0.1" ), bind_port=int(_environment( "straywild_DISCOVERY_PORT", "NETFISHING_DISCOVERY_PORT", "7770" )), room_ttl_seconds=float(_environment( "straywild_DISCOVERY_ROOM_TTL", "NETFISHING_DISCOVERY_ROOM_TTL", "45", )), max_rooms=int( _environment( "straywild_DISCOVERY_MAX_ROOMS", "NETFISHING_DISCOVERY_MAX_ROOMS", str(DEFAULT_MAX_ROOMS), ) ), max_rooms_per_address=int( _environment( "straywild_DISCOVERY_MAX_ROOMS_PER_ADDRESS", "NETFISHING_DISCOVERY_MAX_ROOMS_PER_ADDRESS", str(DEFAULT_MAX_ROOMS_PER_ADDRESS), ) ), max_body_bytes=int( _environment( "straywild_DISCOVERY_MAX_BODY_BYTES", "NETFISHING_DISCOVERY_MAX_BODY_BYTES", str(DEFAULT_MAX_BODY_BYTES), ) ), trusted_proxy_cidrs=tuple(proxy_cidrs), traversal_bind_host=_environment( "straywild_DISCOVERY_TRAVERSAL_HOST", "NETFISHING_DISCOVERY_TRAVERSAL_HOST", "0.0.0.0", ), traversal_port=int( _environment( "straywild_DISCOVERY_TRAVERSAL_PORT", "NETFISHING_DISCOVERY_TRAVERSAL_PORT", "7771", ) ), traversal_public_host=_environment( "straywild_DISCOVERY_TRAVERSAL_PUBLIC_HOST", "NETFISHING_DISCOVERY_TRAVERSAL_PUBLIC_HOST", "127.0.0.1", ), build_revision=( _environment( "straywild_DISCOVERY_BUILD_REVISION", "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, social_registry: SocialRegistry | None = None, ) -> None: self.config = config self.registry = registry or RoomRegistry( config.room_ttl_seconds, config.max_rooms, config.max_rooms_per_address, ) self.social_registry = social_registry or SocialRegistry() 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": presence_count, invitation_count = self.server.social_registry.counts() self._send_json( HTTPStatus.OK, { "status": "ok", "service": "straywild-discovery-server", "version": __version__, "build_revision": self.server.config.build_revision, "active_rooms": self.server.registry.room_count(), "active_presence_channels": presence_count, "pending_invitations": invitation_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 if path in [ "/v1/presence", "/v1/presence/query", "/v1/invitations", "/v1/invitations/poll", ]: self._handle_social_post(path) return 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 _handle_social_post(self, path: str) -> None: payload = self._read_json_object() if payload is None or self._reject_incompatible_game_version(payload): return try: game_version = payload.get("game_version") protocol_version = payload.get("protocol_version") if not isinstance(game_version, str): raise ValidationError("game_version must be a string") if ( isinstance(protocol_version, bool) or not isinstance(protocol_version, int) or protocol_version < 1 ): raise ValidationError("protocol_version must be a positive integer") if path == "/v1/presence": channels = self.server.social_registry.publish_presence( self._client_address(), payload ) self._send_json(HTTPStatus.OK, {"channels": channels}) return if path == "/v1/presence/query": presence = self.server.social_registry.query_presence( payload.get("channels"), game_version.strip(), protocol_version, self.server.registry, ) self._send_json(HTTPStatus.OK, {"presence": presence}) return if path == "/v1/invitations/poll": invitations = self.server.social_registry.poll_invitations( payload.get("inbox_tokens"), game_version.strip(), protocol_version, self.server.registry, ) self._send_json(HTTPStatus.OK, {"invitations": invitations}) return invitation = self.server.social_registry.send_invitation( payload.get("inbox_token"), payload.get("room_id", ""), game_version.strip(), protocol_version, self.server.registry, ) self._send_json( HTTPStatus.CREATED, {"invite_id": invitation.invite_id}, ) except FriendOfflineError as error: self._send_error(HTTPStatus.CONFLICT, "friend_offline", str(error)) except RoomNotFoundError as error: self._send_error(HTTPStatus.NOT_FOUND, "room_not_found", str(error)) except ValidationError as error: self._send_error(HTTPStatus.BAD_REQUEST, "invalid_social_request", str(error)) 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 straywild {__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=_environment( "straywild_DISCOVERY_LOG_LEVEL", "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( "straywild discovery server %s listening on %s:%d (room TTL %.1fs)", __version__, config.bind_host, server.server_address[1], config.room_ttl_seconds, ) LOGGER.info( "straywild 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()