from __future__ import annotations import json import logging import os from http import HTTPStatus from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Any from urllib.parse import urlsplit from .auth import AuthProvider from .errors import AuthenticationError from .mcp import MCPApplication LOGGER = logging.getLogger("project_bus.http") MAX_BODY_BYTES = 1_048_576 class ProjectBusHTTPServer(ThreadingHTTPServer): daemon_threads = True def __init__( self, server_address: tuple[str, int], application: MCPApplication, auth_provider: AuthProvider, allowed_origins: set[str], ): super().__init__(server_address, ProjectBusRequestHandler) self.application = application self.auth_provider = auth_provider self.allowed_origins = allowed_origins class ProjectBusRequestHandler(BaseHTTPRequestHandler): server: ProjectBusHTTPServer protocol_version = "HTTP/1.1" def log_message(self, format_string: str, *args: Any) -> None: LOGGER.info("%s - %s", self.address_string(), format_string % args) def do_GET(self) -> None: path = urlsplit(self.path).path if path == "/healthz": self._json(HTTPStatus.OK, {"status": "ok"}) return if path == "/mcp": self._json( HTTPStatus.METHOD_NOT_ALLOWED, {"error": "This stateless server does not expose an SSE GET stream"}, extra_headers={"Allow": "POST"}, ) return self._json(HTTPStatus.NOT_FOUND, {"error": "not_found"}) def do_POST(self) -> None: if urlsplit(self.path).path != "/mcp": self._json(HTTPStatus.NOT_FOUND, {"error": "not_found"}) return if not self._origin_allowed(): self._json(HTTPStatus.FORBIDDEN, {"error": "origin_not_allowed"}) return content_type = self.headers.get("Content-Type", "").split(";", 1)[0].strip() if content_type != "application/json": self._json(HTTPStatus.UNSUPPORTED_MEDIA_TYPE, {"error": "application/json required"}) return try: length = int(self.headers.get("Content-Length", "0")) except ValueError: self._json(HTTPStatus.BAD_REQUEST, {"error": "invalid_content_length"}) return if length <= 0 or length > MAX_BODY_BYTES: self._json(HTTPStatus.REQUEST_ENTITY_TOO_LARGE, {"error": "invalid_body_size"}) return try: request = json.loads(self.rfile.read(length)) except (json.JSONDecodeError, UnicodeDecodeError): self._json( HTTPStatus.OK, {"jsonrpc": "2.0", "id": None, "error": {"code": -32700, "message": "Parse error"}}, ) return if not isinstance(request, dict): self._json( HTTPStatus.OK, {"jsonrpc": "2.0", "id": None, "error": {"code": -32600, "message": "Invalid Request"}}, ) return principal = None if request.get("method") == "tools/call": try: principal = self.server.auth_provider.authenticate(self.headers.get("Authorization")) except AuthenticationError as error: self._json( HTTPStatus.UNAUTHORIZED, {"error": error.code, "message": error.message}, extra_headers={"WWW-Authenticate": 'Bearer realm="project-bus"'}, ) return response = self.server.application.handle(request, principal) if response is None: self.send_response(HTTPStatus.ACCEPTED) self.send_header("Content-Length", "0") self.end_headers() return self._json(HTTPStatus.OK, response) def _origin_allowed(self) -> bool: origin = self.headers.get("Origin") if not origin: return True return origin in self.server.allowed_origins def _json( self, status: HTTPStatus, payload: dict[str, Any], *, extra_headers: dict[str, str] | None = None, ) -> 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") for name, value in (extra_headers or {}).items(): self.send_header(name, value) self.end_headers() self.wfile.write(body) def allowed_origins_from_environment() -> set[str]: raw = os.environ.get("PROJECT_BUS_ALLOWED_ORIGINS", "") return {item.strip() for item in raw.split(",") if item.strip()}