mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-20 00:11:50 +00:00
fix(mcp): inject stored per-user BYOK credential in REST tools-list
The REST tools-list path (used by the UI MCP Tools playground) only resolved static/passthrough and OBO auth, never the per-user BYOK key, so BYOK servers always showed no tools. Resolve the stored credential when no header is present, matching the protocol tool-call path in execute_mcp_tool.
This commit is contained in:
parent
8f4389246d
commit
babdc8145b
2 changed files with 209 additions and 0 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue