Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 30 additions & 19 deletions nerve/mcp_server/http.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
import logging
from typing import TYPE_CHECKING, Callable

from mcp.server.lowlevel.server import request_ctx
from mcp.server.context import ServerRequestContext
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
from starlette.types import Receive, Scope, Send

Expand Down Expand Up @@ -71,22 +71,28 @@ async def _send_status(send: Send, status: int, message: str) -> None:
await send({"type": "http.response.body", "body": body, "more_body": False})


def _resolve_client_info() -> tuple[str | None, str | None, str | None]:
"""Read client metadata from the active MCP request context.
def _resolve_client_info(
rctx: ServerRequestContext | None,
) -> tuple[str | None, str | None, str | None]:
"""Read client metadata from the supplied MCP request context.

Returns ``(client_name, mcp_session_id, request_path)`` — any field
may be ``None`` if the corresponding data isn't available (e.g.
during the very first ``initialize`` call the session may not yet
have ``client_params``).

mcp 2.x hands the request context to handlers as an argument rather
than exposing it through a contextvar, so callers thread it down from
the handler. ``None`` is accepted for callers with no live request.
"""
try:
rctx = request_ctx.get()
except LookupError:
if rctx is None:
return None, None, None

client_name: str | None = None
if rctx.session and rctx.session.client_params:
info = rctx.session.client_params.clientInfo
# mcp 2.x exposes the model fields as snake_case in Python; the wire
# form is still camelCase (`clientInfo`) via the alias generator.
info = rctx.session.client_params.client_info
if info is not None:
client_name = info.name

Expand All @@ -105,6 +111,7 @@ def _resolve_client_info() -> tuple[str | None, str | None, str | None]:

def _bound_identity_from_request(
config: "NerveConfig",
rctx: ServerRequestContext | None,
) -> tuple[str | None, dict[str, str]]:
"""Session id bound into the request's bearer token, if any.

Expand All @@ -120,9 +127,7 @@ def _bound_identity_from_request(
"""
if not config.auth.jwt_secret:
return None, {}
try:
rctx = request_ctx.get()
except LookupError:
if rctx is None:
return None, {}
request = getattr(rctx, "request", None)
if request is None:
Expand Down Expand Up @@ -150,28 +155,34 @@ def _bound_identity_from_request(
return bound_session_id(payload), runtime


def _bound_session_from_request(config: "NerveConfig") -> str | None:
def _bound_session_from_request(
config: "NerveConfig",
rctx: ServerRequestContext | None,
) -> str | None:
"""Backward-compatible session-only view used by tests/callers."""
return _bound_identity_from_request(config)[0]
return _bound_identity_from_request(config, rctx)[0]


def build_ctx_resolver(engine: "AgentEngine", resolver: SatelliteSessionResolver):
"""Build the per-call_tool ``ToolContext`` resolver closure.

The Server's ``call_tool`` handler invokes this for every tool call
to attribute the call to the correct session: a session-bound token
(backend-managed agents) binds directly to its engine session;
everything else goes through satellite attribution. Per-call
The Server's ``call_tool`` handler invokes this for every tool call,
passing the request's ``ServerRequestContext``, to attribute the call
to the correct session: a session-bound token (backend-managed
agents) binds directly to its engine session; everything else goes
through satellite attribution. Per-call
resolution is cheap (the satellite session id is deterministic and
the underlying ``get_session`` / ``create_session`` check is O(1)
on the indexed primary key).
"""

async def _resolve() -> ToolContext:
session_id, runtime_metadata = _bound_identity_from_request(engine.config)
async def _resolve(rctx: ServerRequestContext | None = None) -> ToolContext:
session_id, runtime_metadata = _bound_identity_from_request(
engine.config, rctx
)

if session_id is None:
client_name, mcp_session_id, _ = _resolve_client_info()
client_name, mcp_session_id, _ = _resolve_client_info(rctx)

if mcp_session_id is None:
# Stateless requests / pre-initialize calls can land here.
Expand Down
118 changes: 81 additions & 37 deletions nerve/mcp_server/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,17 @@

The ``ctx_resolver`` callable is invoked per ``call_tool`` request to
build a fresh :class:`ToolContext` for the satellite session that owns
this MCP connection. It returns an awaitable so the resolver can fetch
state from the DB on first use and cache it for subsequent calls.
this MCP connection. It receives the request's
:class:`~mcp.server.context.ServerRequestContext` — mcp 2.x hands that to
handlers as an argument rather than exposing it through a contextvar — and
returns an awaitable so the resolver can fetch state from the DB on first
use and cache it for subsequent calls.

Argument validation against each tool's ``inputSchema`` happens here,
explicitly. mcp 1.x's ``@server.call_tool()`` decorator did it for us by
default (``validate_input=True``); the 2.x ``on_call_tool`` callback does
not, so relying on the library would have silently dropped validation on an
endpoint external clients can reach.

The ``audit_writer`` callable is invoked after every successful tool
call to persist an ``external_tool_call`` event into ``session_events``.
Expand All @@ -21,15 +30,24 @@
import logging
from typing import Any, Awaitable, Callable

import jsonschema
from mcp.server.context import ServerRequestContext
from mcp.server.lowlevel import Server
from mcp.types import CallToolResult, TextContent, Tool
from mcp.types import (
CallToolRequestParams,
CallToolResult,
ListToolsResult,
PaginatedRequestParams,
TextContent,
Tool,
)

from nerve.agent.tools import ToolContext, ToolRegistry, ToolResult

logger = logging.getLogger(__name__)


CtxResolver = Callable[[], Awaitable[ToolContext]]
CtxResolver = Callable[[ServerRequestContext], Awaitable[ToolContext]]
AuditWriter = Callable[[ToolContext, str, dict, ToolResult, float, bool], Awaitable[None]]


Expand Down Expand Up @@ -67,46 +85,68 @@ def build_mcp_server(
"""
import time

server: Server = Server(name=name, version=version)

@server.list_tools()
async def _list_tools() -> list[Tool]:
return [
Tool(
name=spec.name,
description=spec.description,
inputSchema=spec.input_schema,
)
for spec in registry.list(include_hoa=include_hoa)
]

@server.call_tool()
async def _call_tool(name: str, arguments: dict[str, Any]) -> CallToolResult:
def _error(message: str) -> CallToolResult:
"""A failed call, shaped exactly like every other failure here."""
return CallToolResult(
content=[TextContent(type="text", text=message)],
isError=True,
)

async def _list_tools(
rctx: ServerRequestContext,
params: PaginatedRequestParams | None = None,
) -> ListToolsResult:
return ListToolsResult(
tools=[
Tool(
name=spec.name,
description=spec.description,
inputSchema=spec.input_schema,
)
for spec in registry.list(include_hoa=include_hoa)
]
)

async def _call_tool(
rctx: ServerRequestContext,
params: CallToolRequestParams,
) -> CallToolResult:
name = params.name
arguments: dict[str, Any] = dict(params.arguments or {})

spec = registry.get(name)
if spec is None:
return CallToolResult(
content=[TextContent(type="text", text=f"Unknown tool: {name!r}")],
isError=True,
)
return _error(f"Unknown tool: {name!r}")

# HoA gating: registry.list() filters by include_hoa, but a
# malicious caller could still invoke a HoA tool by name. Enforce
# the same allowlist here.
if not include_hoa and name.startswith("hoa_"):
return CallToolResult(
content=[TextContent(type="text", text=f"Tool not available: {name!r}")],
isError=True,
)
return _error(f"Tool not available: {name!r}")

# Validate arguments against the tool's declared inputSchema before
# any handler sees them. mcp 1.x's call_tool decorator did this by
# default; the 2.x callback does not, and this endpoint is reachable
# by external clients, so the check lives here explicitly rather
# than depending on a library default.
if spec.input_schema:
try:
jsonschema.validate(instance=arguments, schema=spec.input_schema)
except jsonschema.ValidationError as e:
logger.info("Invalid arguments for %s: %s", name, e.message)
return _error(f"Invalid arguments for {name!r}: {e.message}")
except jsonschema.SchemaError:
# A malformed schema is our bug, not the caller's. Refuse the
# call rather than dispatching unvalidated arguments.
logger.exception("Tool %s has an invalid inputSchema", name)
return _error(f"Tool {name!r} has an invalid input schema")

start = time.monotonic()
try:
ctx = await ctx_resolver()
ctx = await ctx_resolver(rctx)
except Exception as e:
logger.exception("Failed to resolve ToolContext for %s", name)
return CallToolResult(
content=[TextContent(type="text", text=f"Context error: {e}")],
isError=True,
)
return _error(f"Context error: {e}")

try:
result = await spec.handler(ctx, arguments)
Expand All @@ -126,10 +166,7 @@ async def _call_tool(name: str, arguments: dict[str, Any]) -> CallToolResult:
)
except Exception:
logger.exception("Audit writer failed for %s", name)
return CallToolResult(
content=[TextContent(type="text", text=f"Tool error: {e}")],
isError=True,
)
return _error(f"Tool error: {e}")

