Skip to content

Commit 9936777

Browse files
fix(session): satisfy strict CI for method classification
1 parent bd1665b commit 9936777

2 files changed

Lines changed: 26 additions & 6 deletions

File tree

src/mcp/shared/session.py

Lines changed: 4 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, get_args
6+
from typing import Any, Generic, Protocol, TypeVar, cast, get_args
77

88
import anyio
99
import httpx
@@ -167,15 +167,15 @@ def _extract_known_request_methods(request_type: type[Any]) -> frozenset[str]:
167167
if hasattr(union_type, "model_fields") and "root" in union_type.model_fields:
168168
union_type = union_type.model_fields["root"].annotation
169169

170-
methods = set()
170+
methods: set[str] = set()
171171

172172
def _unpack_union(t: Any) -> None:
173173
args = get_args(t)
174174
if args:
175175
for arg in args:
176176
_unpack_union(arg)
177177
elif hasattr(t, "model_fields") and "method" in t.model_fields:
178-
m = t.model_fields["method"].default
178+
m = cast(Any, t.model_fields["method"].default)
179179
if isinstance(m, str):
180180
methods.add(m)
181181

@@ -376,9 +376,7 @@ async def _send_response(self, request_id: RequestId, response: SendResultT | Er
376376
await self._write_stream.send(session_message)
377377

378378
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-
):
379+
if self._known_request_methods and request.method not in self._known_request_methods:
382380
return ErrorData(code=METHOD_NOT_FOUND, message="Method not found")
383381
return ErrorData(code=INVALID_PARAMS, message="Invalid request parameters")
384382

tests/shared/test_session.py

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
from collections.abc import AsyncGenerator
2+
from types import SimpleNamespace
23
from typing import Any
34

45
import anyio
@@ -10,6 +11,7 @@
1011
from mcp.shared.exceptions import McpError
1112
from mcp.shared.memory import create_client_server_memory_streams, create_connected_server_and_client_session
1213
from mcp.shared.message import SessionMessage
14+
from mcp.shared.session import _extract_known_request_methods
1315
from mcp.types import (
1416
CancelledNotification,
1517
CancelledNotificationParams,
@@ -363,6 +365,7 @@ async def mock_server():
363365
)
364366
)
365367
resp = await server_read.receive()
368+
assert isinstance(resp, SessionMessage)
366369
assert isinstance(resp.message.root, JSONRPCError)
367370
assert resp.message.root.id == 1
368371
assert resp.message.root.error.code == types.METHOD_NOT_FOUND
@@ -373,3 +376,22 @@ async def mock_server():
373376
anyio.create_task_group() as tg,
374377
):
375378
tg.start_soon(mock_server)
379+
380+
381+
def test_extract_known_request_methods_ignores_non_string_defaults():
382+
class RequestWithNonStringMethod:
383+
model_fields = {"method": SimpleNamespace(default=None)}
384+
385+
assert _extract_known_request_methods(RequestWithNonStringMethod) == frozenset()
386+
387+
388+
def test_extract_known_request_methods_fails_closed_on_schema_introspection_error():
389+
class ExplodingMeta(type):
390+
@property
391+
def __value__(cls) -> type[Any]:
392+
raise RuntimeError("broken schema")
393+
394+
class BrokenRequest(metaclass=ExplodingMeta):
395+
pass
396+
397+
assert _extract_known_request_methods(BrokenRequest) == frozenset()

0 commit comments

Comments
 (0)