diff --git a/tests/integration/_support/conformance.py b/tests/integration/_support/conformance.py index 82317850081..ef0ae6ff66a 100644 --- a/tests/integration/_support/conformance.py +++ b/tests/integration/_support/conformance.py @@ -1,4 +1,5 @@ import hashlib +import json import os import queue import signal @@ -115,7 +116,7 @@ def reference_server(root: Path, directory: Path, port: int) -> Iterator["McpPee @contextmanager -def authenticated_endpoint(target: str, key: str) -> Iterator[str]: +def authenticated_endpoint(target: str, key: str, alias: str) -> Iterator[str]: from integration._support.asgi import asgi_server from starlette.applications import Starlette from starlette.requests import Request @@ -125,10 +126,13 @@ def authenticated_endpoint(target: str, key: str) -> Iterator[str]: async def forward(scope: Scope, receive: Receive, send: Send) -> None: request: Final = Request(scope, receive) + original: Final = await request.body() + body: Final = prefixed_request(original, alias) headers: Final = tuple( - (name, value) for name, value in request.headers.raw if name.lower() != b"authorization" + (name, value) + for name, value in request.headers.raw + if name.lower() != b"authorization" and (name.lower() != b"content-length" or body == original) ) + ((b"authorization", f"Bearer {key}".encode()),) - body: Final = await request.body() async with httpx.AsyncClient(timeout=35, trust_env=False) as client: async with client.stream(request.method, target, headers=headers, content=body) as response: streamed: Final = StreamingResponse(response.aiter_raw(), status_code=response.status_code) @@ -141,3 +145,16 @@ def authenticated_endpoint(target: str, key: str) -> Iterator[str]: with asgi_server(Starlette(routes=[Mount("/mcp", app=forward)])) as url: yield url + "/mcp/" + + +def prefixed_request(body: bytes, alias: str) -> bytes: + try: + request: Final = json.loads(body) + except (ValueError, UnicodeDecodeError): + return body + if not isinstance(request, dict) or request.get("method") not in ("tools/call", "prompts/get"): + return body + params: Final = request.get("params") + if not isinstance(params, dict) or not isinstance(params.get("name"), str): + return body + return json.dumps({**request, "params": {**params, "name": f"{alias}-{params['name']}"}}).encode() diff --git a/tests/integration/mcp/test_mcp_official_conformance.py b/tests/integration/mcp/test_mcp_official_conformance.py index 0877fc9457e..c0375b4223e 100644 --- a/tests/integration/mcp/test_mcp_official_conformance.py +++ b/tests/integration/mcp/test_mcp_official_conformance.py @@ -19,7 +19,7 @@ def test_official_scenario_through_gateway(gateway: Gateway, tmp_path: Path, unu key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) direct: Final = run_scenario(root, reference.url, name, output / "direct") endpoint: Final = str(gateway.client.base_url).rstrip("/") + f"/{alias}/mcp" - with authenticated_endpoint(endpoint, key) as authenticated: + with authenticated_endpoint(endpoint, key, alias) as authenticated: proxied: Final = run_scenario(root, authenticated, name, output / "gateway") assert tuple(check.status for check in direct if check.id == name) == ("SUCCESS",) assert tuple(check.status for check in proxied if check.id == name) == ("SUCCESS",) diff --git a/tests/unit/integration_support/test_conformance.py b/tests/unit/integration_support/test_conformance.py index c7b0795479f..23d950938db 100644 --- a/tests/unit/integration_support/test_conformance.py +++ b/tests/unit/integration_support/test_conformance.py @@ -4,7 +4,7 @@ from pathlib import Path from typing import Final import pytest -from tests.integration._support.conformance import read_checks, verify_archive +from tests.integration._support.conformance import read_checks, verify_archive, prefixed_request from pydantic import ValidationError @@ -97,3 +97,29 @@ def test_official_text_success_requires_the_reference_payload(tmp_path: Path) -> ) ) assert read_checks(tmp_path, "tools-call-simple-text")[0].status == "SUCCESS" + + +@pytest.mark.parametrize("method", ("tools/call", "prompts/get")) +def test_only_fixture_name_is_translated_to_advertised_name(method: str) -> None: + request: Final = { + "jsonrpc": "2.0", + "id": 4, + "method": method, + "params": {"name": "test_tool", "_meta": {"progressToken": "keep"}}, + } + changed: Final = json.loads(prefixed_request(json.dumps(request).encode(), "official")) + assert changed == {**request, "params": {"name": "official-test_tool", "_meta": {"progressToken": "keep"}}} + assert "arguments" not in changed["params"] + + +@pytest.mark.parametrize( + "body", + ( + b"malformed JSON", + b"[]", + b'{"method":"initialize","params":{"protocolVersion":"2025-03-26"}}', + b'{"method":"tools/call","params":{"name":17}}', + ), +) +def test_name_integration_preserves_other_requests(body: bytes) -> None: + assert prefixed_request(body, "official") == body