feat(ai): serveur MCP Streamable HTTP + confirmations two-step (#79 phase E)
CI / lint (push) Successful in 1m3s
CI / security (push) Successful in 41s
CI / test (push) Successful in 1m20s
CI / build (push) Successful in 1m19s
CI / e2e (push) Successful in 10m56s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
CI / lint (push) Successful in 1m3s
CI / security (push) Successful in 41s
CI / test (push) Successful in 1m20s
CI / build (push) Successful in 1m19s
CI / e2e (push) Successful in 10m56s
Desktop Build / build-windows (push) Canceled after 0s
Desktop Build / build-linux (push) Canceled after 0s
This commit is contained in:
@@ -800,6 +800,7 @@ class SSESafeGZipMiddleware(GZipMiddleware):
|
||||
"/api/admin/stream",
|
||||
"/api/ai/bookslm/chat",
|
||||
"/api/ai/bookslm/agent",
|
||||
"/mcp",
|
||||
)
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
@@ -882,6 +883,14 @@ try:
|
||||
except ImportError as e:
|
||||
logger.warning(f"Could not load plugins router: {e}")
|
||||
|
||||
# MCP server (Streamable HTTP) for external clients (#79 phase E)
|
||||
try:
|
||||
from backend.mcp.server import McpMount, mcp_app
|
||||
app.router.routes.append(McpMount(mcp_app))
|
||||
logger.info("MCP server mounted at /mcp")
|
||||
except Exception as e: # pragma: no cover - optional dependency
|
||||
logger.warning(f"Could not mount MCP server: {e}")
|
||||
|
||||
# Resolve frontend path relative to this file
|
||||
FRONTEND_DIR = Path(__file__).resolve().parent.parent / "frontend"
|
||||
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
"""Signed, single-use confirmation tokens for MCP mutations (Phase E4).
|
||||
|
||||
MCP has no "Apply" button, so mutating tools are exposed in two steps:
|
||||
``propose_<tool>`` returns a preview plus a **signed token**, and
|
||||
``apply_<tool>`` consumes that token to execute the mutation.
|
||||
|
||||
The token is a short-lived JWT (same secret as the app) carrying the tool name
|
||||
and its arguments. Single use is enforced by a persisted JTI blacklist, which
|
||||
also protects against replay and TOCTOU (a stale proposal cannot be re-applied).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from jose import JWTError, jwt
|
||||
|
||||
from backend.auth.jwt_handler import get_secret_key
|
||||
|
||||
logger = logging.getLogger("obsigate.mcp.confirmations")
|
||||
|
||||
ALGORITHM = "HS256"
|
||||
TOKEN_TYPE = "mcp_confirmation"
|
||||
|
||||
# Default token lifetime (seconds); override with OBSIGATE_MCP_CONFIRMATION_TTL.
|
||||
DEFAULT_TTL = int(os.environ.get("OBSIGATE_MCP_CONFIRMATION_TTL", "300"))
|
||||
|
||||
_USED_TOKENS_FILE = Path("data/mcp_used_tokens.json")
|
||||
_used_lock = threading.RLock()
|
||||
_used_loaded = False
|
||||
_used_jtis: dict[str, int] = {}
|
||||
|
||||
|
||||
class ConfirmationError(Exception):
|
||||
"""Raised when a confirmation token is invalid, expired, or already used."""
|
||||
|
||||
def __init__(self, message: str, *, code: str = "invalid_confirmation"):
|
||||
super().__init__(message)
|
||||
self.message = message
|
||||
self.code = code
|
||||
|
||||
|
||||
def _load_used() -> None:
|
||||
"""Load used JTIs from disk once, dropping expired entries."""
|
||||
global _used_loaded, _used_jtis
|
||||
if _used_loaded:
|
||||
return
|
||||
with _used_lock:
|
||||
if _used_loaded:
|
||||
return
|
||||
if _USED_TOKENS_FILE.exists():
|
||||
try:
|
||||
data = json.loads(_USED_TOKENS_FILE.read_text(encoding="utf-8"))
|
||||
now = int(time.time())
|
||||
_used_jtis = {jti: exp for jti, exp in data.items() if int(exp) > now}
|
||||
except Exception as e: # pragma: no cover - corrupt store
|
||||
logger.warning(f"Failed to load used MCP tokens: {e}")
|
||||
_used_jtis = {}
|
||||
_used_loaded = True
|
||||
|
||||
|
||||
def _save_used() -> None:
|
||||
"""Persist used JTIs with their expiry (best-effort)."""
|
||||
try:
|
||||
_USED_TOKENS_FILE.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp = _USED_TOKENS_FILE.with_suffix(".tmp")
|
||||
tmp.write_text(json.dumps(_used_jtis), encoding="utf-8")
|
||||
tmp.replace(_USED_TOKENS_FILE)
|
||||
except Exception as e: # pragma: no cover - disk error
|
||||
logger.warning(f"Failed to persist used MCP tokens: {e}")
|
||||
|
||||
|
||||
def create_confirmation_token(
|
||||
tool: str,
|
||||
arguments: dict[str, Any],
|
||||
username: str,
|
||||
*,
|
||||
ttl: int | None = None,
|
||||
) -> str:
|
||||
"""Create a signed confirmation token for a pending mutation."""
|
||||
now = int(time.time())
|
||||
payload = {
|
||||
"type": TOKEN_TYPE,
|
||||
"tool": tool,
|
||||
"arguments": arguments,
|
||||
"sub": username,
|
||||
"jti": str(uuid.uuid4()),
|
||||
"iat": now,
|
||||
"exp": now + (ttl if ttl is not None else DEFAULT_TTL),
|
||||
}
|
||||
return jwt.encode(payload, get_secret_key(), algorithm=ALGORITHM)
|
||||
|
||||
|
||||
def peek_confirmation_token(token: str) -> dict[str, Any]:
|
||||
"""Decode and validate a token without consuming it (signature + expiry)."""
|
||||
try:
|
||||
payload = jwt.decode(token, get_secret_key(), algorithms=[ALGORITHM])
|
||||
except JWTError as e:
|
||||
raise ConfirmationError("Invalid or expired confirmation token", code="invalid_confirmation") from e
|
||||
|
||||
if payload.get("type") != TOKEN_TYPE:
|
||||
raise ConfirmationError("Wrong token type", code="invalid_confirmation")
|
||||
if not payload.get("tool") or not payload.get("sub"):
|
||||
raise ConfirmationError("Malformed confirmation token", code="invalid_confirmation")
|
||||
return payload
|
||||
|
||||
|
||||
def consume_confirmation_token(token: str, username: str) -> dict[str, Any]:
|
||||
"""Validate a token, enforce single use, and return its payload.
|
||||
|
||||
Raises:
|
||||
ConfirmationError: invalid/expired token, wrong user, or replay.
|
||||
"""
|
||||
payload = peek_confirmation_token(token)
|
||||
|
||||
if payload.get("sub") != username:
|
||||
raise ConfirmationError("Confirmation token does not belong to this user", code="forbidden")
|
||||
|
||||
jti = payload.get("jti")
|
||||
if not jti:
|
||||
raise ConfirmationError("Malformed confirmation token", code="invalid_confirmation")
|
||||
|
||||
_load_used()
|
||||
now = int(time.time())
|
||||
with _used_lock:
|
||||
if jti in _used_jtis:
|
||||
raise ConfirmationError("Confirmation token already used", code="token_reused")
|
||||
# Record as used *before* returning so a concurrent replay is rejected.
|
||||
_used_jtis[jti] = int(payload.get("exp", now))
|
||||
# Opportunistic cleanup of expired JTIs.
|
||||
for old_jti in [k for k, exp in _used_jtis.items() if exp <= now]:
|
||||
_used_jtis.pop(old_jti, None)
|
||||
_save_used()
|
||||
|
||||
return payload
|
||||
@@ -0,0 +1,485 @@
|
||||
"""ObsiGate MCP server — Streamable HTTP transport (Phase E).
|
||||
|
||||
Exposes the shared AI tool layer (``backend.tools``) to external MCP clients
|
||||
(Claude Desktop, Cursor…). The same registry that powers the in-app assistant
|
||||
is registered here, so the two fronts never diverge.
|
||||
|
||||
Primitives:
|
||||
- **Tools** — read/search tools directly; mutating tools as a two-step
|
||||
``propose_<tool>`` / ``apply_<tool>`` pair (signed, single-use token).
|
||||
- **Resources** — accessible vaults (``vault://<name>``) and files
|
||||
(``vault://<name>/<path>``), read-only and secret-redacted.
|
||||
- **Prompts** — reusable note/summary templates.
|
||||
|
||||
Transport: Streamable HTTP mounted at ``/mcp`` (auth ``Authorization: Bearer
|
||||
<JWT>``). ``stdio`` is left for later.
|
||||
|
||||
The transport manager is started lazily on the first request so the endpoint
|
||||
also works in tests (where the ASGI lifespan is not run).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import difflib
|
||||
import json
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from fastapi.responses import JSONResponse
|
||||
from mcp import types
|
||||
from mcp.server.lowlevel import Server
|
||||
from mcp.server.lowlevel.helper_types import ReadResourceContents
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
||||
from pydantic import AnyUrl
|
||||
from starlette._utils import get_route_path
|
||||
from starlette.requests import Request
|
||||
from starlette.routing import BaseRoute, Match
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
from backend.auth.middleware import get_current_user, is_auth_enabled
|
||||
from backend.services.files import read_file_text
|
||||
from backend.services.vaults import list_accessible_vaults
|
||||
from backend.tools.api import (
|
||||
ToolContext,
|
||||
ToolError,
|
||||
ToolMode,
|
||||
ToolRisk,
|
||||
ToolScope,
|
||||
call_tool,
|
||||
get_tool,
|
||||
list_tools,
|
||||
)
|
||||
|
||||
logger = logging.getLogger("obsigate.mcp")
|
||||
|
||||
SERVER_NAME = "obsigate"
|
||||
SERVER_INSTRUCTIONS = (
|
||||
"ObsiGate exposes your Obsidian vaults: read, search and (with confirmation) "
|
||||
"create, edit, rename, move or delete notes. Mutating tools require the "
|
||||
"two-step propose_/apply_ flow."
|
||||
)
|
||||
|
||||
# Cap on the file size returned by resources (bytes).
|
||||
MAX_RESOURCE_BYTES = 200_000
|
||||
|
||||
|
||||
def _anonymous_user() -> dict[str, Any]:
|
||||
return {
|
||||
"username": "anonymous",
|
||||
"display_name": "Anonymous",
|
||||
"role": "admin",
|
||||
"vaults": ["*"],
|
||||
"active": True,
|
||||
"_token_vaults": ["*"],
|
||||
}
|
||||
|
||||
|
||||
def _authenticate(request: Request) -> dict[str, Any] | None:
|
||||
"""Resolve the caller from the ``Authorization: Bearer`` header."""
|
||||
if not is_auth_enabled():
|
||||
return _anonymous_user()
|
||||
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
|
||||
header = request.headers.get("authorization", "")
|
||||
if not header.lower().startswith("bearer "):
|
||||
return None
|
||||
credentials = HTTPAuthorizationCredentials(scheme="Bearer", credentials=header[7:].strip())
|
||||
return get_current_user(request, credentials)
|
||||
|
||||
|
||||
def _user_from_context(server: Server) -> dict[str, Any]:
|
||||
"""Return the authenticated user attached to the current request scope."""
|
||||
request = server.request_context.request
|
||||
if request is None: # pragma: no cover - stdio not supported yet
|
||||
raise ToolError("No request context available", code="no_request_context")
|
||||
user = getattr(request.state, "user", None)
|
||||
if not user:
|
||||
raise ToolError("Unauthenticated MCP request", code="unauthenticated")
|
||||
return user
|
||||
|
||||
|
||||
def _text(payload: Any) -> list[types.Content]:
|
||||
"""Serialize a payload as a single text content block."""
|
||||
return [types.TextContent(type="text", text=json.dumps(payload, ensure_ascii=False, default=str))]
|
||||
|
||||
|
||||
def _unified_diff(old: str, new: str) -> str:
|
||||
diff = difflib.unified_diff(
|
||||
old.splitlines(keepends=True),
|
||||
new.splitlines(keepends=True),
|
||||
fromfile="current",
|
||||
tofile="proposed",
|
||||
)
|
||||
return "".join(diff)
|
||||
|
||||
|
||||
def _build_preview(spec: Any, params: Any) -> dict[str, Any]:
|
||||
"""Build a human-readable preview for a proposed mutation."""
|
||||
arguments = params.model_dump()
|
||||
preview: dict[str, Any] = {"tool": spec.name, "arguments": arguments}
|
||||
|
||||
if spec.name in ("create_file", "edit_file", "append_to_file"):
|
||||
vault = arguments.get("vault")
|
||||
path = arguments.get("path", "")
|
||||
content = arguments.get("content", "")
|
||||
try:
|
||||
current = read_file_text(vault, path, redact=True, max_bytes=MAX_RESOURCE_BYTES)["content"]
|
||||
except Exception:
|
||||
current = ""
|
||||
if spec.name == "append_to_file":
|
||||
separator = "" if (not current or current.endswith("\n")) else "\n"
|
||||
proposed = current + separator + content
|
||||
else:
|
||||
proposed = content
|
||||
preview["diff"] = _unified_diff(current, proposed)
|
||||
|
||||
return preview
|
||||
|
||||
|
||||
def _tool_definitions() -> list[types.Tool]:
|
||||
"""Build the MCP tool list from the shared registry."""
|
||||
definitions: list[types.Tool] = []
|
||||
for spec in list_tools(scope=ToolScope.MCP):
|
||||
if spec.risk == ToolRisk.READ:
|
||||
definitions.append(
|
||||
types.Tool(
|
||||
name=spec.name,
|
||||
description=spec.description,
|
||||
inputSchema=spec.parameters_schema(),
|
||||
annotations=types.ToolAnnotations(readOnlyHint=True, openWorldHint=False),
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
destructive = spec.risk == ToolRisk.DANGEROUS
|
||||
definitions.append(
|
||||
types.Tool(
|
||||
name=f"propose_{spec.name}",
|
||||
description=(
|
||||
f"Propose to run '{spec.name}' and return a confirmation token. "
|
||||
"No change is made until apply_ is called."
|
||||
),
|
||||
inputSchema=spec.parameters_schema(),
|
||||
annotations=types.ToolAnnotations(
|
||||
readOnlyHint=True, destructiveHint=False, idempotentHint=True, openWorldHint=False
|
||||
),
|
||||
)
|
||||
)
|
||||
definitions.append(
|
||||
types.Tool(
|
||||
name=f"apply_{spec.name}",
|
||||
description=f"Apply a previously proposed '{spec.name}' using its confirmation token.",
|
||||
inputSchema={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"confirmation_token": {
|
||||
"type": "string",
|
||||
"description": f"Token returned by propose_{spec.name}",
|
||||
}
|
||||
},
|
||||
"required": ["confirmation_token"],
|
||||
},
|
||||
annotations=types.ToolAnnotations(
|
||||
readOnlyHint=False, destructiveHint=destructive, idempotentHint=False, openWorldHint=False
|
||||
),
|
||||
)
|
||||
)
|
||||
return definitions
|
||||
|
||||
|
||||
def _parse_vault_uri(uri: str) -> tuple[str, str]:
|
||||
"""Split ``vault://name/path`` into ``(vault, path)``."""
|
||||
parsed = urlparse(uri)
|
||||
if parsed.scheme != "vault" or not parsed.netloc:
|
||||
raise ValueError(f"Unsupported resource URI: {uri}")
|
||||
return parsed.netloc, parsed.path.lstrip("/")
|
||||
|
||||
|
||||
def build_server() -> Server:
|
||||
"""Create and configure the low-level MCP server."""
|
||||
from backend.version import get_version
|
||||
|
||||
server: Server = Server(SERVER_NAME, version=get_version(), instructions=SERVER_INSTRUCTIONS)
|
||||
|
||||
# ── Tools ──────────────────────────────────────────────────────────
|
||||
@server.list_tools()
|
||||
async def _list_tools() -> list[types.Tool]:
|
||||
return _tool_definitions()
|
||||
|
||||
@server.call_tool()
|
||||
async def _call_tool(name: str, arguments: dict[str, Any]) -> list[types.Content]:
|
||||
user = _user_from_context(server)
|
||||
|
||||
if name.startswith("propose_"):
|
||||
return _propose(user, name[len("propose_"):], arguments)
|
||||
if name.startswith("apply_"):
|
||||
return _apply(user, name[len("apply_"):], arguments)
|
||||
|
||||
spec = get_tool(name)
|
||||
if spec is None:
|
||||
raise ValueError(f"Unknown tool: {name}")
|
||||
if spec.risk != ToolRisk.READ:
|
||||
raise ValueError(f"Tool '{name}' is mutating; use propose_{name}/apply_{name}")
|
||||
|
||||
ctx = ToolContext(user=user, mode=ToolMode.MCP)
|
||||
try:
|
||||
result = call_tool(name, ctx, arguments)
|
||||
except ToolError as e:
|
||||
return _text(e.to_dict())
|
||||
return _text({"ok": True, "data": result.data})
|
||||
|
||||
# ── Resources ──────────────────────────────────────────────────────
|
||||
@server.list_resources()
|
||||
async def _list_resources() -> list[types.Resource]:
|
||||
user = _user_from_context(server)
|
||||
return [
|
||||
types.Resource(
|
||||
uri=AnyUrl(f"vault://{v['name']}"),
|
||||
name=v["name"],
|
||||
description=f"ObsiGate vault '{v['name']}' ({v['file_count']} files)",
|
||||
mimeType="application/x-obsigate-vault",
|
||||
)
|
||||
for v in list_accessible_vaults(user)
|
||||
]
|
||||
|
||||
@server.list_resource_templates()
|
||||
async def _list_resource_templates() -> list[types.ResourceTemplate]:
|
||||
return [
|
||||
types.ResourceTemplate(
|
||||
uriTemplate="vault://{vault}/{path}",
|
||||
name="Vault file",
|
||||
description="Read a text file from a vault (secrets redacted)",
|
||||
mimeType="text/markdown",
|
||||
)
|
||||
]
|
||||
|
||||
@server.read_resource()
|
||||
async def _read_resource(uri: Any) -> list[ReadResourceContents]:
|
||||
user = _user_from_context(server)
|
||||
vault, path = _parse_vault_uri(str(uri))
|
||||
ctx = ToolContext(user=user, mode=ToolMode.MCP)
|
||||
ctx.require_vault_access(vault)
|
||||
data = read_file_text(vault, path, redact=True, max_bytes=MAX_RESOURCE_BYTES)
|
||||
mime = "text/markdown" if Path(path).suffix.lower() == ".md" else "text/plain"
|
||||
return [ReadResourceContents(content=data["content"], mime_type=mime)]
|
||||
|
||||
# ── Prompts ────────────────────────────────────────────────────────
|
||||
@server.list_prompts()
|
||||
async def _list_prompts() -> list[types.Prompt]:
|
||||
return [
|
||||
types.Prompt(
|
||||
name="summarize-directory",
|
||||
description="Summarize every note in a vault directory.",
|
||||
arguments=[
|
||||
types.PromptArgument(name="vault", description="Vault name", required=True),
|
||||
types.PromptArgument(name="path", description="Directory path (empty = root)", required=False),
|
||||
],
|
||||
),
|
||||
types.Prompt(
|
||||
name="generate-note",
|
||||
description="Draft a new note on a topic, using the vault for context.",
|
||||
arguments=[
|
||||
types.PromptArgument(name="topic", description="Note topic", required=True),
|
||||
types.PromptArgument(name="vault", description="Target vault name", required=True),
|
||||
],
|
||||
),
|
||||
types.Prompt(
|
||||
name="find-related",
|
||||
description="Find notes related to a given note.",
|
||||
arguments=[
|
||||
types.PromptArgument(name="vault", description="Vault name", required=True),
|
||||
types.PromptArgument(name="path", description="Reference note path", required=True),
|
||||
],
|
||||
),
|
||||
]
|
||||
|
||||
@server.get_prompt()
|
||||
async def _get_prompt(name: str, arguments: dict[str, str] | None) -> types.GetPromptResult:
|
||||
args = arguments or {}
|
||||
|
||||
def _require(key: str) -> str:
|
||||
value = (args.get(key) or "").strip()
|
||||
if not value:
|
||||
raise ValueError(f"Missing required prompt argument: {key}")
|
||||
return value
|
||||
|
||||
if name == "summarize-directory":
|
||||
vault = _require("vault")
|
||||
path = (args.get("path") or "").strip()
|
||||
text = (
|
||||
f"List the directory '{path or '/'}' of vault '{vault}' (use list_directory), "
|
||||
"read the notes it contains, then write a concise summary with the key ideas."
|
||||
)
|
||||
elif name == "generate-note":
|
||||
topic = _require("topic")
|
||||
vault = _require("vault")
|
||||
text = (
|
||||
f"Draft a well-structured Markdown note about '{topic}' in vault '{vault}'. "
|
||||
"Search the vault first for existing material, then propose the note content."
|
||||
)
|
||||
elif name == "find-related":
|
||||
vault = _require("vault")
|
||||
path = _require("path")
|
||||
text = (
|
||||
f"Read '{path}' in vault '{vault}', then find related notes via backlinks and "
|
||||
"full-text search. Return a short list with why each note is related."
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unknown prompt: {name}")
|
||||
|
||||
return types.GetPromptResult(
|
||||
description=f"ObsiGate prompt: {name}",
|
||||
messages=[types.PromptMessage(role="user", content=types.TextContent(type="text", text=text))],
|
||||
)
|
||||
|
||||
return server
|
||||
|
||||
|
||||
def _propose(user: dict[str, Any], tool: str, arguments: dict[str, Any]) -> list[types.Content]:
|
||||
"""Validate a mutation, return a preview and a confirmation token."""
|
||||
from backend.mcp.confirmations import DEFAULT_TTL, create_confirmation_token
|
||||
|
||||
spec = get_tool(tool)
|
||||
if spec is None:
|
||||
raise ValueError(f"Unknown tool: {tool}")
|
||||
if spec.risk == ToolRisk.READ:
|
||||
raise ValueError(f"Tool '{tool}' is read-only; call it directly")
|
||||
|
||||
params = spec.input_model.model_validate(arguments)
|
||||
vault = getattr(params, "vault", None)
|
||||
ctx = ToolContext(user=user, mode=ToolMode.MCP)
|
||||
if vault and vault != "all":
|
||||
ctx.require_vault_access(vault)
|
||||
if spec.risk == ToolRisk.DANGEROUS:
|
||||
ctx.require_destructive_allowed(vault)
|
||||
|
||||
token = create_confirmation_token(tool, arguments, ctx.username)
|
||||
payload = _build_preview(spec, params)
|
||||
payload["confirmation_token"] = token
|
||||
payload["expires_in"] = DEFAULT_TTL
|
||||
return _text(payload)
|
||||
|
||||
|
||||
def _apply(user: dict[str, Any], tool: str, arguments: dict[str, Any]) -> list[types.Content]:
|
||||
"""Consume a confirmation token and execute the mutation."""
|
||||
from backend.mcp.confirmations import ConfirmationError, consume_confirmation_token
|
||||
|
||||
token = (arguments or {}).get("confirmation_token")
|
||||
if not token:
|
||||
raise ValueError("Missing 'confirmation_token'")
|
||||
|
||||
try:
|
||||
payload = consume_confirmation_token(token, user.get("username", ""))
|
||||
except ConfirmationError as e:
|
||||
return _text({"ok": False, "error": {"code": e.code, "message": e.message}})
|
||||
|
||||
if payload.get("tool") != tool:
|
||||
return _text(
|
||||
{
|
||||
"ok": False,
|
||||
"error": {
|
||||
"code": "token_mismatch",
|
||||
"message": f"Token was issued for '{payload.get('tool')}', not '{tool}'",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
ctx = ToolContext(user=user, mode=ToolMode.MCP, confirmed=True)
|
||||
try:
|
||||
result = call_tool(tool, ctx, payload.get("arguments") or {}, confirm=True)
|
||||
except ToolError as e:
|
||||
return _text(e.to_dict())
|
||||
return _text({"ok": True, "data": result.data})
|
||||
|
||||
|
||||
class McpASGIApp:
|
||||
"""ASGI wrapper: authenticates the request, then runs the MCP transport.
|
||||
|
||||
The session manager is started lazily on first use and kept alive for the
|
||||
process lifetime, so the endpoint works both under uvicorn (with lifespan)
|
||||
and under the test client (without entering the ASGI lifespan).
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._manager: StreamableHTTPSessionManager | None = None
|
||||
self._start_lock: asyncio.Lock | None = None
|
||||
self._run_task: asyncio.Task | None = None
|
||||
|
||||
def _get_manager(self) -> StreamableHTTPSessionManager:
|
||||
if self._manager is None:
|
||||
self._manager = StreamableHTTPSessionManager(
|
||||
app=build_server(),
|
||||
json_response=True,
|
||||
stateless=False,
|
||||
)
|
||||
return self._manager
|
||||
|
||||
async def _ensure_started(self) -> StreamableHTTPSessionManager:
|
||||
manager = self._get_manager()
|
||||
if getattr(manager, "_task_group", None) is not None:
|
||||
return manager
|
||||
if self._start_lock is None:
|
||||
self._start_lock = asyncio.Lock()
|
||||
async with self._start_lock:
|
||||
if getattr(manager, "_task_group", None) is None:
|
||||
self._run_task = asyncio.create_task(self._run_manager(manager))
|
||||
for _ in range(500):
|
||||
if getattr(manager, "_task_group", None) is not None:
|
||||
break
|
||||
await asyncio.sleep(0.005)
|
||||
return manager
|
||||
|
||||
async def _run_manager(self, manager: StreamableHTTPSessionManager) -> None:
|
||||
async with manager.run():
|
||||
await asyncio.Event().wait()
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
return
|
||||
|
||||
request = Request(scope, receive)
|
||||
user = _authenticate(request)
|
||||
if user is None:
|
||||
response = JSONResponse(
|
||||
{"detail": "Authentification requise"},
|
||||
status_code=401,
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
await response(scope, receive, send)
|
||||
return
|
||||
|
||||
scope.setdefault("state", {})["user"] = user
|
||||
manager = await self._ensure_started()
|
||||
await manager.handle_request(scope, receive, send)
|
||||
|
||||
|
||||
# Module-level singleton mounted by ``backend.main`` at ``/mcp``.
|
||||
mcp_app = McpASGIApp()
|
||||
|
||||
|
||||
class McpMount(BaseRoute):
|
||||
"""ASGI route matching ``/mcp`` and ``/mcp/...`` (unlike Starlette's Mount).
|
||||
|
||||
Starlette's :class:`~starlette.routing.Mount` compiles ``/mcp/{path:path}``
|
||||
and therefore does **not** match the bare ``/mcp`` path used by MCP clients.
|
||||
This route matches both forms and delegates to :data:`mcp_app` unchanged.
|
||||
"""
|
||||
|
||||
def __init__(self, app: ASGIApp, path: str = "/mcp") -> None:
|
||||
self.app = app
|
||||
self.path = path.rstrip("/")
|
||||
|
||||
def matches(self, scope: Scope) -> tuple[Match, Scope]:
|
||||
if scope["type"] == "http":
|
||||
route_path = get_route_path(scope)
|
||||
if route_path == self.path or route_path.startswith(self.path + "/"):
|
||||
return Match.FULL, {"endpoint": self.app}
|
||||
return Match.NONE, {}
|
||||
|
||||
async def handle(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
@@ -17,3 +17,5 @@ pyotp>=2.10.0
|
||||
webauthn==2.6.0
|
||||
psutil>=5.9
|
||||
pywebpush>=2.3.0
|
||||
mcp==1.9.4
|
||||
sse-starlette==2.1.3
|
||||
|
||||
Reference in New Issue
Block a user