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
486 lines
18 KiB
Python
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)
|
|
|