diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 2149f079a3d..77d0c2e8432 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -176,6 +176,39 @@ if MCP_AVAILABLE: ) return None + async def _get_user_byok_auth_header( + server, + user_api_key_dict: UserAPIKeyAuth, + ) -> Optional[str]: + """ + For BYOK servers, return the user's stored per-user key as the raw + credential string so the REST tools-list path injects it the same way + the MCP protocol tool-call path does in ``execute_mcp_tool``. The + auth_type formatting (Bearer / x-api-key / ...) is applied downstream by + the MCP client. Returns None for non-BYOK servers or when no credential + is stored. + """ + if not getattr(server, "is_byok", False): + return None + user_id = getattr(user_api_key_dict, "user_id", None) + server_id = getattr(server, "server_id", None) + if not user_id or not server_id: + return None + try: + from litellm.proxy._experimental.mcp_server.db import get_user_credential + from litellm.proxy.utils import get_prisma_client_or_throw + + prisma_client = get_prisma_client_or_throw( + "Database not connected. Connect a database to use BYOK MCP tools." + ) + return await get_user_credential(prisma_client, user_id, server_id) + except Exception as e: + verbose_logger.warning( + f"_get_user_byok_auth_header: failed to retrieve credential for " + f"user={user_id} server={server_id}: {e}" + ) + return None + async def _prefetch_user_oauth_creds( user_api_key_dict: UserAPIKeyAuth, ) -> Dict[str, Dict[str, Any]]: @@ -527,6 +560,10 @@ if MCP_AVAILABLE: server_auth_header = _get_server_auth_header( server, mcp_server_auth_headers, mcp_auth_header ) + if not server_auth_header: + server_auth_header = await _get_user_byok_auth_header( + server, user_api_key_dict + ) user_oauth_extra_headers = await _get_user_oauth_extra_headers( server, user_api_key_dict ) @@ -692,6 +729,10 @@ if MCP_AVAILABLE: server_auth_header = _get_server_auth_header( server, mcp_server_auth_headers, mcp_auth_header ) + if not server_auth_header: + server_auth_header = await _get_user_byok_auth_header( + server, user_api_key_dict + ) user_oauth_extra_headers = await _get_user_oauth_extra_headers( server, user_api_key_dict, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 47c9396f121..614618a433a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -546,6 +546,174 @@ class TestListToolsRestAPI: assert result["error"] is None assert result["message"] == "Successfully retrieved tools" + async def test_injects_stored_byok_credential_for_byok_server(self, monkeypatch): + """A BYOK server with no incoming auth header must have the user's stored + per-user key resolved and forwarded as the server auth header, so the UI + tools playground lists tools the same way the protocol tool-call path works.""" + import litellm.proxy._experimental.mcp_server.db as mcp_db + import litellm.proxy.utils as proxy_utils + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + server_id = "server-1" + alias = "server-1" + server_name = "server-1" + name = "stub" + auth_type = MCPAuth.bearer_token + is_byok = True + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + + stub_server = StubServer() + captured = {} + + async def fake_get_tools( + server, + server_auth_header, + raw_headers=None, + user_api_key_auth=None, + extra_headers=None, + apply_tool_filters=True, + ): + captured["auth_header"] = server_auth_header + return ["tool-1"] + + cred_calls = {} + + async def fake_get_user_credential(prisma_client, user_id, server_id): + cred_calls["args"] = (user_id, server_id) + return "user-byok-key" + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_get_tools_for_single_server", + fake_get_tools, + raising=False, + ) + monkeypatch.setattr( + proxy_utils, + "get_prisma_client_or_throw", + lambda msg: object(), + raising=False, + ) + monkeypatch.setattr( + mcp_db, "get_user_credential", fake_get_user_credential, raising=False + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + result = await rest_endpoints.list_tool_rest_api( + request, + server_id="server-1", + user_api_key_dict=UserAPIKeyAuth(user_id="user-1"), + ) + + assert captured["auth_header"] == "user-byok-key" + assert cred_calls["args"] == ("user-1", "server-1") + assert result["tools"] == ["tool-1"] + + async def test_does_not_resolve_byok_for_non_byok_server(self, monkeypatch): + """A non-BYOK server must not trigger a per-user credential lookup.""" + import litellm.proxy._experimental.mcp_server.db as mcp_db + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + server_id = "server-1" + alias = "server-1" + server_name = "server-1" + name = "stub" + auth_type = MCPAuth.bearer_token + is_byok = False + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + + stub_server = StubServer() + captured = {} + + async def fake_get_tools( + server, + server_auth_header, + raw_headers=None, + user_api_key_auth=None, + extra_headers=None, + apply_tool_filters=True, + ): + captured["auth_header"] = server_auth_header + return [] + + cred_calls = {"count": 0} + + async def fake_get_user_credential(prisma_client, user_id, server_id): + cred_calls["count"] += 1 + return "should-not-be-used" + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_get_tools_for_single_server", + fake_get_tools, + raising=False, + ) + monkeypatch.setattr( + mcp_db, "get_user_credential", fake_get_user_credential, raising=False + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + await rest_endpoints.list_tool_rest_api( + request, + server_id="server-1", + user_api_key_dict=UserAPIKeyAuth(user_id="user-1"), + ) + + assert cred_calls["count"] == 0 + assert captured["auth_header"] is None + async def test_include_disabled_tools_is_admin_only(self, monkeypatch): """include_disabled_tools skips the allowlist filter only for PROXY_ADMIN; a non-admin passing it stays filtered so the REST endpoint can't be used