Skip to content

Commit 86ec15f

Browse files
authored
docs: add a middleware tutorial that refuses tool calls (#3613)
1 parent 17aaf25 commit 86ec15f

3 files changed

Lines changed: 266 additions & 6 deletions

File tree

‎docs/advanced/middleware.md‎

Lines changed: 26 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -10,8 +10,8 @@ You write it as `async (ctx, call_next)` and append it to `server.middleware`. T
1010
*refuse* messages; do not make it the foundation your server stands on.
1111

1212
`MCPServer` takes the list at construction (`MCPServer(name, middleware=[...])`) and exposes it as
13-
`mcp.middleware`; the low-level `Server` exposes the same list as `server.middleware`. The example
14-
below uses the low-level `Server`; if `Server(name, on_call_tool=...)` is new to you, read
13+
`mcp.middleware`; the low-level `Server` exposes the same list as `server.middleware`. The examples
14+
below use the low-level `Server`; if `Server(name, on_call_tool=...)` is new to you, read
1515
**[The low-level Server](low-level-server.md)** first.
1616

1717
## A timing middleware
@@ -56,14 +56,35 @@ That is the point. Middleware wraps **every** inbound message:
5656
* Even a method the server has no handler for: `call_next` raises the
5757
`MCPError(-32601, "Method not found")` *through* your middleware on its way to the client.
5858

59+
## A concurrency cap
60+
61+
A middleware doesn't have to call `call_next(ctx)`. Raise an `MCPError` instead and that one
62+
message is **refused**: the connection stays up and the next message goes through.
63+
64+
Say every search holds a connection from a pool of four. This middleware lets four tool calls run
65+
at once and refuses the fifth:
66+
67+
```python title="server.py" hl_lines="15-16 40-55 59"
68+
--8<-- "docs_src/middleware/tutorial002.py"
69+
```
70+
71+
* Only `tools/call` is counted, so the server keeps answering `server/discover` and `tools/list`
72+
while it refuses tool calls.
73+
* MCP defines no "server busy" error code, so `SERVER_BUSY` is this server's own.
74+
* Refusing tells the client straight away that the server is overloaded. If you'd rather make
75+
callers wait, hold an `anyio.CapacityLimiter` around `call_next(ctx)` instead.
76+
77+
A raised `MCPError` goes to the client application, not to the model. If the model should read the
78+
message, return a tool result with `is_error=True` instead: that is **Answer**, below.
79+
5980
## What you can do inside one
6081

6182
In increasing order of how much you should hesitate:
6283

63-
* **Observe.** Time it, count it, log it. The example above.
84+
* **Observe.** Time it, count it, log it. The timing middleware above.
6485
* **Refuse.** Raise an `MCPError` *instead of* calling `call_next(ctx)` and that one message is
65-
answered with a JSON-RPC error. The connection stays up; the next message goes through. This is
66-
how a server gates `subscriptions/listen` per caller:
86+
answered with a JSON-RPC error. The connection stays up; the next message goes through. The
87+
concurrency cap above. It is also how a server gates `subscriptions/listen` per caller:
6788
**[Deciding who may watch](../handlers/subscriptions.md#deciding-who-may-watch)** on the
6889
Subscriptions page walks through it.
6990
* **Rewrite.** `ctx` is a dataclass: `await call_next(dataclasses.replace(ctx, params=...))`

‎docs_src/middleware/tutorial002.py‎

Lines changed: 59 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,59 @@
1+
from typing import Any
2+
3+
from mcp import MCPError
4+
from mcp.server import Server, ServerRequestContext
5+
from mcp.server.context import CallNext, HandlerResult, ServerMiddleware
6+
from mcp.types import (
7+
CallToolRequestParams,
8+
CallToolResult,
9+
ListToolsResult,
10+
PaginatedRequestParams,
11+
TextContent,
12+
Tool,
13+
)
14+
15+
# MCP defines no "busy" error, so this server picks its own code.
16+
SERVER_BUSY = 1
17+
18+
19+
async def on_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams | None) -> ListToolsResult:
20+
return ListToolsResult(
21+
tools=[
22+
Tool(
23+
name="search_books",
24+
description="Search the catalog by title or author.",
25+
input_schema={
26+
"type": "object",
27+
"properties": {"query": {"type": "string"}},
28+
"required": ["query"],
29+
},
30+
)
31+
]
32+
)
33+
34+
35+
async def on_call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult:
36+
query = (params.arguments or {})["query"]
37+
return CallToolResult(content=[TextContent(type="text", text=f"Found 3 books matching {query!r}.")])
38+
39+
40+
def max_concurrent_tool_calls(limit: int) -> ServerMiddleware[Any]:
41+
running = 0
42+
43+
async def middleware(ctx: ServerRequestContext, call_next: CallNext) -> HandlerResult:
44+
nonlocal running
45+
if ctx.method != "tools/call":
46+
return await call_next(ctx)
47+
if running >= limit:
48+
raise MCPError(code=SERVER_BUSY, message=f"Server busy: tool call limit reached ({limit} in progress).")
49+
running += 1
50+
try:
51+
return await call_next(ctx)
52+
finally:
53+
running -= 1
54+
55+
return middleware
56+
57+
58+
server = Server("Bookshop", on_list_tools=on_list_tools, on_call_tool=on_call_tool)
59+
server.middleware.append(max_concurrent_tool_calls(4))

‎tests/docs_src/test_middleware.py‎

Lines changed: 181 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,17 +3,20 @@
33
import logging
44
import re
55

6+
import anyio
67
import pytest
78
from mcp_types import (
9+
INTERNAL_ERROR,
810
INVALID_REQUEST,
911
METHOD_NOT_FOUND,
1012
CallToolRequestParams,
13+
CallToolResult,
1114
ErrorData,
1215
RequestId,
1316
TextContent,
1417
)
1518

16-
from docs_src.middleware import tutorial001
19+
from docs_src.middleware import tutorial001, tutorial002
1720
from mcp import Client, MCPError
1821
from mcp.server import Server, ServerRequestContext
1922
from mcp.server.context import CallNext, HandlerResult
@@ -27,6 +30,29 @@ def _is_timing_record(record: logging.LogRecord) -> bool:
2730
return record.name == tutorial001.logger.name
2831

2932

33+
class _HeldSearches:
34+
"""An `on_call_tool` for `search_books` that holds each call open until the test finishes it, keyed by query.
35+
36+
A query the test did not name is a `KeyError`: a call the cap should have refused cannot quietly succeed.
37+
"""
38+
39+
def __init__(self, *queries: str) -> None:
40+
self.started = {query: anyio.Event() for query in queries}
41+
self.finish = {query: anyio.Event() for query in queries}
42+
43+
async def __call__(self, ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult:
44+
assert params.name == "search_books"
45+
query = (params.arguments or {})["query"]
46+
self.started[query].set()
47+
await self.finish[query].wait()
48+
return CallToolResult(content=[TextContent(type="text", text=f"Found {query}.")])
49+
50+
51+
async def _search(client: Client, query: str, results: dict[str, CallToolResult]) -> None:
52+
"""Call `search_books` and file the result under its query, for calls a test runs in the background."""
53+
results[query] = await client.call_tool("search_books", {"query": query})
54+
55+
3056
def test_timing_record_predicate() -> None:
3157
"""The caplog filter keeps the middleware's own records and drops everyone else's."""
3258
args = (logging.INFO, __file__, 1, "msg", None, None)
@@ -114,3 +140,157 @@ async def test_initialize_cannot_be_replaced_only_wrapped() -> None:
114140
)
115141
with pytest.raises(ValueError, match=re.escape(expected)):
116142
tutorial001.server.add_request_handler("initialize", CallToolRequestParams, tutorial001.on_call_tool)
143+
144+
145+
async def test_a_tool_call_over_the_cap_is_refused_while_the_earlier_calls_still_run() -> None:
146+
"""tutorial002: with `limit` tool calls in flight, the next one is answered with the busy error at once.
147+
148+
Steps:
149+
1. One client's four calls enter the handler and are held there.
150+
2. A second client connects (its `server/discover` is not a tool call) and makes a fifth call, which is
151+
refused before any of the four has returned: the count belongs to the server, not to a connection.
152+
3. The four are finished and each returns its own result.
153+
"""
154+
held = _HeldSearches("dune", "emma", "ulysses", "walden")
155+
server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=held)
156+
server.middleware.append(tutorial002.max_concurrent_tool_calls(4))
157+
results: dict[str, CallToolResult] = {}
158+
with anyio.fail_after(5):
159+
async with Client(server) as client:
160+
async with anyio.create_task_group() as tg:
161+
for query in held.started:
162+
tg.start_soon(_search, client, query, results)
163+
for started in held.started.values():
164+
await started.wait()
165+
async with Client(server) as latecomer:
166+
with pytest.raises(MCPError) as exc_info:
167+
await latecomer.call_tool("search_books", {"query": "middlemarch"})
168+
assert exc_info.value.error == ErrorData(
169+
code=tutorial002.SERVER_BUSY, message="Server busy: tool call limit reached (4 in progress)."
170+
)
171+
assert results == {}
172+
for finish in held.finish.values():
173+
finish.set()
174+
assert {query: result.content for query, result in results.items()} == {
175+
query: [TextContent(type="text", text=f"Found {query}.")] for query in held.started
176+
}
177+
178+
179+
async def test_other_requests_are_answered_with_the_cap_reached() -> None:
180+
"""tutorial002: only `tools/call` is capped, so `tools/list` is answered with the cap reached."""
181+
held = _HeldSearches("dune")
182+
server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=held)
183+
server.middleware.append(tutorial002.max_concurrent_tool_calls(1))
184+
results: dict[str, CallToolResult] = {}
185+
with anyio.fail_after(5):
186+
async with Client(server) as client:
187+
async with anyio.create_task_group() as tg:
188+
tg.start_soon(_search, client, "dune", results)
189+
await held.started["dune"].wait()
190+
tools = (await client.list_tools()).tools
191+
assert [tool.name for tool in tools] == ["search_books"]
192+
assert results == {}
193+
held.finish["dune"].set()
194+
assert results["dune"].content == [TextContent(type="text", text="Found dune.")]
195+
196+
197+
async def test_a_refused_call_is_accepted_once_a_running_call_finishes() -> None:
198+
"""tutorial002: a finished call gives its slot back, so the caller that was refused gets through on a retry."""
199+
held = _HeldSearches("dune", "emma")
200+
server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=held)
201+
server.middleware.append(tutorial002.max_concurrent_tool_calls(1))
202+
results: dict[str, CallToolResult] = {}
203+
# Only `dune` is held open; `emma` returns as soon as it is let in.
204+
held.finish["emma"].set()
205+
with anyio.fail_after(5):
206+
async with Client(server) as client:
207+
async with anyio.create_task_group() as tg:
208+
tg.start_soon(_search, client, "dune", results)
209+
await held.started["dune"].wait()
210+
with pytest.raises(MCPError) as exc_info:
211+
await client.call_tool("search_books", {"query": "emma"})
212+
assert exc_info.value.error.code == tutorial002.SERVER_BUSY
213+
assert not held.started["emma"].is_set()
214+
held.finish["dune"].set()
215+
# The task group has joined, so the first call's response has arrived.
216+
retried = await client.call_tool("search_books", {"query": "emma"})
217+
assert results["dune"].content == [TextContent(type="text", text="Found dune.")]
218+
assert retried.content == [TextContent(type="text", text="Found emma.")]
219+
220+
221+
async def test_a_tool_call_that_raises_gives_its_slot_back() -> None:
222+
"""tutorial002: the `finally` releases the slot when the handler raises, so the next call is not refused."""
223+
224+
async def on_call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult:
225+
assert params.name == "search_books"
226+
query = (params.arguments or {})["query"]
227+
if query == "necronomicon":
228+
raise RuntimeError("the shelf collapsed")
229+
return CallToolResult(content=[TextContent(type="text", text=f"Found {query}.")])
230+
231+
server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=on_call_tool)
232+
server.middleware.append(tutorial002.max_concurrent_tool_calls(1))
233+
async with Client(server) as client:
234+
with pytest.raises(MCPError) as exc_info:
235+
await client.call_tool("search_books", {"query": "necronomicon"})
236+
assert exc_info.value.error.code == INTERNAL_ERROR
237+
result = await client.call_tool("search_books", {"query": "dune"})
238+
assert result.content == [TextContent(type="text", text="Found dune.")]
239+
240+
241+
async def test_a_cancelled_tool_call_gives_its_slot_back() -> None:
242+
"""tutorial002: a call the client abandons is cancelled out of `call_next`, and the `finally` frees its slot."""
243+
started = anyio.Event()
244+
cancelled = anyio.Event()
245+
246+
async def on_call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult:
247+
assert params.name == "search_books"
248+
query = (params.arguments or {})["query"]
249+
if query == "dune":
250+
started.set()
251+
try:
252+
await anyio.sleep_forever()
253+
finally:
254+
cancelled.set()
255+
return CallToolResult(content=[TextContent(type="text", text=f"Found {query}.")])
256+
257+
server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=on_call_tool)
258+
server.middleware.append(tutorial002.max_concurrent_tool_calls(1))
259+
results: dict[str, CallToolResult] = {}
260+
with anyio.fail_after(5):
261+
async with Client(server) as client:
262+
async with anyio.create_task_group() as tg:
263+
tg.start_soon(_search, client, "dune", results)
264+
await started.wait()
265+
tg.cancel_scope.cancel()
266+
await cancelled.wait()
267+
result = await client.call_tool("search_books", {"query": "emma"})
268+
assert results == {}
269+
assert result.content == [TextContent(type="text", text="Found emma.")]
270+
271+
272+
async def test_answering_with_an_error_result_returns_it_to_the_caller_instead_of_raising() -> None:
273+
"""A middleware that answers `tools/call` with an `is_error=True` result gives the client a result, not an error."""
274+
275+
async def busy(ctx: ServerRequestContext, call_next: CallNext) -> HandlerResult:
276+
if ctx.method == "tools/call":
277+
return CallToolResult(
278+
content=[TextContent(type="text", text="Server busy. Try again shortly.")], is_error=True
279+
)
280+
return await call_next(ctx)
281+
282+
server = Server("Bookshop", on_list_tools=tutorial002.on_list_tools, on_call_tool=tutorial002.on_call_tool)
283+
server.middleware.append(busy)
284+
async with Client(server) as client:
285+
result = await client.call_tool("search_books", {"query": "dune"})
286+
assert result.is_error
287+
assert result.content == [TextContent(type="text", text="Server busy. Try again shortly.")]
288+
289+
290+
async def test_calls_made_one_after_another_never_reach_the_cap() -> None:
291+
"""tutorial002's own server: the cap counts calls in flight, so more than four in sequence all succeed."""
292+
async with Client(tutorial002.server) as client:
293+
results = [await client.call_tool("search_books", {"query": "dune"}) for _ in range(5)]
294+
assert [result.content for result in results] == [
295+
[TextContent(type="text", text="Found 3 books matching 'dune'.")]
296+
] * 5

0 commit comments

Comments
 (0)