From b1affbd72e5bf547361406fb976ec62cae6c406b Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Thu, 11 Jun 2026 13:17:40 -0700 Subject: [PATCH] test(mcp): cover BYOK credential helper's guard, failure, and aggregator paths codecov/patch flagged uncovered new lines in the REST tools-list BYOK injection. The existing tests covered only the single-server happy path, leaving the helper's missing-user_id short-circuit, its credential-lookup exception handler, and the multi-server aggregator call site untested. Adds three tests so every new line in rest_endpoints.py is exercised: a BYOK server with no resolvable user_id skips the lookup, a lookup that raises is swallowed without injecting a header or failing the request, and the no-server_id aggregator path injects the stored credential the same way the single-server path does. --- .../mcp_server/test_rest_endpoints.py | 249 ++++++++++++++++++ 1 file changed, 249 insertions(+) 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 614618a433a..744f7c4836f 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 @@ -714,6 +714,255 @@ class TestListToolsRestAPI: assert cred_calls["count"] == 0 assert captured["auth_header"] is None + async def test_byok_skips_lookup_when_user_id_missing(self, monkeypatch): + """A BYOK server still must not query a credential when the caller has no + resolvable user_id; the helper short-circuits before touching the DB.""" + 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 = 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 [] + + 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=None), + ) + + assert cred_calls["count"] == 0 + assert captured["auth_header"] is None + + async def test_byok_credential_lookup_failure_is_swallowed(self, monkeypatch): + """If the per-user credential lookup raises, the helper logs and returns + None rather than failing the whole tools-list request.""" + 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"] + + async def fake_get_user_credential(prisma_client, user_id, server_id): + raise RuntimeError("db unavailable") + + 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"), + ) + + # Lookup failure must not inject a header nor break the request. + assert captured["auth_header"] is None + assert result["tools"] == ["tool-1"] + + async def test_injects_stored_byok_credential_in_aggregator_path(self, monkeypatch): + """The multi-server aggregator path (no server_id) injects the stored + per-user BYOK credential the same way the single-server path does.""" + 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"] + + async def fake_get_user_credential(prisma_client, 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=None, + user_api_key_dict=UserAPIKeyAuth(user_id="user-1"), + ) + + assert captured["auth_header"] == "user-byok-key" + assert result["tools"] == ["tool-1"] + 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