Skip to content
11 changes: 10 additions & 1 deletion src/mcp/client/_probe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand Down
38 changes: 35 additions & 3 deletions src/mcp/client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -366,6 +366,13 @@ async def main():
derived)."""

_entered: bool = field(init=False, default=False)
protocol_version_override: str | None = None
Comment thread
cubic-dev-ai[bot] marked this conversation as resolved.
"""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)
Expand All @@ -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 = (
Comment thread
cubic-dev-ai[bot] marked this conversation as resolved.
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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Comment thread
cubic-dev-ai[bot] marked this conversation as resolved.
else:
session.adopt(self.prior_discover or _synthesize_discover(self.mode))

Expand Down
7 changes: 3 additions & 4 deletions src/mcp/client/session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2: Older protocol overrides still advertise form and URL elicitation capabilities when the callback is configured. Gate each capability by the requested protocol version; otherwise the server can use an advertisement for features unavailable in the selected protocol.

Prompt for AI agents
Check if this issue is valid — if so, understand the root cause and fix it. At src/mcp/client/session.py, line 658:

<comment>Older protocol overrides still advertise form and URL elicitation capabilities when the callback is configured. Gate each capability by the requested protocol version; otherwise the server can use an advertisement for features unavailable in the selected protocol.</comment>

<file context>
@@ -648,15 +648,14 @@ def _build_capabilities(self, version: str) -> types.ClientCapabilities:
-                    # 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,
                 ),
</file context>

client_info=self._client_info,
),
),
Expand Down
6 changes: 5 additions & 1 deletion src/mcp/client/session_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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.
Expand Down
43 changes: 42 additions & 1 deletion tests/client/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -130,6 +130,47 @@ 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_unknown_protocol_version_override(app: MCPServer):
"""A `protocol_version_override` that is neither a handshake-era nor a modern version
is rejected with no extra hint, unlike the modern-version case above."""
with pytest.raises(ValueError, match=r"protocol_version_override must be one of .*got '1999-01-01'$"):
Client(app, mode="auto", protocol_version_override="1999-01-01")


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:
Expand Down
25 changes: 22 additions & 3 deletions tests/client/test_probe.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand All @@ -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]:
Expand Down Expand Up @@ -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
82 changes: 82 additions & 0 deletions tests/client/test_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
36 changes: 35 additions & 1 deletion tests/client/test_session_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Comment thread
cubic-dev-ai[bot] marked this conversation as resolved.
9 changes: 9 additions & 0 deletions tests/interaction/_requirements.py
Original file line number Diff line number Diff line change
Expand Up @@ -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=<version> 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
# ═══════════════════════════════════════════════════════════════════════════
Expand Down
Loading
Loading