diff --git a/docs/advanced/middleware.md b/docs/advanced/middleware.md index 6865ae2112..46a5f98fce 100644 --- a/docs/advanced/middleware.md +++ b/docs/advanced/middleware.md @@ -10,8 +10,8 @@ You write it as `async (ctx, call_next)` and append it to `server.middleware`. T *refuse* messages; do not make it the foundation your server stands on. `MCPServer` takes the list at construction (`MCPServer(name, middleware=[...])`) and exposes it as -`mcp.middleware`; the low-level `Server` exposes the same list as `server.middleware`. The example -below uses the low-level `Server`; if `Server(name, on_call_tool=...)` is new to you, read +`mcp.middleware`; the low-level `Server` exposes the same list as `server.middleware`. The examples +below use the low-level `Server`; if `Server(name, on_call_tool=...)` is new to you, read **[The low-level Server](low-level-server.md)** first. ## A timing middleware @@ -56,14 +56,35 @@ That is the point. Middleware wraps **every** inbound message: * Even a method the server has no handler for: `call_next` raises the `MCPError(-32601, "Method not found")` *through* your middleware on its way to the client. +## A concurrency cap + +A middleware doesn't have to call `call_next(ctx)`. Raise an `MCPError` instead and that one +message is **refused**: the connection stays up and the next message goes through. + +Say every search holds a connection from a pool of four. This middleware lets four tool calls run +at once and refuses the fifth: + +```python title="server.py" hl_lines="15-16 40-55 59" +--8<-- "docs_src/middleware/tutorial002.py" +``` + +* Only `tools/call` is counted, so the server keeps answering `server/discover` and `tools/list` + while it refuses tool calls. +* MCP defines no "server busy" error code, so `SERVER_BUSY` is this server's own. +* Refusing tells the client straight away that the server is overloaded. If you'd rather make + callers wait, hold an `anyio.CapacityLimiter` around `call_next(ctx)` instead. + +A raised `MCPError` goes to the client application, not to the model. If the model should read the +message, return a tool result with `is_error=True` instead: that is **Answer**, below. + ## What you can do inside one In increasing order of how much you should hesitate: -* **Observe.** Time it, count it, log it. The example above. +* **Observe.** Time it, count it, log it. The timing middleware above. * **Refuse.** Raise an `MCPError` *instead of* calling `call_next(ctx)` and that one message is - answered with a JSON-RPC error. The connection stays up; the next message goes through. This is - how a server gates `subscriptions/listen` per caller: + answered with a JSON-RPC error. The connection stays up; the next message goes through. The + concurrency cap above. It is also how a server gates `subscriptions/listen` per caller: **[Deciding who may watch](../handlers/subscriptions.md#deciding-who-may-watch)** on the Subscriptions page walks through it. * **Rewrite.** `ctx` is a dataclass: `await call_next(dataclasses.replace(ctx, params=...))` diff --git a/docs_src/middleware/tutorial002.py b/docs_src/middleware/tutorial002.py new file mode 100644 index 0000000000..bfc72c6920 --- /dev/null +++ b/docs_src/middleware/tutorial002.py @@ -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"] + return CallToolResult(content=[TextContent(type="text", text=f"Found 3 books matching {query!r}.")]) + + +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)) diff --git a/tests/docs_src/test_middleware.py b/tests/docs_src/test_middleware.py index 97d9e96086..39a635e12d 100644 --- a/tests/docs_src/test_middleware.py +++ b/tests/docs_src/test_middleware.py @@ -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.")] + + +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