diff --git a/src/mcp/client/_probe.py b/src/mcp/client/_probe.py index 0e46ae57d9..2e0f416924 100644 --- a/src/mcp/client/_probe.py +++ b/src/mcp/client/_probe.py @@ -49,7 +49,7 @@ def _parse_supported(data: Any) -> list[str] | None: return None -async def negotiate_auto(session: ClientSession) -> None: +async def negotiate_auto(session: ClientSession, protocol_version: str | None = None) -> None: """Drive the ``mode='auto'`` connect-time policy on ``session``. Probes ``server/discover`` once (twice if the server names a mutual @@ -58,12 +58,21 @@ async def negotiate_auto(session: ClientSession) -> None: ``session.discover_result`` / ``session.initialize_result`` is set on return. + ``protocol_version`` pins the legacy handshake to a specific version. A + caller supplying it wants that exact version, so this skips the + ``server/discover`` probe entirely and goes straight to the handshake — + otherwise a server with modern support would win discovery and the pin + would be silently ignored. + Raises: MCPError: The server is modern-only and shares no version with this client (-32022 with a disjoint ``supported`` list), or the fallback handshake failed and one corrective re-probe did too. Exception: Any transport/network error from the probe propagates as-is. """ + if protocol_version is not None: + await session.initialize(protocol_version=protocol_version) + return version = LATEST_MODERN_VERSION for attempt in range(2): try: diff --git a/src/mcp/client/client.py b/src/mcp/client/client.py index f921c7e30b..a32f19f167 100644 --- a/src/mcp/client/client.py +++ b/src/mcp/client/client.py @@ -366,6 +366,13 @@ async def main(): derived).""" _entered: bool = field(init=False, default=False) + protocol_version_override: str | None = None + """Pin the legacy `initialize` handshake to a specific handshake-era version. + + Only meaningful with `mode='legacy'` or `mode='auto'` (where it skips `server/discover` + and negotiates directly); raises at construction with any other `mode`, since a version + pin already fixes the negotiated version. Must be a member of `HANDSHAKE_PROTOCOL_VERSIONS`. + `None` (the default) negotiates the latest version each `mode` would otherwise pick.""" _session: ClientSession | None = field(init=False, default=None) _exit_stack: AsyncExitStack | None = field(init=False, default=None) _connect: _Connector = field(init=False, repr=False, compare=False) @@ -383,6 +390,23 @@ def __post_init__(self) -> None: f"mode must be 'legacy', 'auto', or one of {list(MODERN_PROTOCOL_VERSIONS)}; got {self.mode!r}{hint}" ) + if self.protocol_version_override is not None: + if self.protocol_version_override not in HANDSHAKE_PROTOCOL_VERSIONS: + hint = ( + f" ({self.protocol_version_override!r} is a modern version; mode='auto' already negotiates it)" + if self.protocol_version_override in MODERN_PROTOCOL_VERSIONS + else "" + ) + raise ValueError( + "protocol_version_override must be one of " + f"{list(HANDSHAKE_PROTOCOL_VERSIONS)}; got {self.protocol_version_override!r}{hint}" + ) + if self.mode not in ("legacy", "auto"): + raise ValueError( + f"protocol_version_override has no effect with mode={self.mode!r} " + "(a version pin already fixes the negotiated version); use mode='legacy' or mode='auto'" + ) + self._folded_extensions = _fold_extensions(self.extensions) srv = self.server @@ -424,7 +448,12 @@ def __post_init__(self) -> None: async def _build_session(self, exit_stack: AsyncExitStack) -> ClientSession: """Enter the resolved connector and return an un-entered ClientSession.""" - dispatcher = await self._connect(exit_stack, self.mode, self.raise_exceptions) + # An override on mode='auto' skips discovery and drives `initialize()` directly + # (see `negotiate_auto`), so the in-proc connector must hand back the legacy, + # stream-backed dispatcher for this combination too, not the handshake-less + # DirectDispatcher it otherwise picks for every non-'legacy' mode. + connect_mode = "legacy" if self.mode == "auto" and self.protocol_version_override is not None else self.mode + dispatcher = await self._connect(exit_stack, connect_mode, self.raise_exceptions) message_handler = self.message_handler if self._response_cache is not None: message_handler = _evicting_message_handler(self._response_cache, self.message_handler) @@ -455,9 +484,12 @@ async def __aenter__(self) -> Client: session = await exit_stack.enter_async_context(session) if self.mode == "legacy": - await session.initialize() + if self.protocol_version_override is not None: + await session.initialize(protocol_version=self.protocol_version_override) + else: + await session.initialize() elif self.mode == "auto": - await negotiate_auto(session) + await negotiate_auto(session, protocol_version=self.protocol_version_override) else: session.adopt(self.prior_discover or _synthesize_discover(self.mode)) diff --git a/src/mcp/client/session.py b/src/mcp/client/session.py index a618112153..c3f22c9670 100644 --- a/src/mcp/client/session.py +++ b/src/mcp/client/session.py @@ -648,15 +648,14 @@ def _build_capabilities(self, version: str) -> types.ClientCapabilities: sampling=sampling, elicitation=elicitation, experimental=None, extensions=extensions, roots=roots ) - async def initialize(self) -> types.InitializeResult: + async def initialize(self, protocol_version: str = LATEST_HANDSHAKE_VERSION) -> types.InitializeResult: if self._initialize_result is not None: return self._initialize_result result = await self.send_request( types.InitializeRequest( params=types.InitializeRequestParams( - protocol_version=LATEST_HANDSHAKE_VERSION, - # The handshake negotiates only legacy versions, where no claim is active. - capabilities=self._build_capabilities(LATEST_HANDSHAKE_VERSION), + protocol_version=protocol_version, + capabilities=self._build_capabilities(protocol_version), client_info=self._client_info, ), ), diff --git a/src/mcp/client/session_group.py b/src/mcp/client/session_group.py index a544cecbe8..9d7e725b7e 100644 --- a/src/mcp/client/session_group.py +++ b/src/mcp/client/session_group.py @@ -80,6 +80,7 @@ class ClientSessionParameters: logging_callback: LoggingFnT | None = None message_handler: MessageHandlerFnT | None = None client_info: types.Implementation | None = None + protocol_version: str | None = None class ClientSessionGroup: @@ -352,7 +353,10 @@ async def _establish_session( ) ) - result = await session.initialize() + if session_params.protocol_version is not None: + result = await session.initialize(protocol_version=session_params.protocol_version) + else: + result = await session.initialize() # Session successfully initialized. # Store its stack and register the stack with the main group stack. diff --git a/tests/client/test_client.py b/tests/client/test_client.py index d7278e3a81..adc109552a 100644 --- a/tests/client/test_client.py +++ b/tests/client/test_client.py @@ -33,7 +33,7 @@ Tool, ToolsCapability, ) -from mcp_types.version import LATEST_HANDSHAKE_VERSION +from mcp_types.version import LATEST_HANDSHAKE_VERSION, LATEST_MODERN_VERSION from pydantic import FileUrl from mcp import MCPDeprecationWarning, MCPError, StdioServerParameters @@ -130,6 +130,40 @@ async def test_client_exposes_negotiated_protocol_version(app: MCPServer): assert client.protocol_version == LATEST_HANDSHAKE_VERSION +async def test_client_custom_protocol_version(app: MCPServer): + """Test that the client negotiates a custom protocol version when configured.""" + async with Client(app, mode="legacy", protocol_version_override="2024-11-05") as client: + assert client.protocol_version == "2024-11-05" + assert client.server_info is not None + assert client.server_info.name == "test" + + +async def test_client_auto_mode_with_override_against_in_process_server(app: MCPServer): + """Regression: `mode='auto'` with `protocol_version_override` against an in-process + `Server`/`MCPServer` used to always get the handshake-less `DirectDispatcher` (every + non-'legacy' mode picked it), so `negotiate_auto`'s direct `initialize()` call for the + override case had no JSON-RPC dispatcher to run on and the connect failed. + """ + async with Client(app, mode="auto", protocol_version_override="2024-11-05") as client: + assert client.protocol_version == "2024-11-05" + assert client.server_info is not None + assert client.server_info.name == "test" + + +def test_client_rejects_modern_protocol_version_override(app: MCPServer): + """`protocol_version_override` only pins the legacy handshake; a modern version string + is a construction-time error rather than a confusing failure once connected.""" + with pytest.raises(ValueError, match="protocol_version_override must be one of"): + Client(app, mode="auto", protocol_version_override=LATEST_MODERN_VERSION) + + +def test_client_rejects_protocol_version_override_with_a_version_pin_mode(app: MCPServer): + """`protocol_version_override` has no effect once `mode` already pins a version, so it's + rejected at construction instead of being silently ignored.""" + with pytest.raises(ValueError, match="protocol_version_override has no effect with mode="): + Client(app, mode=LATEST_MODERN_VERSION, protocol_version_override="2024-11-05") + + async def test_client_with_simple_server(simple_server: Server): """Test that from_server works with a basic Server instance.""" async with Client(simple_server) as client: diff --git a/tests/client/test_probe.py b/tests/client/test_probe.py index 354f8fd0c1..0cb724326e 100644 --- a/tests/client/test_probe.py +++ b/tests/client/test_probe.py @@ -57,6 +57,7 @@ def __init__(self, *script: dict[str, Any] | Exception, handshake: list[Exceptio self.probed_at: list[str] = [] self.initialize_calls: int = 0 self.initialized: bool = False + self.initialize_version: str | None = None self.adopted: types.DiscoverResult | None = None async def send_discover(self, version: str) -> dict[str, Any]: @@ -66,19 +67,20 @@ async def send_discover(self, version: str) -> dict[str, Any]: raise step return step - async def initialize(self) -> None: + async def initialize(self, protocol_version: str | None = None) -> None: self.initialize_calls += 1 if self._handshake: raise self._handshake.pop(0) self.initialized = True + self.initialize_version = protocol_version def adopt(self, result: types.DiscoverResult) -> None: self.adopted = result -async def _negotiate(session: _StubSession) -> None: +async def _negotiate(session: _StubSession, protocol_version: str | None = None) -> None: """Drive `negotiate_auto` against the stub; cast at one seam so the tests stay suppression-free.""" - await negotiate_auto(cast("ClientSession", session)) + await negotiate_auto(cast("ClientSession", session), protocol_version=protocol_version) def _discover_dict(versions: list[str] | None = None) -> dict[str, Any]: @@ -331,3 +333,20 @@ def test_parse_supported_returns_none_for_anything_not_shaped_like_the_spec_erro """`_parse_supported` returns the `supported` list when `error.data` validates as `UnsupportedProtocolVersionErrorData`, and `None` otherwise — never raises.""" assert _parse_supported(data) == expected + + +# --- protocol_version override forces the legacy handshake, unconditionally --- + + +async def test_a_protocol_version_override_skips_discovery_and_forces_the_legacy_handshake() -> None: + """`protocol_version` pins an explicit legacy version, so the caller wants exactly that + version - the probe is skipped entirely and the handshake runs unconditionally, even + though the stub's discover script would otherwise return a valid modern result (regression: + the override used to only reach `initialize()` via the fallback paths, so a successful + discover silently dropped it).""" + session = _StubSession(_discover_dict()) + await _negotiate(session, protocol_version="2024-11-05") + assert session.probed_at == [] + assert session.initialized + assert session.initialize_version == "2024-11-05" + assert session.adopted is None diff --git a/tests/client/test_session.py b/tests/client/test_session.py index 6663fb47a2..a07f65fde4 100644 --- a/tests/client/test_session.py +++ b/tests/client/test_session.py @@ -154,6 +154,88 @@ async def message_handler(message: IncomingMessage) -> None: # pragma: no cover assert isinstance(initialized_notification, InitializedNotification) +@pytest.mark.anyio +async def test_client_session_initialize_custom_protocol_version(): + client_to_server_send, client_to_server_receive = anyio.create_memory_object_stream[SessionMessage](1) + server_to_client_send, server_to_client_receive = anyio.create_memory_object_stream[SessionMessage](1) + + initialized_notification = None + result = None + + async def mock_server(): + nonlocal initialized_notification + + session_message = await client_to_server_receive.receive() + jsonrpc_request = session_message.message + assert isinstance(jsonrpc_request, JSONRPCRequest) + request = client_request_adapter.validate_python( + jsonrpc_request.model_dump(by_alias=True, mode="json", exclude_none=True) + ) + assert isinstance(request, InitializeRequest) + assert request.params.protocol_version == "2024-11-05" + + result = InitializeResult( + protocol_version="2024-11-05", + capabilities=ServerCapabilities( + logging=None, + resources=None, + tools=None, + experimental=None, + prompts=None, + ), + server_info=Implementation(name="mock-server", version="0.1.0"), + instructions="The server instructions.", + ) + + async with server_to_client_send: + await server_to_client_send.send( + SessionMessage( + JSONRPCResponse( + jsonrpc="2.0", + id=jsonrpc_request.id, + result=result.model_dump(by_alias=True, mode="json", exclude_none=True), + ) + ) + ) + session_notification = await client_to_server_receive.receive() + jsonrpc_notification = session_notification.message + assert isinstance(jsonrpc_notification, JSONRPCNotification) + initialized_notification = client_notification_adapter.validate_python( + jsonrpc_notification.model_dump(by_alias=True, mode="json", exclude_none=True) + ) + + # Create a message handler to catch exceptions + async def message_handler(message: IncomingMessage) -> None: # pragma: no cover + if isinstance(message, Exception): + raise message + + async with ( + ClientSession( + server_to_client_receive, + client_to_server_send, + message_handler=message_handler, + ) as session, + anyio.create_task_group() as tg, + client_to_server_send, + client_to_server_receive, + server_to_client_send, + server_to_client_receive, + ): + tg.start_soon(mock_server) + result = await session.initialize(protocol_version="2024-11-05") + + # Assert the result + assert isinstance(result, InitializeResult) + assert result.protocol_version == "2024-11-05" + assert isinstance(result.capabilities, ServerCapabilities) + assert result.server_info == Implementation(name="mock-server", version="0.1.0") + assert result.instructions == "The server instructions." + + # Check that the client sent the initialized notification + assert initialized_notification + assert isinstance(initialized_notification, InitializedNotification) + + @pytest.mark.anyio async def test_client_session_custom_client_info(): client_to_server_send, client_to_server_receive = anyio.create_memory_object_stream[SessionMessage](1) diff --git a/tests/client/test_session_group.py b/tests/client/test_session_group.py index b75d22b7a0..45a6db5a64 100644 --- a/tests/client/test_session_group.py +++ b/tests/client/test_session_group.py @@ -397,8 +397,42 @@ async def test_client_session_group_establish_session_parameterized( client_info=None, ) mock_raw_session_cm.__aenter__.assert_awaited_once() - mock_entered_session.initialize.assert_awaited_once() + mock_entered_session.initialize.assert_awaited_once_with() # 3. Assert returned values assert returned_server_info is mock_initialize_result.server_info assert returned_session is mock_entered_session + + +@pytest.mark.anyio +async def test_client_session_group_establish_session_custom_protocol_version(): + with mock.patch("mcp.client.session_group.mcp.ClientSession") as mock_ClientSession_class: + with mock.patch("mcp.client.session_group.mcp.stdio_client") as mock_stdio_client: + mock_client_cm_instance = mock.AsyncMock(name="stdioClientCM") + mock_read_stream = mock.AsyncMock(name="stdioRead") + mock_write_stream = mock.AsyncMock(name="stdioWrite") + + mock_client_cm_instance.__aenter__.return_value = (mock_read_stream, mock_write_stream) + mock_client_cm_instance.__aexit__ = mock.AsyncMock(return_value=None) + mock_stdio_client.return_value = mock_client_cm_instance + + mock_raw_session_cm = mock.AsyncMock(name="RawSessionCM") + mock_ClientSession_class.return_value = mock_raw_session_cm + + mock_entered_session = mock.AsyncMock(name="EnteredSessionInstance") + mock_raw_session_cm.__aenter__.return_value = mock_entered_session + mock_raw_session_cm.__aexit__ = mock.AsyncMock(return_value=None) + + mock_initialize_result = mock.AsyncMock(name="InitializeResult") + mock_initialize_result.server_info = types.Implementation(name="foo", version="1") + mock_entered_session.initialize.return_value = mock_initialize_result + + group = ClientSessionGroup() + server_params = StdioServerParameters(command="test_stdio_cmd") + session_params = ClientSessionParameters(protocol_version="2024-11-05") + + async with contextlib.AsyncExitStack() as stack: + group._exit_stack = stack + await group._establish_session(server_params, session_params) + + mock_entered_session.initialize.assert_awaited_once_with(protocol_version="2024-11-05") diff --git a/tests/interaction/_requirements.py b/tests/interaction/_requirements.py index 235fb65cd4..5e525dbbbf 100644 --- a/tests/interaction/_requirements.py +++ b/tests/interaction/_requirements.py @@ -465,6 +465,15 @@ def __post_init__(self) -> None: ), added_in="2026-07-28", ), + "lifecycle:mode:auto-override-skips-discover": Requirement( + source="sdk", + behavior=( + "A Client constructed with mode='auto' and protocol_version_override= sends " + "initialize at that version as its first request and never sends server/discover, even " + "when the server would answer discover successfully." + ), + added_in="2026-07-28", + ), # ═══════════════════════════════════════════════════════════════════════════ # Protocol primitives: cancellation, timeout, progress, errors, _meta # ═══════════════════════════════════════════════════════════════════════════ diff --git a/tests/interaction/lowlevel/test_client_connect.py b/tests/interaction/lowlevel/test_client_connect.py index a992027a23..cb48744082 100644 --- a/tests/interaction/lowlevel/test_client_connect.py +++ b/tests/interaction/lowlevel/test_client_connect.py @@ -178,6 +178,35 @@ async def test_auto_mode_probes_server_discover_and_adopts_the_result() -> None: assert "initialize" not in [b["method"] for b in bodies] +@requirement("lifecycle:mode:auto-override-skips-discover") +async def test_auto_mode_with_a_protocol_version_override_skips_discover_and_initializes() -> None: + """`Client(..., mode='auto', protocol_version_override=...)` sends `initialize` at the + override version and never probes `server/discover`, even though the mounted server answers + discover successfully. Regression: the override used to only reach `negotiate_auto`'s + `initialize()` fallback calls, so a successful discover silently dropped it and the client + ended up modern-negotiated at the server's latest version instead of the pinned one. + """ + requests, on_request = _request_recorder() + server = _tools_server("discoverable") + + with anyio.fail_after(5): + async with ( + mounted_app(server, on_request=on_request) as (http, _), + Client( + streamable_http_client(f"{BASE_URL}/mcp", http_client=http), + mode="auto", + protocol_version_override="2024-11-05", + ) as client, + ): + assert client.protocol_version == "2024-11-05" + assert client.server_info is not None + assert client.server_info.name == "discoverable" + + bodies = [json.loads(r.content)["method"] for r in requests if r.method == "POST"] + assert bodies[0] == "initialize" + assert "server/discover" not in bodies + + @requirement("lifecycle:discover:retry-on-32022") async def test_auto_mode_retries_discover_once_on_unsupported_protocol_version() -> None: """A -32022 from `server/discover` triggers exactly one retry at the highest mutual modern version.