Skip to content

Commit bd1665b

Browse files
fix(session): return METHOD_NOT_FOUND (-32601) for unknown request methods
1 parent c0c5a9d commit bd1665b

3 files changed

Lines changed: 142 additions & 6 deletions

File tree

src/mcp/shared/session.py

Lines changed: 36 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
from contextlib import AsyncExitStack
44
from datetime import timedelta
55
from types import TracebackType
6-
from typing import Any, Generic, Protocol, TypeVar
6+
from typing import Any, Generic, Protocol, TypeVar, get_args
77

88
import anyio
99
import httpx
@@ -17,6 +17,7 @@
1717
from mcp.types import (
1818
CONNECTION_CLOSED,
1919
INVALID_PARAMS,
20+
METHOD_NOT_FOUND,
2021
CancelledNotification,
2122
ClientNotification,
2223
ClientRequest,
@@ -159,6 +160,31 @@ def cancelled(self) -> bool: # pragma: no cover
159160
return self._cancel_scope.cancel_called
160161

161162

163+
def _extract_known_request_methods(request_type: type[Any]) -> frozenset[str]:
164+
"""Extract default method names from a Pydantic RootModel or Union of request models."""
165+
try:
166+
union_type = getattr(request_type, "__value__", None) or request_type
167+
if hasattr(union_type, "model_fields") and "root" in union_type.model_fields:
168+
union_type = union_type.model_fields["root"].annotation
169+
170+
methods = set()
171+
172+
def _unpack_union(t: Any) -> None:
173+
args = get_args(t)
174+
if args:
175+
for arg in args:
176+
_unpack_union(arg)
177+
elif hasattr(t, "model_fields") and "method" in t.model_fields:
178+
m = t.model_fields["method"].default
179+
if isinstance(m, str):
180+
methods.add(m)
181+
182+
_unpack_union(union_type)
183+
return frozenset(methods)
184+
except Exception:
185+
return frozenset()
186+
187+
162188
class BaseSession(
163189
Generic[
164190
SendRequestT,
@@ -197,6 +223,7 @@ def __init__(
197223
self._request_id = 0
198224
self._receive_request_type = receive_request_type
199225
self._receive_notification_type = receive_notification_type
226+
self._known_request_methods = _extract_known_request_methods(receive_request_type)
200227
self._session_read_timeout_seconds = read_timeout_seconds
201228
self._in_flight = {}
202229
self._progress_callbacks = {}
@@ -348,6 +375,13 @@ async def _send_response(self, request_id: RequestId, response: SendResultT | Er
348375
session_message = SessionMessage(message=JSONRPCMessage(jsonrpc_response))
349376
await self._write_stream.send(session_message)
350377

378+
def _get_request_validation_error(self, request: JSONRPCRequest) -> ErrorData:
379+
if self._known_request_methods and (
380+
not isinstance(request.method, str) or request.method not in self._known_request_methods
381+
):
382+
return ErrorData(code=METHOD_NOT_FOUND, message="Method not found")
383+
return ErrorData(code=INVALID_PARAMS, message="Invalid request parameters")
384+
351385
async def _receive_loop(self) -> None:
352386
async with (
353387
self._read_stream,
@@ -385,11 +419,7 @@ async def _receive_loop(self) -> None:
385419
error_response = JSONRPCError(
386420
jsonrpc="2.0",
387421
id=message.message.root.id,
388-
error=ErrorData(
389-
code=INVALID_PARAMS,
390-
message="Invalid request parameters",
391-
data="",
392-
),
422+
error=self._get_request_validation_error(message.message.root),
393423
)
394424
session_message = SessionMessage(message=JSONRPCMessage(error_response))
395425
await self._write_stream.send(session_message)

tests/server/test_session.py

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -521,3 +521,75 @@ async def mock_client():
521521

522522
assert error_response_received
523523
assert error_code == types.INVALID_PARAMS
524+
525+
526+
@pytest.mark.anyio
527+
async def test_unknown_method_returns_method_not_found():
528+
"""Test that unknown request methods return METHOD_NOT_FOUND (-32601),
529+
while malformed known methods return INVALID_PARAMS (-32602).
530+
"""
531+
server_to_client_send, server_to_client_receive = anyio.create_memory_object_stream[SessionMessage](1)
532+
client_to_server_send, client_to_server_receive = anyio.create_memory_object_stream[SessionMessage | Exception](1)
533+
534+
async def run_server():
535+
async with ServerSession(
536+
client_to_server_receive,
537+
server_to_client_send,
538+
InitializationOptions(
539+
server_name="mcp",
540+
server_version="0.1.0",
541+
capabilities=ServerCapabilities(),
542+
),
543+
):
544+
await anyio.sleep(0.2)
545+
546+
async def mock_client():
547+
# 1. Send unknown method request
548+
await client_to_server_send.send(
549+
SessionMessage(
550+
types.JSONRPCMessage(
551+
types.JSONRPCRequest(
552+
jsonrpc="2.0",
553+
id=1,
554+
method="totally/bogus",
555+
params={},
556+
)
557+
)
558+
)
559+
)
560+
561+
resp1 = await server_to_client_receive.receive()
562+
assert isinstance(resp1.message.root, types.JSONRPCError)
563+
assert resp1.message.root.id == 1
564+
assert resp1.message.root.error.code == types.METHOD_NOT_FOUND
565+
assert resp1.message.root.error.message == "Method not found"
566+
567+
# 2. Send malformed known method request (initialize with invalid params shape)
568+
await client_to_server_send.send(
569+
SessionMessage(
570+
types.JSONRPCMessage(
571+
types.JSONRPCRequest(
572+
jsonrpc="2.0",
573+
id=2,
574+
method="initialize",
575+
params={"protocolVersion": 12345}, # invalid type for protocolVersion
576+
)
577+
)
578+
)
579+
)
580+
581+
resp2 = await server_to_client_receive.receive()
582+
assert isinstance(resp2.message.root, types.JSONRPCError)
583+
assert resp2.message.root.id == 2
584+
assert resp2.message.root.error.code == types.INVALID_PARAMS
585+
assert resp2.message.root.error.message == "Invalid request parameters"
586+
587+
async with (
588+
client_to_server_send,
589+
client_to_server_receive,
590+
server_to_client_send,
591+
server_to_client_receive,
592+
anyio.create_task_group() as tg,
593+
):
594+
tg.start_soon(run_server)
595+
tg.start_soon(mock_client)

tests/shared/test_session.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -339,3 +339,37 @@ async def mock_server():
339339
await ev_closed.wait()
340340
with anyio.fail_after(1): # pragma: no cover
341341
await ev_response.wait()
342+
343+
344+
@pytest.mark.anyio
345+
async def test_client_session_unknown_method_returns_method_not_found():
346+
"""Test that ClientSession returns METHOD_NOT_FOUND (-32601) when receiving an unknown server request method."""
347+
async with create_client_server_memory_streams() as (client_stream, server_stream):
348+
client_read, client_write = client_stream
349+
server_read, server_write = server_stream
350+
351+
async def mock_server():
352+
# Send unknown request method to client
353+
await server_write.send(
354+
SessionMessage(
355+
message=JSONRPCMessage(
356+
JSONRPCRequest(
357+
jsonrpc="2.0",
358+
id=1,
359+
method="invalid/server_method",
360+
params={},
361+
)
362+
)
363+
)
364+
)
365+
resp = await server_read.receive()
366+
assert isinstance(resp.message.root, JSONRPCError)
367+
assert resp.message.root.id == 1
368+
assert resp.message.root.error.code == types.METHOD_NOT_FOUND
369+
assert resp.message.root.error.message == "Method not found"
370+
371+
async with (
372+
ClientSession(read_stream=client_read, write_stream=client_write),
373+
anyio.create_task_group() as tg,
374+
):
375+
tg.start_soon(mock_server)

0 commit comments

Comments
 (0)