From 6079a5f55b27070eba0dfec812ecefee3f2bce54 Mon Sep 17 00:00:00 2001 From: gym-cmd <186399764+gym-cmd@users.noreply.github.com> Date: Fri, 15 May 2026 23:17:10 +0100 Subject: [PATCH] test(mcp): split oauth passthrough regressions --- .../mcp_server/test_mcp_oauth_passthrough.py | 309 +----------------- .../test_mcp_oauth_passthrough_cold_start.py | 153 +++++++++ .../test_mcp_oauth_passthrough_tools.py | 140 ++++++++ 3 files changed, 294 insertions(+), 308 deletions(-) create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py index ab237d58d2c..b79b38ae0ba 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py @@ -1,13 +1,10 @@ -"""Unit tests for the MCP OAuth pass-through patch (EAI-506 / Idea G). +"""Unit tests for MCP OAuth passthrough metadata behavior. Covers: - `MCPServer.is_oauth_passthrough` property semantics. - `/.well-known/oauth-protected-resource/...` pass-through branch (proxies upstream metadata, normalizes the `resource` field, caches, and surfaces network errors as HTTP 502). -- `MCPServerManager._fetch_tools_with_timeout` converting upstream 401s into - `MCPUpstreamAuthError` for pass-through servers while keeping the silent - empty-list fallback for gateway-managed / aggregator paths. """ import sys @@ -25,11 +22,6 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( _OAUTH_METADATA_CACHE, _build_oauth_protected_resource_response, ) -from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError -from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - MCPServerManager, - _extract_upstream_auth_failure, -) from litellm.proxy._types import MCPTransport from litellm.types.mcp import MCPAuth from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -388,302 +380,3 @@ async def test_oauth_protected_resource_gateway_managed_unchanged(): "https://gateway.example.com/keycloak_whoami" ] assert result["scopes_supported"] == ["read"] - - -# -------------------------------------------------------------------------- -# _extract_upstream_auth_failure helper -# -------------------------------------------------------------------------- - - -def test_extract_upstream_auth_failure_finds_401_in_http_status_error(): - response = httpx.Response( - status_code=401, - headers={"www-authenticate": 'Bearer resource_metadata="https://x"'}, - request=httpx.Request("GET", "https://upstream/mcp"), - ) - exc = httpx.HTTPStatusError("401", request=response.request, response=response) - - result = _extract_upstream_auth_failure(exc) - assert result == (401, 'Bearer resource_metadata="https://x"') - - -def test_extract_upstream_auth_failure_walks_exception_group(): - response = httpx.Response( - status_code=401, - headers={"www-authenticate": "Bearer"}, - request=httpx.Request("GET", "https://upstream/mcp"), - ) - inner = httpx.HTTPStatusError("401", request=response.request, response=response) - - try: - raise ExceptionGroup("wrapped", [inner]) # noqa: F821 (PEP 654, py3.11+) - except Exception as group: - result = _extract_upstream_auth_failure(group) - - assert result == (401, "Bearer") - - -def test_extract_upstream_auth_failure_returns_none_for_non_auth(): - assert _extract_upstream_auth_failure(RuntimeError("boom")) is None - - -# -------------------------------------------------------------------------- -# _fetch_tools_with_timeout pass-through behaviour -# -------------------------------------------------------------------------- - - -@pytest.mark.asyncio -async def test_fetch_tools_from_passthrough_raises_on_upstream_401(): - manager = MCPServerManager() - passthrough_server = MCPServer( - server_id="p1", - name="sample_docs", - url="https://upstream/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.none, - extra_headers=["Authorization"], - ) - - response = httpx.Response( - status_code=401, - headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'}, - request=httpx.Request("GET", "https://upstream/mcp"), - ) - upstream_error = httpx.HTTPStatusError( - "401", request=response.request, response=response - ) - - mock_client = MagicMock() - mock_client.list_tools = AsyncMock(side_effect=upstream_error) - - with pytest.raises(MCPUpstreamAuthError) as exc_info: - await manager._fetch_tools_with_timeout( - mock_client, passthrough_server.name, server=passthrough_server - ) - - assert exc_info.value.status_code == 401 - assert exc_info.value.www_authenticate == ( - 'Bearer resource_metadata="https://upstream"' - ) - assert exc_info.value.server_name == "sample_docs" - mock_client.list_tools.assert_awaited_with(raise_on_error=True) - - -@pytest.mark.asyncio -async def test_fetch_tools_from_passthrough_returns_tools_on_success(): - manager = MCPServerManager() - passthrough_server = MCPServer( - server_id="p1", - name="sample_docs", - url="https://upstream/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.none, - extra_headers=["Authorization"], - ) - - # list_tools returns a pre-baked tools list directly (MCPClient contract). - tool = MagicMock() - tool.name = "list_knowledge_bases" - mock_client = MagicMock() - mock_client.list_tools = AsyncMock(return_value=[tool]) - - tools = await manager._fetch_tools_with_timeout( - mock_client, passthrough_server.name, server=passthrough_server - ) - assert tools == [tool] - - -@pytest.mark.asyncio -async def test_fetch_tools_from_gateway_managed_swallows_errors(): - """Regression guard: non-pass-through servers keep returning [] on errors - so the multi-server aggregator isn't tainted by a single bad server.""" - manager = MCPServerManager() - oauth2_server = MCPServer( - server_id="o1", - name="keycloak_whoami", - url="https://upstream/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - ) - - response = httpx.Response( - status_code=401, - headers={}, - request=httpx.Request("GET", "https://upstream/mcp"), - ) - upstream_error = httpx.HTTPStatusError( - "401", request=response.request, response=response - ) - mock_client = MagicMock() - mock_client.list_tools = AsyncMock(side_effect=upstream_error) - - tools = await manager._fetch_tools_with_timeout( - mock_client, oauth2_server.name, server=oauth2_server - ) - assert tools == [] - mock_client.list_tools.assert_awaited_with(raise_on_error=False) - - -# -------------------------------------------------------------------------- -# §2.1 — Admission cold-start: 401 + matching resource_metadata URL -# -------------------------------------------------------------------------- - - -def _make_scope(path: str, headers: list = None) -> dict: - """Build a minimal ASGI HTTP scope for testing.""" - raw_headers = [(k.encode(), v.encode()) for k, v in (headers or [])] - return { - "type": "http", - "method": "POST", - "path": path, - "headers": raw_headers, - "query_string": b"", - "server": ("localhost", 4000), - "scheme": "http", - } - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "route,expected_metadata_path", - [ - ( - "/mcp/sample_docs", - "/.well-known/oauth-protected-resource/mcp/sample_docs", - ), - ( - "/sample_docs/mcp", - "/.well-known/oauth-protected-resource/sample_docs/mcp", - ), - ], -) -async def test_passthrough_cold_start_emits_401_with_matching_resource_metadata( - route, expected_metadata_path -): - """No auth headers on a pass-through server route → 401 with resource_metadata URL - that matches the inbound path so RFC 9728 §3.2 strict clients accept it.""" - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( - _is_mcp_passthrough_cold_start, - _parse_mcp_server_names_from_path, - ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - - global_mcp_server_manager.registry.clear() - passthrough_server = MCPServer( - server_id="pt-cold-start", - name="sample_docs", - server_name="sample_docs", - alias="sample_docs", - url="https://upstream.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.none, - extra_headers=["Authorization"], - ) - global_mcp_server_manager.registry[passthrough_server.server_id] = ( - passthrough_server - ) - - # For /mcp/{name}: scope path stays as-is. - # For /{name}/mcp: dynamic_mcp_route rewrites to /mcp/{name} and sets _original_path. - if route.startswith("/mcp/"): - scope = _make_scope(route) - else: - scope = _make_scope("/mcp/sample_docs") - scope["_original_path"] = route - - # Verify cold-start detection fires for this path - servers = _parse_mcp_server_names_from_path( - scope.get("path", "") # always /mcp/{name} by the time admission runs - ) - assert _is_mcp_passthrough_cold_start(scope, servers, client_ip=None) is True - - # Verify resource_metadata_url form selection - server_name = "sample_docs" - base_url = "http://localhost:4000" - path = scope.get("_original_path") or scope.get("path", "") or "" - if path.startswith(f"/{server_name}/mcp"): - resource_metadata_url = ( - f"{base_url}/.well-known/oauth-protected-resource/{server_name}/mcp" - ) - else: - resource_metadata_url = ( - f"{base_url}/.well-known/oauth-protected-resource/mcp/{server_name}" - ) - - assert resource_metadata_url == f"{base_url}{expected_metadata_path}", ( - f"resource_metadata_url {resource_metadata_url!r} does not match " - f"expected {base_url + expected_metadata_path!r}" - ) - - -# -------------------------------------------------------------------------- -# §2.1 — Regression: non-pass-through servers bypass is NOT applied -# -------------------------------------------------------------------------- - - -def test_is_mcp_passthrough_cold_start_false_for_oauth2_server(): - """Gateway-managed OAuth2 servers must not trigger the cold-start bypass.""" - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( - _is_mcp_passthrough_cold_start, - ) - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) - - global_mcp_server_manager.registry.clear() - oauth2_server = MCPServer( - server_id="oauth2-cold", - name="keycloak_whoami", - server_name="keycloak_whoami", - alias="keycloak_whoami", - url="https://upstream.example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - client_id="cid", - client_secret="cs", - authorization_url="https://keycloak/auth", - token_url="https://keycloak/token", - ) - global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server - - scope = _make_scope("/mcp/keycloak_whoami") - result = _is_mcp_passthrough_cold_start(scope, ["keycloak_whoami"], client_ip=None) - assert result is False - - -def test_is_mcp_passthrough_cold_start_false_for_empty_servers(): - """Aggregate /mcp route (no server list) must not trigger bypass.""" - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( - _is_mcp_passthrough_cold_start, - ) - - scope = _make_scope("/mcp") - assert _is_mcp_passthrough_cold_start(scope, None, client_ip=None) is False - assert _is_mcp_passthrough_cold_start(scope, [], client_ip=None) is False - - -# -------------------------------------------------------------------------- -# §2.1 — _parse_mcp_server_names_from_path -# -------------------------------------------------------------------------- - - -@pytest.mark.parametrize( - "path,expected", - [ - ("/mcp/sample_docs", ["sample_docs"]), - ("/mcp/sample_docs/tools/list", ["sample_docs"]), - ("/sample_docs/mcp", ["sample_docs"]), - ("/sample_docs/mcp/tools/list", ["sample_docs"]), - ("/mcp", None), - ("/mcp/", None), - ("/other/path", None), - ], -) -def test_parse_mcp_server_names_from_path(path, expected): - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( - _parse_mcp_server_names_from_path, - ) - - assert _parse_mcp_server_names_from_path(path) == expected diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py new file mode 100644 index 00000000000..bcfb442ba23 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py @@ -0,0 +1,153 @@ +"""Unit tests for MCP OAuth passthrough cold-start route behavior.""" + +import sys + +import pytest + +sys.path.insert(0, "../../../../../") + +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def _make_scope(path: str, headers: list = None) -> dict: + """Build a minimal ASGI HTTP scope for testing.""" + raw_headers = [(key.encode(), value.encode()) for key, value in (headers or [])] + return { + "type": "http", + "method": "POST", + "path": path, + "headers": raw_headers, + "query_string": b"", + "server": ("localhost", 4000), + "scheme": "http", + } + + +@pytest.mark.parametrize( + "route,expected_metadata_path", + [ + ( + "/mcp/sample_docs", + "/.well-known/oauth-protected-resource/mcp/sample_docs", + ), + ( + "/sample_docs/mcp", + "/.well-known/oauth-protected-resource/sample_docs/mcp", + ), + ], +) +def test_passthrough_cold_start_emits_401_with_matching_resource_metadata( + route, expected_metadata_path +): + """No auth headers on a passthrough server route emits matching metadata.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + _parse_mcp_server_names_from_path, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="pt-cold-start", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + if route.startswith("/mcp/"): + scope = _make_scope(route) + else: + scope = _make_scope("/mcp/sample_docs") + scope["_original_path"] = route + + servers = _parse_mcp_server_names_from_path(scope.get("path", "")) + assert _is_mcp_passthrough_cold_start(scope, servers, client_ip=None) is True + + server_name = "sample_docs" + base_url = "http://localhost:4000" + path = scope.get("_original_path") or scope.get("path", "") or "" + if path.startswith(f"/{server_name}/mcp"): + resource_metadata_url = ( + f"{base_url}/.well-known/oauth-protected-resource/{server_name}/mcp" + ) + else: + resource_metadata_url = ( + f"{base_url}/.well-known/oauth-protected-resource/mcp/{server_name}" + ) + + assert resource_metadata_url == f"{base_url}{expected_metadata_path}", ( + f"resource_metadata_url {resource_metadata_url!r} does not match " + f"expected {base_url + expected_metadata_path!r}" + ) + + +def test_is_mcp_passthrough_cold_start_false_for_oauth2_server(): + """Gateway-managed OAuth2 servers must not trigger the cold-start bypass.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + oauth2_server = MCPServer( + server_id="oauth2-cold", + name="keycloak_whoami", + server_name="keycloak_whoami", + alias="keycloak_whoami", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="cs", + authorization_url="https://keycloak/auth", + token_url="https://keycloak/token", + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + scope = _make_scope("/mcp/keycloak_whoami") + result = _is_mcp_passthrough_cold_start(scope, ["keycloak_whoami"], client_ip=None) + assert result is False + + +def test_is_mcp_passthrough_cold_start_false_for_empty_servers(): + """Aggregate /mcp route (no server list) must not trigger bypass.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + ) + + scope = _make_scope("/mcp") + assert _is_mcp_passthrough_cold_start(scope, None, client_ip=None) is False + assert _is_mcp_passthrough_cold_start(scope, [], client_ip=None) is False + + +@pytest.mark.parametrize( + "path,expected", + [ + ("/mcp/sample_docs", ["sample_docs"]), + ("/mcp/sample_docs/tools/list", ["sample_docs"]), + ("/sample_docs/mcp", ["sample_docs"]), + ("/sample_docs/mcp/tools/list", ["sample_docs"]), + ("/mcp", None), + ("/mcp/", None), + ("/other/path", None), + ], +) +def test_parse_mcp_server_names_from_path(path, expected): + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _parse_mcp_server_names_from_path, + ) + + assert _parse_mcp_server_names_from_path(path) == expected diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py new file mode 100644 index 00000000000..59029339583 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -0,0 +1,140 @@ +"""Unit tests for MCP OAuth passthrough tool-fetch behavior.""" + +import sys +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest + +sys.path.insert(0, "../../../../../") + +from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + _extract_upstream_auth_failure, +) +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def test_extract_upstream_auth_failure_finds_401_in_http_status_error(): + response = httpx.Response( + status_code=401, + headers={"www-authenticate": 'Bearer resource_metadata="https://x"'}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + exc = httpx.HTTPStatusError("401", request=response.request, response=response) + + result = _extract_upstream_auth_failure(exc) + assert result == (401, 'Bearer resource_metadata="https://x"') + + +def test_extract_upstream_auth_failure_walks_exception_group(): + response = httpx.Response( + status_code=401, + headers={"www-authenticate": "Bearer"}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + inner = httpx.HTTPStatusError("401", request=response.request, response=response) + + try: + raise ExceptionGroup("wrapped", [inner]) # noqa: F821 (PEP 654, py3.11+) + except Exception as group: + result = _extract_upstream_auth_failure(group) + + assert result == (401, "Bearer") + + +def test_extract_upstream_auth_failure_returns_none_for_non_auth(): + assert _extract_upstream_auth_failure(RuntimeError("boom")) is None + + +@pytest.mark.asyncio +async def test_fetch_tools_from_passthrough_raises_on_upstream_401(): + manager = MCPServerManager() + passthrough_server = MCPServer( + server_id="p1", + name="sample_docs", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + ) + + response = httpx.Response( + status_code=401, + headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + upstream_error = httpx.HTTPStatusError( + "401", request=response.request, response=response + ) + + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(side_effect=upstream_error) + + with pytest.raises(MCPUpstreamAuthError) as exc_info: + await manager._fetch_tools_with_timeout( + mock_client, passthrough_server.name, server=passthrough_server + ) + + assert exc_info.value.status_code == 401 + assert exc_info.value.www_authenticate == ( + 'Bearer resource_metadata="https://upstream"' + ) + assert exc_info.value.server_name == "sample_docs" + mock_client.list_tools.assert_awaited_with(raise_on_error=True) + + +@pytest.mark.asyncio +async def test_fetch_tools_from_passthrough_returns_tools_on_success(): + manager = MCPServerManager() + passthrough_server = MCPServer( + server_id="p1", + name="sample_docs", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + ) + + tool = MagicMock() + tool.name = "list_documents" + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(return_value=[tool]) + + tools = await manager._fetch_tools_with_timeout( + mock_client, passthrough_server.name, server=passthrough_server + ) + assert tools == [tool] + + +@pytest.mark.asyncio +async def test_fetch_tools_from_gateway_managed_swallows_errors(): + """Regression guard: non-pass-through servers keep returning [] on errors.""" + manager = MCPServerManager() + oauth2_server = MCPServer( + server_id="o1", + name="keycloak_whoami", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + + response = httpx.Response( + status_code=401, + headers={}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + upstream_error = httpx.HTTPStatusError( + "401", request=response.request, response=response + ) + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(side_effect=upstream_error) + + tools = await manager._fetch_tools_with_timeout( + mock_client, oauth2_server.name, server=oauth2_server + ) + assert tools == [] + mock_client.list_tools.assert_awaited_with(raise_on_error=False)