mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-20 00:11:50 +00:00
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.
This commit is contained in:
parent
babdc8145b
commit
b1affbd72e
1 changed files with 249 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue