|
9 | 9 | import json |
10 | 10 | import logging |
11 | 11 | from collections.abc import Callable |
12 | | -from typing import Any |
| 12 | +from typing import Annotated, Any |
13 | 13 |
|
14 | 14 | import anyio |
15 | 15 | import httpx2 |
|
40 | 40 | Tool, |
41 | 41 | ) |
42 | 42 | from mcp_types.version import LATEST_MODERN_VERSION, MODERN_PROTOCOL_VERSIONS |
| 43 | +from pydantic import Field |
43 | 44 | from starlette.types import Message, Receive, Scope, Send |
44 | 45 | from trio.testing import MockClock |
45 | 46 |
|
|
49 | 50 | _to_jsonrpc_response, |
50 | 51 | handle_modern_request, |
51 | 52 | ) |
| 53 | +from mcp.server.context import CallNext, HandlerResult |
| 54 | +from mcp.server.mcpserver import MCPServer |
52 | 55 | from mcp.server.subscriptions import InMemorySubscriptionBus, ListenHandler, ServerEvent |
53 | 56 | from mcp.server.transport_security import TransportSecuritySettings |
54 | 57 | from mcp.shared.exceptions import MCPError, NoBackChannelError |
@@ -1007,6 +1010,101 @@ async def recording_list(ctx: ServerRequestContext, params: PaginatedRequestPara |
1007 | 1010 | assert seen[0].client_info.name == "raw" |
1008 | 1011 |
|
1009 | 1012 |
|
| 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 | + |
1010 | 1108 | async def test_modern_tools_call_leaves_mis_shaped_name_and_arguments_to_dispatch() -> None: |
1011 | 1109 | """A missing `name` or non-mapping `arguments` is dispatch's INVALID_PARAMS, never a header mismatch.""" |
1012 | 1110 | async with _asgi_client(_x_mcp_server()) as http: |
|
0 commit comments