|
2 | 2 |
|
3 | 3 | from collections.abc import AsyncIterator |
4 | 4 | from contextlib import asynccontextmanager |
5 | | -from typing import Annotated, Literal |
| 5 | +from typing import Annotated, Any, Literal |
6 | 6 |
|
7 | 7 | import httpx2 |
8 | 8 | import pytest |
9 | 9 | from mcp_types import HEADER_MISMATCH, ListToolsResult, PaginatedRequestParams |
10 | 10 | from pydantic import Field, WithJsonSchema |
11 | 11 | from starlette.applications import Starlette |
12 | 12 |
|
13 | | -from docs_src.header_parameters import tutorial001, tutorial002 |
| 13 | +from docs_src.header_parameters import tutorial001, tutorial002, tutorial003 |
14 | 14 | from mcp import Client |
15 | 15 | from mcp.client.streamable_http import streamable_http_client |
16 | 16 | from mcp.server import MCPServer, Server, ServerRequestContext |
| 17 | +from mcp.server.context import CallNext, HandlerResult |
17 | 18 | from mcp.server.mcpserver.exceptions import InvalidSignature |
18 | 19 |
|
19 | 20 | # See test_index.py for why this is a per-module mark and not a conftest hook. |
@@ -140,3 +141,35 @@ async def list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams | |
140 | 141 | async with Client(server) as modern: |
141 | 142 | assert modern.protocol_version == "2026-07-28" |
142 | 143 | assert (await modern.list_tools()).tools == [] |
| 144 | + |
| 145 | + |
| 146 | +@pytest.mark.parametrize( |
| 147 | + ("server", "expected"), |
| 148 | + [(tutorial002.server, ["tools/list", "tools/call"]), (tutorial003.server, ["tools/call"])], |
| 149 | + ids=["tutorial002", "tutorial003"], |
| 150 | +) |
| 151 | +async def test_a_call_runs_the_list_handler_unless_the_server_looks_schemas_up_by_name( |
| 152 | + server: Server, expected: list[str], monkeypatch: pytest.MonkeyPatch |
| 153 | +) -> None: |
| 154 | + """tutorial002 and tutorial003: the client's own `tools/call`, replayed, dispatches a `tools/list` first |
| 155 | + on the server without `get_tool_input_schema` and only itself on the server with it.""" |
| 156 | + dispatched: list[str] = [] |
| 157 | + |
| 158 | + async def record(ctx: ServerRequestContext[Any, Any], call_next: CallNext) -> HandlerResult: |
| 159 | + dispatched.append(ctx.method) |
| 160 | + return await call_next(ctx) |
| 161 | + |
| 162 | + monkeypatch.setattr(server, "middleware", [*server.middleware, record]) |
| 163 | + async with check_stock_over_http(server.streamable_http_app()) as (http, call): |
| 164 | + dispatched.clear() |
| 165 | + replayed = await http.post(URL, content=call.content, headers=call.headers) |
| 166 | + assert replayed.status_code == 200 |
| 167 | + assert dispatched == expected |
| 168 | + |
| 169 | + |
| 170 | +async def test_the_schema_the_lookup_returns_is_the_one_the_header_is_checked_against() -> None: |
| 171 | + """tutorial003: the client's own request, replayed with a different `Mcp-Param-Region`, is a 400.""" |
| 172 | + async with check_stock_over_http(tutorial003.app) as (http, call): |
| 173 | + tampered = await http.post(URL, content=call.content, headers={**call.headers, "mcp-param-region": "us"}) |
| 174 | + assert tampered.status_code == 400 |
| 175 | + assert tampered.json()["error"]["code"] == HEADER_MISMATCH |
0 commit comments