diff --git a/app/routes/mcp.py b/app/routes/mcp.py index e50a529..143c1a3 100644 --- a/app/routes/mcp.py +++ b/app/routes/mcp.py @@ -2,13 +2,13 @@ import json import logging import contextvars from fastapi import APIRouter, Depends, Request, HTTPException -from sse_starlette.sse import EventSourceResponse +from fastapi.responses import Response, JSONResponse from mcp.server import Server from mcp.server.sse import SseServerTransport from mcp.types import Tool, TextContent from app.auth import get_current_agent from app.scopes import has_scope_access, ALL_SCOPES -from app.models import get_rule, get_all_rules, get_all_ports, get_port, set_rule +from app.models import get_rule, get_all_rules, get_all_ports, get_port, set_rule, get_agent_by_key, update_key_last_used logger = logging.getLogger(__name__) @@ -165,30 +165,38 @@ async def handle_call_tool(name: str, arguments: dict) -> list[TextContent]: return [TextContent(type="text", text=json.dumps({"error": "Unknown tool"}))] # FastMCP / SSE Integration -# The Python MCP SDK uses SseServerTransport. We need a global dictionary to hold transports. -sse_transports = {} +# connect_sse() ET handle_post_message() envoient chacun leur reponse ASGI +# completement par eux-memes (via le send() qu'on leur passe). Les faire +# passer par des routes FastAPI classiques (Depends + return Response) fait +# que FastAPI tente un second envoi une fois le SDK termine -> RuntimeError +# uvicorn ("Unexpected ASGI message 'http.response.start' sent, after response +# already completed"). Un seul montage ASGI brut gere GET (ouverture SSE) et +# POST (messages) sans jamais repasser par le wrapping FastAPI ; l'auth +# X-API-Key est donc verifiee ici a la main plutot que via Depends. +sse_transport = SseServerTransport("/messages") -@router.get("/mcp") -async def mcp_sse(request: Request, agent: str = Depends(get_current_agent)): - transport = SseServerTransport("/mcp/messages") - sse_transports[agent] = transport - - async def run_server(): +async def _mcp_asgi_app(scope, receive, send): + if scope["type"] != "http": + return + + headers = dict(scope.get("headers") or []) + api_key = headers.get(b"x-api-key", b"").decode() + agent = await get_agent_by_key(api_key) if api_key else None + if not agent: + response = JSONResponse({"detail": "Invalid or missing API Key"}, status_code=401) + await response(scope, receive, send) + return + await update_key_last_used(api_key) + + if scope["method"] == "POST": + await sse_transport.handle_post_message(scope, receive, send) + return + + async with sse_transport.connect_sse(scope, receive, send) as (read_stream, write_stream): token = current_agent_var.set(agent) try: - await mcp_server.run(transport.read_stream(), transport.write_stream(), mcp_server.create_initialization_options()) + await mcp_server.run(read_stream, write_stream, mcp_server.create_initialization_options()) finally: current_agent_var.reset(token) - - import asyncio - asyncio.create_task(run_server()) - - return EventSourceResponse(transport.handle_sse(request)) -@router.post("/mcp/messages") -async def mcp_messages(request: Request, agent: str = Depends(get_current_agent)): - transport = sse_transports.get(agent) - if not transport: - raise HTTPException(status_code=400, detail="SSE connection not found") - await transport.handle_post_message(request.scope, request.receive, request._send) - return {} +router.mount("/mcp", _mcp_asgi_app)