Files
ObsiGate/backend/mcp/server.py
T
bruno 88eecd7671
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
feat(ai): serveur MCP Streamable HTTP + confirmations two-step (#79 phase E)
2026-09-11 21:28:05 -04:00

486 lines
18 KiB
Python

"""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)