diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 8dbee1daa36..5ff237124f0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -4568,5 +4568,130 @@ class TestGetPublicMCPServersLegacyMode: assert sorted(s.server_id for s in result) == ["a", "b"] +class TestCreateMcpClientV2Graft: + """The PR4 v2-resolver graft in ``_create_mcp_client``. + + Migrated HTTP/SSE modes (``none`` plus the static ``api_key`` family) resolve through the + injected ``UpstreamCredentialProvider`` into the ``resolved_auth`` slot; every other mode, + and every stdio server, defers to v1's ``auth_value`` path unchanged. + """ + + def _http_server(self, **overrides: Any) -> MCPServer: + base: Dict[str, Any] = dict( + server_id="http-graft", + name="graft_server", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + ) + base.update(overrides) + return MCPServer(**base) + + async def test_none_mode_resolves_to_noop_auth(self): + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + NoOpAuth, + ) + + client = await MCPServerManager()._create_mcp_client( + self._http_server(auth_type=None) + ) + + assert isinstance(client._resolved_auth, NoOpAuth) + assert client._mcp_auth_value is None + + @pytest.mark.parametrize( + "auth_type, token, expected_name, expected_value", + [ + (MCPAuth.api_key, "k-123", "X-API-Key", "k-123"), + (MCPAuth.bearer_token, "b-123", "Authorization", "Bearer b-123"), + (MCPAuth.token, "t-123", "Authorization", "token t-123"), + (MCPAuth.authorization, "raw-123", "Authorization", "raw-123"), + ], + ) + async def test_static_family_emits_expected_header( + self, auth_type, token, expected_name, expected_value + ): + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + StaticHeaderAuth, + ) + + client = await MCPServerManager()._create_mcp_client( + self._http_server(auth_type=auth_type, authentication_token=token) + ) + + assert isinstance(client._resolved_auth, StaticHeaderAuth) + assert client._resolved_auth.header_name == expected_name + assert client._resolved_auth._header_value.get_secret_value() == expected_value + assert client._mcp_auth_value is None + + async def test_basic_mode_base64_encodes(self): + import base64 + + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + StaticHeaderAuth, + ) + + client = await MCPServerManager()._create_mcp_client( + self._http_server(auth_type=MCPAuth.basic, authentication_token="user:pass") + ) + + encoded = base64.b64encode(b"user:pass").decode() + assert isinstance(client._resolved_auth, StaticHeaderAuth) + assert client._resolved_auth.header_name == "Authorization" + assert ( + client._resolved_auth._header_value.get_secret_value() == f"Basic {encoded}" + ) + + async def test_deferred_mode_uses_v1_auth_value(self): + client = await MCPServerManager()._create_mcp_client( + self._http_server( + auth_type=MCPAuth.oauth2, authentication_token="legacy-token" + ) + ) + + assert client._resolved_auth is None + assert client._mcp_auth_value == "legacy-token" + + async def test_static_token_missing_defers_to_v1(self): + client = await MCPServerManager()._create_mcp_client( + self._http_server(auth_type=MCPAuth.api_key, authentication_token=None) + ) + + assert client._resolved_auth is None + + async def test_stdio_migrated_auth_type_still_defers_to_v1(self): + client = await MCPServerManager()._create_mcp_client( + MCPServer( + server_id="stdio-graft", + name="stdio_graft", + transport=MCPTransport.stdio, + command="node", + args=["server.js"], + auth_type=MCPAuth.api_key, + authentication_token="k-stdio", + ) + ) + + assert client.transport_type == MCPTransport.stdio + assert client._resolved_auth is None + assert client._mcp_auth_value == "k-stdio" + + async def test_resolver_error_maps_to_http_exception(self): + from litellm.proxy._experimental.mcp_server.outbound_credentials import Error + from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + CredError, + ) + + class _UnauthorizedProvider: + async def resolve_credentials(self, subject, server): + return Error(CredError.of_unauthorized("denied")) + + manager = MCPServerManager(cred_provider=_UnauthorizedProvider()) + + with pytest.raises(HTTPException) as exc: + await manager._create_mcp_client(self._http_server(auth_type=None)) + + assert exc.value.status_code == 401 + + if __name__ == "__main__": pytest.main([__file__])