-
Notifications
You must be signed in to change notification settings - Fork 4k
docs: add a middleware tutorial that refuses tool calls #3613
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change | ||||
|---|---|---|---|---|---|---|
| @@ -0,0 +1,59 @@ | ||||||
| from typing import Any | ||||||
|
|
||||||
| from mcp import MCPError | ||||||
| from mcp.server import Server, ServerRequestContext | ||||||
| from mcp.server.context import CallNext, HandlerResult, ServerMiddleware | ||||||
| from mcp.types import ( | ||||||
| CallToolRequestParams, | ||||||
| CallToolResult, | ||||||
| ListToolsResult, | ||||||
| PaginatedRequestParams, | ||||||
| TextContent, | ||||||
| Tool, | ||||||
| ) | ||||||
|
|
||||||
| # MCP defines no "busy" error, so this server picks its own code. | ||||||
| SERVER_BUSY = 1 | ||||||
|
|
||||||
|
|
||||||
| async def on_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams | None) -> ListToolsResult: | ||||||
| return ListToolsResult( | ||||||
| tools=[ | ||||||
| Tool( | ||||||
| name="search_books", | ||||||
| description="Search the catalog by title or author.", | ||||||
| input_schema={ | ||||||
| "type": "object", | ||||||
| "properties": {"query": {"type": "string"}}, | ||||||
| "required": ["query"], | ||||||
| }, | ||||||
| ) | ||||||
| ] | ||||||
| ) | ||||||
|
|
||||||
|
|
||||||
| async def on_call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: | ||||||
| query = (params.arguments or {})["query"] | ||||||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. P3: Prompt for AI agents
Suggested change
|
||||||
| return CallToolResult(content=[TextContent(type="text", text=f"Found 3 books matching {query!r}.")]) | ||||||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. P2: This sample handler never awaits any work, so it releases Prompt for AI agents |
||||||
|
|
||||||
|
|
||||||
| def max_concurrent_tool_calls(limit: int) -> ServerMiddleware[Any]: | ||||||
| running = 0 | ||||||
|
|
||||||
| async def middleware(ctx: ServerRequestContext, call_next: CallNext) -> HandlerResult: | ||||||
| nonlocal running | ||||||
| if ctx.method != "tools/call": | ||||||
| return await call_next(ctx) | ||||||
| if running >= limit: | ||||||
| raise MCPError(code=SERVER_BUSY, message=f"Server busy: tool call limit reached ({limit} in progress).") | ||||||
| running += 1 | ||||||
| try: | ||||||
| return await call_next(ctx) | ||||||
| finally: | ||||||
| running -= 1 | ||||||
|
|
||||||
| return middleware | ||||||
|
|
||||||
|
|
||||||
| server = Server("Bookshop", on_list_tools=on_list_tools, on_call_tool=on_call_tool) | ||||||
| server.middleware.append(max_concurrent_tool_calls(4)) | ||||||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -3,17 +3,20 @@ | |
| import logging | ||
| import re | ||
|
|
||
| import anyio | ||
| import pytest | ||
| from mcp_types import ( | ||
| INTERNAL_ERROR, | ||
| INVALID_REQUEST, | ||
| METHOD_NOT_FOUND, | ||
| CallToolRequestParams, | ||
| CallToolResult, | ||
| ErrorData, | ||
| RequestId, | ||
| TextContent, | ||
| ) | ||
|
|
||
| from docs_src.middleware import tutorial001 | ||
| from docs_src.middleware import tutorial001, tutorial002 | ||
| from mcp import Client, MCPError | ||
| from mcp.server import Server, ServerRequestContext | ||
| from mcp.server.context import CallNext, HandlerResult | ||
|
|
@@ -27,6 +30,29 @@ def _is_timing_record(record: logging.LogRecord) -> bool: | |
| return record.name == tutorial001.logger.name | ||
|
|
||
|
|
||
| class _HeldSearches: | ||
| """An `on_call_tool` for `search_books` that holds each call open until the test finishes it, keyed by query. | ||
| A query the test did not name is a `KeyError`: a call the cap should have refused cannot quietly succeed. | ||
| """ | ||
|
|
||
| def __init__(self, *queries: str) -> None: | ||
| self.started = {query: anyio.Event() for query in queries} | ||
| self.finish = {query: anyio.Event() for query in queries} | ||
|
|
||
| async def __call__(self, ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: | ||
| assert params.name == "search_books" | ||
| query = (params.arguments or {})["query"] | ||
| self.started[query].set() | ||
| await self.finish[query].wait() | ||
| return CallToolResult(content=[TextContent(type="text", text=f"Found {query}.")]) | ||
|
|
||
|
|
||
| async def _search(client: Client, query: str, results: dict[str, CallToolResult]) -> None: | ||
| """Call `search_books` and file the result under its query, for calls a test runs in the background.""" | ||
| results[query] = await client.call_tool("search_books", {"query": query}) | ||
|
|
||
|
|
||
| def test_timing_record_predicate() -> None: | ||
| """The caplog filter keeps the middleware's own records and drops everyone else's.""" | ||
| args = (logging.INFO, __file__, 1, "msg", None, None) | ||
|
|
@@ -114,3 +140,157 @@ async def test_initialize_cannot_be_replaced_only_wrapped() -> None: | |
| ) | ||
| with pytest.raises(ValueError, match=re.escape(expected)): | ||
| tutorial001.server.add_request_handler("initialize", CallToolRequestParams, tutorial001.on_call_tool) | ||
|
|
||
|
|
||
| async def test_a_tool_call_over_the_cap_is_refused_while_the_earlier_calls_still_run() -> None: | ||
| """tutorial002: with `limit` tool calls in flight, the next one is answered with the busy error at once. | ||
| Steps: | ||
| 1. One client's four calls enter the handler and are held there. | ||
| 2. A second client connects (its `server/discover` is not a tool call) and makes a fifth call, which is | ||
| refused before any of the four has returned: the count belongs to the server, not to a connection. | ||
| 3. The four are finished and each returns its own result. | ||
| """ | ||
| held = _HeldSearches("dune", "emma", "ulysses", "walden") | ||
| server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=held) | ||
| server.middleware.append(tutorial002.max_concurrent_tool_calls(4)) | ||
| results: dict[str, CallToolResult] = {} | ||
| with anyio.fail_after(5): | ||
| async with Client(server) as client: | ||
| async with anyio.create_task_group() as tg: | ||
| for query in held.started: | ||
| tg.start_soon(_search, client, query, results) | ||
| for started in held.started.values(): | ||
| await started.wait() | ||
| async with Client(server) as latecomer: | ||
| with pytest.raises(MCPError) as exc_info: | ||
| await latecomer.call_tool("search_books", {"query": "middlemarch"}) | ||
| assert exc_info.value.error == ErrorData( | ||
| code=tutorial002.SERVER_BUSY, message="Server busy: tool call limit reached (4 in progress)." | ||
| ) | ||
| assert results == {} | ||
| for finish in held.finish.values(): | ||
| finish.set() | ||
| assert {query: result.content for query, result in results.items()} == { | ||
| query: [TextContent(type="text", text=f"Found {query}.")] for query in held.started | ||
| } | ||
|
|
||
|
|
||
| async def test_other_requests_are_answered_with_the_cap_reached() -> None: | ||
| """tutorial002: only `tools/call` is capped, so `tools/list` is answered with the cap reached.""" | ||
| held = _HeldSearches("dune") | ||
| server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=held) | ||
| server.middleware.append(tutorial002.max_concurrent_tool_calls(1)) | ||
| results: dict[str, CallToolResult] = {} | ||
| with anyio.fail_after(5): | ||
| async with Client(server) as client: | ||
| async with anyio.create_task_group() as tg: | ||
| tg.start_soon(_search, client, "dune", results) | ||
| await held.started["dune"].wait() | ||
| tools = (await client.list_tools()).tools | ||
| assert [tool.name for tool in tools] == ["search_books"] | ||
| assert results == {} | ||
| held.finish["dune"].set() | ||
| assert results["dune"].content == [TextContent(type="text", text="Found dune.")] | ||
|
|
||
|
|
||
| async def test_a_refused_call_is_accepted_once_a_running_call_finishes() -> None: | ||
| """tutorial002: a finished call gives its slot back, so the caller that was refused gets through on a retry.""" | ||
| held = _HeldSearches("dune", "emma") | ||
| server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=held) | ||
| server.middleware.append(tutorial002.max_concurrent_tool_calls(1)) | ||
| results: dict[str, CallToolResult] = {} | ||
| # Only `dune` is held open; `emma` returns as soon as it is let in. | ||
| held.finish["emma"].set() | ||
| with anyio.fail_after(5): | ||
| async with Client(server) as client: | ||
| async with anyio.create_task_group() as tg: | ||
| tg.start_soon(_search, client, "dune", results) | ||
| await held.started["dune"].wait() | ||
| with pytest.raises(MCPError) as exc_info: | ||
| await client.call_tool("search_books", {"query": "emma"}) | ||
| assert exc_info.value.error.code == tutorial002.SERVER_BUSY | ||
| assert not held.started["emma"].is_set() | ||
| held.finish["dune"].set() | ||
| # The task group has joined, so the first call's response has arrived. | ||
| retried = await client.call_tool("search_books", {"query": "emma"}) | ||
| assert results["dune"].content == [TextContent(type="text", text="Found dune.")] | ||
| assert retried.content == [TextContent(type="text", text="Found emma.")] | ||
|
|
||
|
|
||
| async def test_a_tool_call_that_raises_gives_its_slot_back() -> None: | ||
| """tutorial002: the `finally` releases the slot when the handler raises, so the next call is not refused.""" | ||
|
|
||
| async def on_call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: | ||
| assert params.name == "search_books" | ||
| query = (params.arguments or {})["query"] | ||
| if query == "necronomicon": | ||
| raise RuntimeError("the shelf collapsed") | ||
| return CallToolResult(content=[TextContent(type="text", text=f"Found {query}.")]) | ||
|
|
||
| server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=on_call_tool) | ||
| server.middleware.append(tutorial002.max_concurrent_tool_calls(1)) | ||
| async with Client(server) as client: | ||
| with pytest.raises(MCPError) as exc_info: | ||
| await client.call_tool("search_books", {"query": "necronomicon"}) | ||
| assert exc_info.value.error.code == INTERNAL_ERROR | ||
| result = await client.call_tool("search_books", {"query": "dune"}) | ||
| assert result.content == [TextContent(type="text", text="Found dune.")] | ||
|
|
||
|
|
||
| async def test_a_cancelled_tool_call_gives_its_slot_back() -> None: | ||
| """tutorial002: a call the client abandons is cancelled out of `call_next`, and the `finally` frees its slot.""" | ||
| started = anyio.Event() | ||
| cancelled = anyio.Event() | ||
|
|
||
| async def on_call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: | ||
| assert params.name == "search_books" | ||
| query = (params.arguments or {})["query"] | ||
| if query == "dune": | ||
| started.set() | ||
| try: | ||
| await anyio.sleep_forever() | ||
| finally: | ||
| cancelled.set() | ||
| return CallToolResult(content=[TextContent(type="text", text=f"Found {query}.")]) | ||
|
|
||
| server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=on_call_tool) | ||
| server.middleware.append(tutorial002.max_concurrent_tool_calls(1)) | ||
| results: dict[str, CallToolResult] = {} | ||
| with anyio.fail_after(5): | ||
| async with Client(server) as client: | ||
| async with anyio.create_task_group() as tg: | ||
| tg.start_soon(_search, client, "dune", results) | ||
| await started.wait() | ||
| tg.cancel_scope.cancel() | ||
| await cancelled.wait() | ||
| result = await client.call_tool("search_books", {"query": "emma"}) | ||
| assert results == {} | ||
| assert result.content == [TextContent(type="text", text="Found emma.")] | ||
|
|
||
|
|
||
| async def test_answering_with_an_error_result_returns_it_to_the_caller_instead_of_raising() -> None: | ||
| """A middleware that answers `tools/call` with an `is_error=True` result gives the client a result, not an error.""" | ||
|
|
||
| async def busy(ctx: ServerRequestContext, call_next: CallNext) -> HandlerResult: | ||
| if ctx.method == "tools/call": | ||
| return CallToolResult( | ||
| content=[TextContent(type="text", text="Server busy. Try again shortly.")], is_error=True | ||
| ) | ||
| return await call_next(ctx) | ||
|
|
||
| server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=tutorial002.on_call_tool) | ||
| server.middleware.append(busy) | ||
| async with Client(server) as client: | ||
| result = await client.call_tool("search_books", {"query": "dune"}) | ||
| assert result.is_error | ||
| assert result.content == [TextContent(type="text", text="Server busy. Try again shortly.")] | ||
|
|
||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🟡 (optional) A reviewer following the repo's test rules finds a registered handler that can never run in Why this was flaggedIn tests/docs_src/test_middleware.py:279-293 the Verification: nit. Triggering condition: any run of |
||
|
|
||
| async def test_calls_made_one_after_another_never_reach_the_cap() -> None: | ||
| """tutorial002's own server: the cap counts calls in flight, so more than four in sequence all succeed.""" | ||
| async with Client(tutorial002.server) as client: | ||
| results = [await client.call_tool("search_books", {"query": "dune"}) for _ in range(5)] | ||
| assert [result.content for result in results] == [ | ||
| [TextContent(type="text", text="Found 3 books matching 'dune'.")] | ||
| ] * 5 | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
P3: The Refuse bullet contains the fragment
The concurrency cap above., so the new cross-reference reads as incomplete. Change it toThis is the concurrency cap above..Prompt for AI agents