duration_ms = (time.monotonic() - start) * 1000.0
if audit_writer is not None:
Expand Down Expand Up @@ -161,4 +198,11 @@ async def _call_tool(name: str, arguments: dict[str, Any]) -> CallToolResult:

return CallToolResult(content=content, isError=result.is_error)

return server
# mcp 2.x registers handlers as constructor callbacks; the 1.x
# ``@server.list_tools()`` / ``@server.call_tool()`` decorators are gone.
return Server(
name=name,
version=version,
on_list_tools=_list_tools,
on_call_tool=_call_tool,
)
20 changes: 19 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,25 @@ dependencies = [
"aiosqlite>=0.21.0",
"pyyaml>=6.0",
"python-telegram-bot>=21.0",
"claude-agent-sdk>=0.2.82",
# Floor raised from 0.2.82 for mcp 2.x: 0.2.140 is the first release whose
# own constraint is `mcp<3.0.0` (0.2.100 through 0.2.139 all declare
# `mcp<2.0.0`). It matters because the SDK builds the in-process MCP server
# that serves EVERY agent tool, so an SDK without 2.x support breaks the
# whole Claude backend, not just /mcp/v1. 0.2.82 is the dangerous case: it
# declares `mcp>=1.23.0` with no upper bound, so `0.2.82 + mcp 2.x` is a
# resolvable combination that cannot work.
"claude-agent-sdk>=0.2.140",
# Declared explicitly because nerve/mcp_server/ imports `mcp` directly
# rather than only through claude-agent-sdk. Both bounds are load-bearing:
# nerve/mcp_server/ is written against the 2.x lowlevel API (constructor
# callbacks, request context as a handler argument), which does not exist
# in 1.x, and 3.x has not been vetted.
"mcp>=2,<3",
# Used directly by nerve/mcp_server/server.py to validate tool arguments
# against each tool's inputSchema. mcp 1.x's call_tool decorator did this
# for us; the 2.x callback does not, so the check is ours now and the
# dependency is declared rather than borrowed from claude-agent-sdk.
"jsonschema>=4.20",
"apscheduler>=3.11.0",
"pyjwt>=2.10.0",
"bcrypt>=4.2.0",
Expand Down
Loading
Loading