Files
context-hub/app/routes/mcp.py
T

195 lines
7.9 KiB
Python

import json
import logging
import contextvars
from fastapi import APIRouter, Depends, Request, HTTPException
from sse_starlette.sse import EventSourceResponse
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
logger = logging.getLogger(__name__)
router = APIRouter()
mcp_server = Server("context-hub")
current_agent_var: contextvars.ContextVar[str] = contextvars.ContextVar("current_agent", default="")
@mcp_server.list_tools()
async def handle_list_tools() -> list[Tool]:
return [
Tool(
name="get_rules",
description="Get rules for a specific scope if authorized",
inputSchema={
"type": "object",
"properties": {
"scope": {"type": "string", "description": "Scope name (infra, llm, nyora, perso, tt)"}
},
"required": ["scope"]
}
),
Tool(
name="get_ports",
description="Get the full ports registry",
inputSchema={
"type": "object",
"properties": {}
}
),
Tool(
name="search_rules",
description="Search rules across all authorized scopes",
inputSchema={
"type": "object",
"properties": {
"query": {"type": "string", "description": "Search term"}
},
"required": ["query"]
}
),
Tool(
name="get_agent_config",
description="Get the compiled configuration for a specific agent based on its scopes",
inputSchema={
"type": "object",
"properties": {
"name": {"type": "string", "description": "Agent name (e.g., GEMINI, HERMES_TT)"}
},
"required": ["name"]
}
),
Tool(
name="check_port",
description="Check if a port is free or occupied",
inputSchema={
"type": "object",
"properties": {
"port": {"type": "integer", "description": "Port number"}
},
"required": ["port"]
}
),
Tool(
name="get_my_context",
description="Get all scopes authorized for this agent — call at START of session",
inputSchema={"type": "object", "properties": {"agent": {"type": "string"}}, "required": ["agent"]}
),
Tool(
name="update_rule",
description="Merge-update a scope with new key-value data — call at END of session",
inputSchema={
"type": "object",
"properties": {
"scope": {"type": "string", "description": "Scope to update (infra/llm/nyora/perso/tt)"},
"agent": {"type": "string", "description": "Agent name for scope access check"},
"data": {"type": "object", "description": "Key-value pairs to merge into the scope"}
},
"required": ["scope", "agent", "data"]
}
)
]
@mcp_server.call_tool()
async def handle_call_tool(name: str, arguments: dict) -> list[TextContent]:
# Autorisation liee a l'identite reelle de la connexion SSE (current_agent_var),
# jamais a un champ "agent"/"name" declare par le client dans les arguments.
caller_agent = current_agent_var.get()
if name == "get_rules":
scope = arguments.get("scope")
if not has_scope_access(caller_agent, scope):
return [TextContent(type="text", text=json.dumps({"error": f"Agent {caller_agent} forbidden on scope {scope}"}))]
rule = await get_rule(scope)
if not rule:
return [TextContent(type="text", text=json.dumps({"error": "Scope not found"}))]
return [TextContent(type="text", text=json.dumps(rule))]
elif name == "get_ports":
if not has_scope_access(caller_agent, "infra"):
return [TextContent(type="text", text=json.dumps({"error": f"Agent {caller_agent} forbidden: requires infra scope"}))]
ports = await get_all_ports()
return [TextContent(type="text", text=json.dumps(ports))]
elif name == "search_rules":
query = arguments.get("query", "").lower()
all_rules = await get_all_rules()
results = []
for scope, content_ in all_rules.items():
if not has_scope_access(caller_agent, scope):
continue
if query in json.dumps(content_).lower():
results.append({"scope": scope, "content": content_})
return [TextContent(type="text", text=json.dumps(results))]
elif name == "get_agent_config":
all_rules = await get_all_rules()
config = {}
for scope in ALL_SCOPES:
if has_scope_access(caller_agent, scope):
config[scope] = all_rules.get(scope, {})
return [TextContent(type="text", text=json.dumps(config))]
elif name == "check_port":
if not has_scope_access(caller_agent, "infra"):
return [TextContent(type="text", text=json.dumps({"error": f"Agent {caller_agent} forbidden: requires infra scope"}))]
port = arguments.get("port")
port_info = await get_port(port)
if port_info:
return [TextContent(type="text", text=json.dumps({"status": "occupied", "info": port_info}))]
else:
return [TextContent(type="text", text=json.dumps({"status": "free", "port": port}))]
elif name == "get_my_context":
all_rules = await get_all_rules()
result = {}
for scope in ALL_SCOPES:
if has_scope_access(caller_agent, scope):
result[scope] = all_rules.get(scope, {})
return [TextContent(type="text", text=json.dumps({"agent": caller_agent, "scopes": result}))]
elif name == "update_rule":
scope = arguments.get("scope")
data = arguments.get("data", {})
if not has_scope_access(caller_agent, scope):
return [TextContent(type="text", text=json.dumps({"error": f"Agent {caller_agent} forbidden on scope {scope}"}))]
if not isinstance(data, dict):
return [TextContent(type="text", text=json.dumps({"error": "data must be a JSON object"}))]
existing = await get_rule(scope) or {}
existing.update(data)
await set_rule(scope, existing)
return [TextContent(type="text", text=json.dumps({"status": "ok", "scope": scope, "updated_keys": list(data.keys())}))]
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 = {}
@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():
token = current_agent_var.set(agent)
try:
await mcp_server.run(transport.read_stream(), transport.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 {}