fix(mcp): endpoint /mcp casse (SDK 1.28 connect_sse/handle_post_message envoient leur propre reponse ASGI) -- montage ASGI unique GET+POST, plus de double-envoi, teste end-to-end
This commit is contained in:
+30
-22
@@ -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 _mcp_asgi_app(scope, receive, send):
|
||||
if scope["type"] != "http":
|
||||
return
|
||||
|
||||
async def run_server():
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user