diff --git a/tests/integration/_support/mcp.py b/tests/integration/_support/mcp.py index 6f58a2d5e29..73eb73ab214 100644 --- a/tests/integration/_support/mcp.py +++ b/tests/integration/_support/mcp.py @@ -77,7 +77,7 @@ def math_service( if ctx.method == "initialize": requested: Final = InitializeRequestParams.model_validate(ctx.params).protocol_version returned: Final = InitializeResult.model_validate(result).protocol_version - record({"body": {}, "negotiation": {"requested": requested, "returned": returned}}) + record({"body": {}, "headers": {}, "negotiation": {"requested": requested, "returned": returned}}) return result service.middleware.append(capture_negotiation) diff --git a/tests/unit/integration_support/test_conformance.py b/tests/unit/integration_support/test_conformance.py index e1cf0610411..db102b9cfea 100644 --- a/tests/unit/integration_support/test_conformance.py +++ b/tests/unit/integration_support/test_conformance.py @@ -18,6 +18,38 @@ from tests.integration._support.conformance import ( from pydantic import ValidationError +@pytest.mark.asyncio +async def test_negotiation_records_preserve_the_peer_call_contract() -> None: + from integration._support.mcp import math_service + from integration.mcp.test_mcp_user_env_vars import UPSTREAM_CALLS + from mcp.server.context import ServerRequestContext + from mcp.types import InitializeResult + + records: Final[list[dict[str, object]]] = [] + service: Final = math_service(record=records.append) + context: Final = ServerRequestContext( + session=Mock(), + lifespan_context=None, + protocol_version="2025-03-26", + method="initialize", + params={"protocolVersion": "2025-03-26", "capabilities": {}, "clientInfo": {"name": "client", "version": "1"}}, + ) + + async def initialize(request: ServerRequestContext[object, object]) -> InitializeResult: + return InitializeResult.model_validate( + { + "protocolVersion": request.protocol_version, + "capabilities": {}, + "serverInfo": {"name": "peer", "version": "1"}, + } + ) + + await service.middleware[-1](context, initialize) + calls: Final = UPSTREAM_CALLS.validate_python(tuple(records)) + assert len(calls) == 1 and calls[0].headers == {} and calls[0].body == {} + assert records[0]["negotiation"] == {"requested": "2025-03-26", "returned": "2025-03-26"} + + @pytest.mark.parametrize("explicit", (False, True)) def test_owned_proxy_isolates_automatic_coverage_unless_requested( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, explicit: bool