feat: implement project bus MVP
This commit is contained in:
@@ -0,0 +1,136 @@
|
||||
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()}
|
||||
|
||||
Reference in New Issue
Block a user