mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): degrade gracefully when BYOK credential lookup errors on protocol tools-list
The protocol tools-list aggregates servers in an asyncio.gather without return_exceptions, so a raise from the per-user BYOK credential lookup would abort the entire listing, not just the one server. A missing credential already returns None (lists without a key); guard the lookup so a DB error does the same instead of propagating, matching the REST helper's behavior.
This commit is contained in:
parent
21c69751c2
commit
4d9fabe11d
2 changed files with 90 additions and 3 deletions
|
|
@ -1771,9 +1771,18 @@ if MCP_AVAILABLE:
|
|||
# does. Without this a BYOK server lists with the shared token,
|
||||
# which is None, so the upstream rejects the listing.
|
||||
if server.is_byok and not server_auth_header:
|
||||
server_auth_header = await _get_byok_credential(
|
||||
server, user_api_key_auth
|
||||
)
|
||||
# Don't let a credential-lookup error abort the whole
|
||||
# multi-server listing; degrade to listing without the key,
|
||||
# matching the REST path's behavior.
|
||||
try:
|
||||
server_auth_header = await _get_byok_credential(
|
||||
server, user_api_key_auth
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"BYOK credential lookup failed for {server.server_id}; "
|
||||
f"listing without it: {e}"
|
||||
)
|
||||
|
||||
# Prefer server-stored per-user OAuth when configured, so a stale
|
||||
# Authorization header from the MCP client cannot override Redis/DB
|
||||
|
|
|
|||
|
|
@ -1023,6 +1023,84 @@ async def test_get_tools_from_mcp_servers_injects_byok_credential():
|
|||
assert captured["mcp_auth_header"] == "user-stored-byok-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_from_mcp_servers_byok_lookup_error_does_not_abort_listing():
|
||||
"""A DB error during the BYOK credential lookup must degrade to listing
|
||||
without the key, not propagate out of the gather and abort the whole list."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_server_module
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_get_tools_from_mcp_servers,
|
||||
set_auth_context,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="byok_user")
|
||||
set_auth_context(user_api_key_auth)
|
||||
|
||||
byok_server = MagicMock()
|
||||
byok_server.name = "byok_server"
|
||||
byok_server.alias = "byok"
|
||||
byok_server.allowed_tools = None
|
||||
byok_server.disallowed_tools = None
|
||||
byok_server.server_id = "byok_server"
|
||||
byok_server.server_name = "byok_server"
|
||||
byok_server.auth_type = "authorization"
|
||||
byok_server.extra_headers = None
|
||||
byok_server.is_byok = True
|
||||
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["byok_server"])
|
||||
mock_manager.get_mcp_server_by_id = lambda server_id: byok_server
|
||||
mock_manager.filter_server_ids_by_ip_with_info = lambda server_ids, client_ip: (
|
||||
server_ids,
|
||||
0,
|
||||
)
|
||||
|
||||
captured = {}
|
||||
|
||||
async def mock_get_tools_from_server(
|
||||
server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=None,
|
||||
add_prefix=True,
|
||||
raw_headers=None,
|
||||
**kwargs,
|
||||
):
|
||||
captured["mcp_auth_header"] = mcp_auth_header
|
||||
tool = MagicMock()
|
||||
tool.name = "byok_tool"
|
||||
tool.description = "BYOK tool"
|
||||
tool.inputSchema = {}
|
||||
return [tool]
|
||||
|
||||
mock_manager._get_tools_from_server = mock_get_tools_from_server
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
mcp_server_module,
|
||||
"_get_byok_credential",
|
||||
AsyncMock(side_effect=Exception("transient DB error")),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
):
|
||||
# Must not raise even though the credential lookup blew up.
|
||||
result = await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=None,
|
||||
mcp_servers=["byok_server"],
|
||||
)
|
||||
|
||||
# The listing still completed; the server was listed without the BYOK key.
|
||||
assert len(result) == 1
|
||||
assert result[0].name == "byok_tool"
|
||||
assert captured["mcp_auth_header"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
|
||||
"""Test that _get_tools_from_mcp_servers handles all servers failing gracefully"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue