Skip to content

Commit c54075c

Browse files
authored
Look the tool schema up by name for Mcp-Param-* validation instead of running tools/list (#3630)
1 parent 9afccae commit c54075c

8 files changed

Lines changed: 228 additions & 7 deletions

File tree

‎docs/advanced/header-parameters.md‎

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,11 +38,23 @@ There you write `input_schema` by hand, so the key goes straight in:
3838

3939
* Nothing checks the annotation for you: an invalid one is served, and `2026-07-28` clients leave the tool out of their listing.
4040

41+
### Schemas by name
42+
43+
To check the header, the SDK needs the tool's input schema before it dispatches the call. Without `get_tool_input_schema` it gets it by running your `on_list_tools` handler on every call that carries arguments, whether or not any tool is marked.
44+
45+
```python title="server.py" hl_lines="26 39-41 48"
46+
--8<-- "docs_src/header_parameters/tutorial003.py"
47+
```
48+
49+
* Pass the function to answer from what you already have.
50+
* Return `None` for a tool with nothing to check.
51+
4152
## Recap
4253

4354
* `x-mcp-header` on a tool argument makes `2026-07-28` clients repeat it as an `Mcp-Param-*` HTTP header.
4455
* The server rejects a call whose header and body disagree.
4556
* Only `str`, `int` and `bool` arguments can be marked. `MCPServer` raises `InvalidSignature` for anything else.
4657
* The low-level `Server` checks nothing, and clients drop a tool whose annotation is invalid.
58+
* `get_tool_input_schema` keeps the low-level `Server` from running `on_list_tools` on every call.
4759

4860
The rest of the hand-written `Server` API is **[The low-level Server](low-level-server.md)**.

‎docs/advanced/low-level-server.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -204,6 +204,7 @@ Each of these is one idea you now have the vocabulary for; each has its own page
204204
* `on_call_tool`, `on_get_prompt`, and `on_read_resource` may return an `InputRequiredResult` instead of their normal result to pause the call and ask the client for input; see **[Multi-round-trip requests](../handlers/multi-round-trip.md)**. True to this tier, nothing is installed for you: where `MCPServer` seals `requestState` by default, here the `request_state` you set crosses the wire exactly as written until you opt in with `server.middleware.append(RequestStateBoundary(RequestStateSecurity(keys=[...]), default_audience=server.name))`: one line (both names import from `mcp.server.request_state`) for the identical sealing and verification `MCPServer` performs (**[Protecting `requestState`](../handlers/multi-round-trip.md#protecting-requeststate)**).
205205
* `on_list_resources`, `on_read_resource`, `on_list_prompts`, `on_get_prompt`, `on_completion` are the same `(ctx, params) -> result` shape for the other primitives.
206206
* `on_subscriptions_listen` serves the 2026-07-28 `subscriptions/listen` stream. Pass a `ListenHandler` built over a `SubscriptionBus` and publish events to the bus from your other handlers; see **[Subscriptions](../handlers/subscriptions.md)** for the full composition.
207+
* `get_tool_input_schema` keeps `on_list_tools` off the call path; see **[Header parameters](header-parameters.md#schemas-by-name)**.
207208
* `server.streamable_http_app()` returns the same Starlette app `MCPServer`'s does; deploy it the way **[Running your server](../run/index.md)** deploys any other ASGI app. There is no `server.run(transport=...)` down here: `server.run(read_stream, write_stream, server.create_initialization_options())` drives one connection over a pair of streams, and that one line is the whole story.
208209

209210
## Recap
Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,50 @@
1+
from typing import Any
2+
3+
from mcp.server import Server, ServerRequestContext
4+
from mcp.types import (
5+
CallToolRequestParams,
6+
CallToolResult,
7+
ListToolsResult,
8+
PaginatedRequestParams,
9+
TextContent,
10+
Tool,
11+
)
12+
13+
CHECK_STOCK = Tool(
14+
name="check_stock",
15+
description="Count the copies of a book in one region's warehouses.",
16+
input_schema={
17+
"type": "object",
18+
"properties": {
19+
"title": {"type": "string"},
20+
"region": {"type": "string", "x-mcp-header": "Region"},
21+
},
22+
"required": ["title", "region"],
23+
},
24+
)
25+
26+
TOOLS = {CHECK_STOCK.name: CHECK_STOCK}
27+
28+
29+
async def list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams | None) -> ListToolsResult:
30+
return ListToolsResult(tools=list(TOOLS.values()))
31+
32+
33+
async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult:
34+
args = params.arguments or {}
35+
text = f"{args['title']}: 3 copies in {args['region']}."
36+
return CallToolResult(content=[TextContent(type="text", text=text)])
37+
38+
39+
def tool_input_schema(name: str) -> dict[str, Any] | None:
40+
tool = TOOLS.get(name)
41+
return tool.input_schema if tool else None
42+
43+
44+
server = Server(
45+
"Bookshop",
46+
on_list_tools=list_tools,
47+
on_call_tool=call_tool,
48+
get_tool_input_schema=tool_input_schema,
49+
)
50+
app = server.streamable_http_app()

‎src/mcp/server/_streamable_http_modern.py‎

Lines changed: 14 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -339,10 +339,12 @@ async def _mcp_param_rejection(
339339
"""Validate a `tools/call` request's `Mcp-Param-*` headers against the called tool's schema.
340340
341341
Runs pre-dispatch, before any SSE machinery, so a rejection is always a
342-
plain `application/json` 400 (the spec's MUST). With no `tools/list` handler
343-
the catalog is undiscoverable and there is no recognized header to validate.
342+
plain `application/json` 400 (the spec's MUST). The schema comes from the
343+
server's `get_tool_input_schema` when set, else from its `tools/list` handler;
344+
with neither there is no recognized header to validate.
344345
"""
345-
if req.method != "tools/call" or app.get_request_handler("tools/list") is None:
346+
lookup = app.get_tool_input_schema
347+
if req.method != "tools/call" or (lookup is None and app.get_request_handler("tools/list") is None):
346348
return None
347349
params = req.params or {}
348350
name = params.get("name")
@@ -356,7 +358,15 @@ async def _mcp_param_rejection(
356358
if not arguments and not any(header.startswith(_MCP_PARAM_PREFIX_LOWER) for header in request.headers):
357359
# No argument values and no `Mcp-Param-*` headers: no declaration can be violated either way.
358360
return None
359-
input_schema = await _tool_input_schema(app, request, req.id, verdict, lifespan_state, name)
361+
if lookup is None:
362+
input_schema = await _tool_input_schema(app, request, req.id, verdict, lifespan_state, name)
363+
else:
364+
try:
365+
input_schema = lookup(name)
366+
except Exception:
367+
# Fail-open like a failed listing: header validation must never break a working call path.
368+
logger.exception("Mcp-Param header validation skipped: get_tool_input_schema raised")
369+
return None
360370
if input_schema is None:
361371
return None
362372
return validate_mcp_param_headers(input_schema, arguments, request.headers)

‎src/mcp/server/lowlevel/server.py‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -143,6 +143,7 @@ def __init__(
143143
[Server[LifespanResultT]],
144144
AbstractAsyncContextManager[LifespanResultT],
145145
] = lifespan,
146+
get_tool_input_schema: Callable[[str], Mapping[str, Any] | None] | None = None,
146147
# Request handlers
147148
on_list_tools: Callable[
148149
[ServerRequestContext[LifespanResultT], types.PaginatedRequestParams | None],
@@ -226,6 +227,7 @@ def __init__(
226227
[Server[LifespanResultT]],
227228
AbstractAsyncContextManager[LifespanResultT],
228229
] = lifespan,
230+
get_tool_input_schema: Callable[[str], Mapping[str, Any] | None] | None = None,
229231
# Request handlers
230232
on_list_tools: Callable[
231233
[ServerRequestContext[LifespanResultT], types.PaginatedRequestParams | None],
@@ -318,6 +320,7 @@ def __init__(
318320
[Server[LifespanResultT]],
319321
AbstractAsyncContextManager[LifespanResultT],
320322
] = lifespan,
323+
get_tool_input_schema: Callable[[str], Mapping[str, Any] | None] | None = None,
321324
# Request handlers
322325
on_list_tools: Callable[
323326
[ServerRequestContext[LifespanResultT], types.PaginatedRequestParams | None],
@@ -425,6 +428,14 @@ def __init__(
425428
# after the handler returns; fields the handler set explicitly win.
426429
self.cache_hints: dict[str, CacheHint] = validate_cache_hints(cache_hints)
427430
self.lifespan = lifespan
431+
self.get_tool_input_schema = get_tool_input_schema
432+
"""Returns a tool's input schema by name, or `None` when there is nothing to validate.
433+
434+
When set, `Mcp-Param-*` header validation on the 2026-07-28 Streamable HTTP
435+
path calls this instead of running the `tools/list` handler. It is called
436+
before middleware runs, so it is not scoped to the caller. If it raises,
437+
the error is logged and the call is served unvalidated.
438+
"""
428439
self._request_handlers: dict[str, HandlerEntry[LifespanResultT]] = {}
429440
self._notification_handlers: dict[str, HandlerEntry[LifespanResultT]] = {}
430441
self._session_manager: StreamableHTTPSessionManager | None = None

‎src/mcp/server/mcpserver/server.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -218,6 +218,7 @@ def __init__(
218218
icons=icons,
219219
version=version,
220220
cache_hints=cache_hints,
221+
get_tool_input_schema=self._tool_input_schema,
221222
on_list_tools=self._handle_list_tools,
222223
on_call_tool=self._handle_call_tool,
223224
on_list_resources=self._handle_list_resources,
@@ -428,6 +429,11 @@ async def _handle_list_tools(
428429
) -> ListToolsResult:
429430
return ListToolsResult(tools=await self.list_tools())
430431

432+
def _tool_input_schema(self, name: str) -> dict[str, Any] | None:
433+
"""Called before middleware runs, so it also finds a tool that middleware hides from the caller."""
434+
tool = self._tool_manager.get_tool(name)
435+
return None if tool is None else tool.parameters
436+
431437
async def _handle_call_tool(
432438
self, ctx: ServerRequestContext[LifespanResultT], params: CallToolRequestParams
433439
) -> CallToolResult | InputRequiredResult:

‎tests/docs_src/test_header_parameters.py‎

Lines changed: 35 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,18 +2,19 @@
22

33
from collections.abc import AsyncIterator
44
from contextlib import asynccontextmanager
5-
from typing import Annotated, Literal
5+
from typing import Annotated, Any, Literal
66

77
import httpx2
88
import pytest
99
from mcp_types import HEADER_MISMATCH, ListToolsResult, PaginatedRequestParams
1010
from pydantic import Field, WithJsonSchema
1111
from starlette.applications import Starlette
1212

13-
from docs_src.header_parameters import tutorial001, tutorial002
13+
from docs_src.header_parameters import tutorial001, tutorial002, tutorial003
1414
from mcp import Client
1515
from mcp.client.streamable_http import streamable_http_client
1616
from mcp.server import MCPServer, Server, ServerRequestContext
17+
from mcp.server.context import CallNext, HandlerResult
1718
from mcp.server.mcpserver.exceptions import InvalidSignature
1819

1920
# See test_index.py for why this is a per-module mark and not a conftest hook.
@@ -165,3 +166,35 @@ async def list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams |
165166
async with Client(server) as modern:
166167
assert modern.protocol_version == "2026-07-28"
167168
assert (await modern.list_tools()).tools == []
169+
170+
171+
@pytest.mark.parametrize(
172+
("server", "expected"),
173+
[(tutorial002.server, ["tools/list", "tools/call"]), (tutorial003.server, ["tools/call"])],
174+
ids=["tutorial002", "tutorial003"],
175+
)
176+
async def test_a_call_runs_the_list_handler_unless_the_server_looks_schemas_up_by_name(
177+
server: Server, expected: list[str], monkeypatch: pytest.MonkeyPatch
178+
) -> None:
179+
"""tutorial002 and tutorial003: the client's own `tools/call`, replayed, dispatches a `tools/list` first
180+
on the server without `get_tool_input_schema` and only itself on the server with it."""
181+
dispatched: list[str] = []
182+
183+
async def record(ctx: ServerRequestContext[Any, Any], call_next: CallNext) -> HandlerResult:
184+
dispatched.append(ctx.method)
185+
return await call_next(ctx)
186+
187+
monkeypatch.setattr(server, "middleware", [*server.middleware, record])
188+
async with check_stock_over_http(server.streamable_http_app()) as (http, call):
189+
dispatched.clear()
190+
replayed = await http.post(URL, content=call.content, headers=call.headers)
191+
assert replayed.status_code == 200
192+
assert dispatched == expected
193+
194+
195+
async def test_the_schema_the_lookup_returns_is_the_one_the_header_is_checked_against() -> None:
196+
"""tutorial003: the client's own request, replayed with a different `Mcp-Param-Region`, is a 400."""
197+
async with check_stock_over_http(tutorial003.app) as (http, call):
198+
tampered = await http.post(URL, content=call.content, headers={**call.headers, "mcp-param-region": "us"})
199+
assert tampered.status_code == 400
200+
assert tampered.json()["error"]["code"] == HEADER_MISMATCH

‎tests/server/test_streamable_http_modern.py‎

Lines changed: 99 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
import json
1010
import logging
1111
from collections.abc import Callable
12-
from typing import Any
12+
from typing import Annotated, Any
1313

1414
import anyio
1515
import httpx2
@@ -40,6 +40,7 @@
4040
Tool,
4141
)
4242
from mcp_types.version import LATEST_MODERN_VERSION, MODERN_PROTOCOL_VERSIONS
43+
from pydantic import Field
4344
from starlette.types import Message, Receive, Scope, Send
4445
from trio.testing import MockClock
4546

@@ -49,6 +50,8 @@
4950
_to_jsonrpc_response,
5051
handle_modern_request,
5152
)
53+
from mcp.server.context import CallNext, HandlerResult
54+
from mcp.server.mcpserver import MCPServer
5255
from mcp.server.subscriptions import InMemorySubscriptionBus, ListenHandler, ServerEvent
5356
from mcp.server.transport_security import TransportSecuritySettings
5457
from mcp.shared.exceptions import MCPError, NoBackChannelError
@@ -1007,6 +1010,101 @@ async def recording_list(ctx: ServerRequestContext, params: PaginatedRequestPara
10071010
assert seen[0].client_info.name == "raw"
10081011

10091012

1013+
async def test_modern_tools_call_on_mcpserver_validates_mcp_param_headers_without_listing() -> None:
1014+
"""SDK-defined: `MCPServer` reads the called tool's schema from its registry, so the spec's
1015+
`Mcp-Param-*` verdicts hold while no `tools/list` runs for a `tools/call`. Raw HTTP because
1016+
a `Client` cannot send a mismatched header and issues listings of its own."""
1017+
dispatched: list[str] = []
1018+
1019+
async def record(ctx: ServerRequestContext[Any, Any], call_next: CallNext) -> HandlerResult:
1020+
dispatched.append(ctx.method)
1021+
return await call_next(ctx)
1022+
1023+
mcp = MCPServer("test", middleware=[record])
1024+
1025+
@mcp.tool()
1026+
def search(region: Annotated[str, Field(json_schema_extra={"x-mcp-header": "Region"})]) -> str:
1027+
return region
1028+
1029+
body = _tool_call_body({"region": "eu"})
1030+
async with _asgi_client(mcp._lowlevel_server) as http:
1031+
matched = await http.post("/mcp", json=body, headers=_TOOL_CALL_HEADERS | {"mcp-param-region": "eu"})
1032+
mismatched = await http.post("/mcp", json=body, headers=_TOOL_CALL_HEADERS | {"mcp-param-region": "us"})
1033+
missing = await http.post("/mcp", json=body, headers=_TOOL_CALL_HEADERS)
1034+
unknown = await http.post(
1035+
"/mcp",
1036+
json=_tool_call_body({"region": "eu"}, name="unregistered"),
1037+
headers={MCP_METHOD_HEADER: "tools/call", MCP_NAME_HEADER: "unregistered", "mcp-param-region": "us"},
1038+
)
1039+
1040+
assert matched.status_code == 200
1041+
assert matched.json()["result"]["structuredContent"] == {"result": "eu"}
1042+
assert (mismatched.status_code, mismatched.json()["error"]["code"]) == (400, HEADER_MISMATCH)
1043+
assert (missing.status_code, missing.json()["error"]["code"]) == (400, HEADER_MISMATCH)
1044+
# An unregistered tool has no schema to validate against; dispatch owns the unknown-tool answer.
1045+
assert unknown.status_code == 200
1046+
assert unknown.json()["result"]["isError"] is True
1047+
assert dispatched == ["tools/call", "tools/call"]
1048+
1049+
1050+
async def test_modern_tools_call_asks_get_tool_input_schema_instead_of_listing() -> None:
1051+
"""SDK-defined: a low-level server that passes `get_tool_input_schema` is asked for the called
1052+
tool's schema by name, so the spec's `Mcp-Param-*` verdicts hold while its `tools/list` handler
1053+
never runs for a `tools/call`."""
1054+
dispatched: list[str] = []
1055+
1056+
async def record(ctx: ServerRequestContext[Any, Any], call_next: CallNext) -> HandlerResult:
1057+
dispatched.append(ctx.method)
1058+
return await call_next(ctx)
1059+
1060+
async def list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams | None) -> ListToolsResult:
1061+
raise NotImplementedError
1062+
1063+
schemas = {_REGION_TOOL.name: _REGION_TOOL.input_schema}
1064+
server: Server[Any] = Server(
1065+
"gateway", on_list_tools=list_tools, on_call_tool=_ok_call_tool, get_tool_input_schema=schemas.get
1066+
)
1067+
server.middleware.append(record)
1068+
1069+
body = _tool_call_body({"region": "eu"})
1070+
async with _asgi_client(server) as http:
1071+
matched = await http.post("/mcp", json=body, headers=_TOOL_CALL_HEADERS | {"mcp-param-region": "eu"})
1072+
mismatched = await http.post("/mcp", json=body, headers=_TOOL_CALL_HEADERS | {"mcp-param-region": "us"})
1073+
missing = await http.post("/mcp", json=body, headers=_TOOL_CALL_HEADERS)
1074+
unknown = await http.post(
1075+
"/mcp",
1076+
json=_tool_call_body({"region": "eu"}, name="unregistered"),
1077+
headers={MCP_METHOD_HEADER: "tools/call", MCP_NAME_HEADER: "unregistered", "mcp-param-region": "us"},
1078+
)
1079+
1080+
assert matched.status_code == 200
1081+
assert (mismatched.status_code, mismatched.json()["error"]["code"]) == (400, HEADER_MISMATCH)
1082+
assert (missing.status_code, missing.json()["error"]["code"]) == (400, HEADER_MISMATCH)
1083+
# `None` from the lookup means nothing to validate; the mismatched header is ignored.
1084+
assert unknown.status_code == 200
1085+
assert dispatched == ["tools/call", "tools/call"]
1086+
1087+
1088+
async def test_modern_tools_call_skips_validation_when_get_tool_input_schema_raises(
1089+
caplog: pytest.LogCaptureFixture,
1090+
) -> None:
1091+
"""A raising `get_tool_input_schema` fails open like a raising listing: the call is served and the skip logged."""
1092+
1093+
def unavailable(name: str) -> dict[str, Any] | None:
1094+
raise RuntimeError("catalog unavailable")
1095+
1096+
server: Server[Any] = Server("gateway", on_call_tool=_ok_call_tool, get_tool_input_schema=unavailable)
1097+
with caplog.at_level(logging.ERROR, logger=_streamable_http_modern.__name__):
1098+
async with _asgi_client(server) as http:
1099+
response = await http.post(
1100+
"/mcp",
1101+
json=_tool_call_body({"region": "us"}),
1102+
headers=_TOOL_CALL_HEADERS | {"mcp-param-region": "eu"},
1103+
)
1104+
assert response.status_code == 200
1105+
assert "Mcp-Param header validation skipped: get_tool_input_schema raised" in caplog.text
1106+
1107+
10101108
async def test_modern_tools_call_leaves_mis_shaped_name_and_arguments_to_dispatch() -> None:
10111109
"""A missing `name` or non-mapping `arguments` is dispatch's INVALID_PARAMS, never a header mismatch."""
10121110
async with _asgi_client(_x_mcp_server()) as http:

0 commit comments

Comments
 (0)