mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
Merge pull request #42072 from BerriAI/litellm_mcp_cold_worker_tools_call
fix(mcp): tools/call no longer 404s on a worker that has not served tools/list
This commit is contained in:
commit
ef7da9b49f
4 changed files with 306 additions and 12 deletions
|
|
@ -2867,7 +2867,7 @@ class MCPServerManager:
|
|||
normalize_server_name(value) for value in (*iter_known_server_prefixes(server), server.name) if value
|
||||
)
|
||||
|
||||
def _server_exposes_tool(self, server: MCPServer, tool_name: str) -> bool:
|
||||
def server_exposes_tool(self, server: MCPServer, tool_name: str) -> bool:
|
||||
owned: Final = self._owned_mapping_values(server)
|
||||
mapped_owners: Final = (
|
||||
self.tool_name_to_mcp_server_name_mapping.get(spelling)
|
||||
|
|
@ -2875,6 +2875,20 @@ class MCPServerManager:
|
|||
)
|
||||
return any(owner is not None and normalize_server_name(owner) in owned for owner in mapped_owners)
|
||||
|
||||
def _known_prefix_to_server(self) -> Mapping[str, MCPServer]:
|
||||
"""Every prefix form a tool name may carry, keyed to its server; a form two servers share
|
||||
stays with the one registered first."""
|
||||
return {
|
||||
normalize_server_name(known_prefix): server
|
||||
for server in reversed(tuple(self.get_registry().values()))
|
||||
for known_prefix in iter_known_server_prefixes(server)
|
||||
}
|
||||
|
||||
def server_owning_tool_name_prefix(self, tool_name: str) -> MCPServer | None:
|
||||
prefix_to_server: Final = self._known_prefix_to_server()
|
||||
matched: Final = match_known_server_prefix(tool_name, prefix_to_server.keys())
|
||||
return None if matched is None else prefix_to_server.get(matched[0])
|
||||
|
||||
def remove_server(self, mcp_server: LiteLLM_MCPServerTable):
|
||||
"""
|
||||
Remove a server from the registry
|
||||
|
|
@ -6114,7 +6128,7 @@ class MCPServerManager:
|
|||
if mcp_server is None:
|
||||
raise ValueError(f"Tool {name} not found")
|
||||
|
||||
if resolved_by_server_name_only and not self._server_exposes_tool(mcp_server, name):
|
||||
if resolved_by_server_name_only and not self.server_exposes_tool(mcp_server, name):
|
||||
raise ValueError(f"Tool {name} not found")
|
||||
|
||||
return mcp_server
|
||||
|
|
@ -6475,15 +6489,7 @@ class MCPServerManager:
|
|||
MCPServer if found, None otherwise
|
||||
"""
|
||||
registry_servers: Final = list(self.get_registry().values())
|
||||
|
||||
# Build prefix → server lookup covering every known form a tool name
|
||||
# may take (alias / server_name / server_id / short ID). This is what
|
||||
# makes the short-prefix mode work without breaking historical names.
|
||||
prefix_to_server: Final[dict[str, MCPServer]] = {}
|
||||
for server in registry_servers:
|
||||
for known_prefix in iter_known_server_prefixes(server):
|
||||
normalised = normalize_server_name(known_prefix)
|
||||
prefix_to_server.setdefault(normalised, server)
|
||||
prefix_to_server: Final = self._known_prefix_to_server()
|
||||
|
||||
# First try with the original tool name
|
||||
if tool_name in self.tool_name_to_mcp_server_name_mapping:
|
||||
|
|
@ -6501,7 +6507,7 @@ class MCPServerManager:
|
|||
if matched is not None:
|
||||
matched_prefix, original_tool_name = matched
|
||||
matched_server: Final = prefix_to_server.get(matched_prefix)
|
||||
if matched_server is not None and self._server_exposes_tool(matched_server, original_tool_name):
|
||||
if matched_server is not None and self.server_exposes_tool(matched_server, original_tool_name):
|
||||
return matched_server
|
||||
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -2888,6 +2888,40 @@ if MCP_AVAILABLE:
|
|||
headers={"WWW-Authenticate": get_byok_www_authenticate()},
|
||||
)
|
||||
|
||||
async def _list_tools_before_first_call(
|
||||
server: MCPServer | None,
|
||||
tool_name: str,
|
||||
allowed_mcp_servers: list[MCPServer],
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
mcp_auth_header: str | None,
|
||||
mcp_server_auth_headers: dict[str, dict[str, str]] | None,
|
||||
oauth2_headers: dict[str, str] | None,
|
||||
raw_headers: dict[str, str] | None,
|
||||
) -> None:
|
||||
"""List ``server`` with the caller's own credentials when it does not yet expose ``tool_name`` here.
|
||||
|
||||
The startup fill skips a server whose upstream wants the caller's token, and mcp 2 no
|
||||
longer lists before an uncached tools/call, so a worker that has not served tools/list
|
||||
for this caller would otherwise answer 404 for a tool the caller can see. Gating on the
|
||||
requested tool, not on any prior listing, keeps callers with different upstream catalogs
|
||||
from masking each other.
|
||||
"""
|
||||
if server is None or global_mcp_server_manager.server_exposes_tool(server, tool_name):
|
||||
return
|
||||
if all(allowed.server_id != server.server_id for allowed in allowed_mcp_servers):
|
||||
return
|
||||
try:
|
||||
await _get_tools_from_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_servers=[server.server_id],
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before
|
||||
verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e)
|
||||
|
||||
async def execute_mcp_tool(
|
||||
name: str,
|
||||
arguments: dict[str, object],
|
||||
|
|
@ -2948,6 +2982,27 @@ if MCP_AVAILABLE:
|
|||
all_registry_prefixes.add(normalize_server_name(known_prefix))
|
||||
name_is_prefixed = is_tool_name_prefixed(name, known_server_prefixes=all_registry_prefixes)
|
||||
|
||||
first_call_target: Final = (
|
||||
requested_server
|
||||
if requested_server is not None and not name_is_prefixed
|
||||
else global_mcp_server_manager.server_owning_tool_name_prefix(name)
|
||||
)
|
||||
first_call_tool_name: Final = (
|
||||
name
|
||||
if first_call_target is None or (requested_server is not None and not name_is_prefixed)
|
||||
else strip_known_server_prefix(name, first_call_target)
|
||||
)
|
||||
await _list_tools_before_first_call(
|
||||
server=first_call_target,
|
||||
tool_name=first_call_tool_name,
|
||||
allowed_mcp_servers=allowed_mcp_servers,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
)
|
||||
|
||||
if requested_server is not None and not name_is_prefixed:
|
||||
# REST callers may pass server_id with the upstream tool name (no
|
||||
# LiteLLM prefix). The first segment is not a registered server
|
||||
|
|
|
|||
|
|
@ -7437,6 +7437,186 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool
|
|||
assert captured["name"] == "echo"
|
||||
|
||||
|
||||
def _never_listed_passthrough_server() -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id="lazy-map-1",
|
||||
name="lazy_map",
|
||||
server_name="lazy_map",
|
||||
url="https://up.example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.true_passthrough,
|
||||
)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _worker_that_never_listed(server: MCPServer, upstream_tools: tuple[str, ...]):
|
||||
"""A worker whose tool rows for ``server`` are empty, in front of an upstream that answers
|
||||
tools/list with ``upstream_tools`` and a managed dispatch that records what reaches it."""
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
mcp_module.global_mcp_server_manager.registry[server.server_id] = server
|
||||
dispatched: dict[str, object] = {}
|
||||
|
||||
async def fake_handle_managed_mcp_tool(**kwargs):
|
||||
dispatched.update(kwargs)
|
||||
return CallToolResult(content=[TextContent(type="text", text="ok")], is_error=False)
|
||||
|
||||
async def fake_fetch_tools(client, server_name):
|
||||
return [MCPTool(name=tool_name, inputSchema={}) for tool_name in upstream_tools]
|
||||
|
||||
with (
|
||||
patch.object( # test-quality-ok: the upstream MCP session is the boundary; a real one needs an initialize handshake over a live server
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"_create_mcp_client",
|
||||
new=AsyncMock(return_value=MagicMock()),
|
||||
) as create_client,
|
||||
patch.object( # test-quality-ok: same boundary, this is the tools/list answer the upstream would give
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"_fetch_tools_with_timeout",
|
||||
side_effect=fake_fetch_tools,
|
||||
) as fetch_tools,
|
||||
patch.object( # test-quality-ok: records the resolved server and bare name the managed call would forward upstream
|
||||
mcp_module, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool
|
||||
),
|
||||
):
|
||||
yield SimpleNamespace(create_client=create_client, fetch_tools=fetch_tools, dispatched=dispatched)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_lists_never_listed_passthrough_server_with_caller_token_first():
|
||||
"""A prefixed tools/call on a worker that has not served tools/list must list that server once
|
||||
with the caller's own credentials and then dispatch, instead of answering 404."""
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
server = _never_listed_passthrough_server()
|
||||
with _worker_that_never_listed(server, upstream_tools=("add",)) as worker:
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="lazy_map-add",
|
||||
arguments={"a": 1, "b": 2},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
raw_headers={"authorization": "Bearer caller-token"},
|
||||
)
|
||||
|
||||
assert worker.fetch_tools.await_count == 1
|
||||
assert "caller-token" in str(worker.create_client.await_args.kwargs.get("mcp_auth_header"))
|
||||
assert worker.dispatched["server_name"] == "lazy_map"
|
||||
assert worker.dispatched["name"] == "add"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_rest_server_id_lists_never_listed_server_first():
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
server = _never_listed_passthrough_server()
|
||||
with _worker_that_never_listed(server, upstream_tools=("add",)) as worker:
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="add",
|
||||
arguments={"a": 1, "b": 2},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
requested_server_id=server.server_id,
|
||||
)
|
||||
|
||||
assert worker.fetch_tools.await_count == 1
|
||||
assert worker.dispatched["server_name"] == "lazy_map"
|
||||
assert worker.dispatched["name"] == "add"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_unknown_tool_on_never_listed_server_lists_once_then_404s():
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
server = _never_listed_passthrough_server()
|
||||
with (
|
||||
_worker_that_never_listed(server, upstream_tools=("add",)) as worker,
|
||||
pytest.raises(HTTPException) as exc_info,
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="lazy_map-nope",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert worker.fetch_tools.await_count == 1
|
||||
assert worker.dispatched == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_does_not_relist_a_server_this_worker_already_listed():
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
server = _never_listed_passthrough_server()
|
||||
with _worker_that_never_listed(server, upstream_tools=("add",)) as worker:
|
||||
mcp_module.global_mcp_server_manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server)
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="lazy_map-add",
|
||||
arguments={"a": 1, "b": 2},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
)
|
||||
|
||||
assert worker.fetch_tools.await_count == 0
|
||||
assert worker.dispatched["name"] == "add"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_lists_a_tool_this_worker_has_not_yet_seen_on_a_listed_server():
|
||||
"""A worker that already holds one of the server's tools must still list when a caller asks
|
||||
for a different tool it has not cached, so callers with wider upstream catalogs are not 404ed."""
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
server = _never_listed_passthrough_server()
|
||||
with _worker_that_never_listed(server, upstream_tools=("add", "multiply")) as worker:
|
||||
mcp_module.global_mcp_server_manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server)
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="lazy_map-multiply",
|
||||
arguments={"a": 1, "b": 2},
|
||||
allowed_mcp_servers=[server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
)
|
||||
|
||||
assert worker.fetch_tools.await_count == 1
|
||||
assert worker.dispatched["server_name"] == "lazy_map"
|
||||
assert worker.dispatched["name"] == "multiply"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_never_lists_a_server_the_caller_cannot_access():
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
server = _never_listed_passthrough_server()
|
||||
other_server = MCPServer(server_id="other-1", name="other", transport=MCPTransport.http)
|
||||
with (
|
||||
_worker_that_never_listed(server, upstream_tools=("add",)) as worker,
|
||||
pytest.raises(HTTPException) as exc_info,
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="lazy_map-add",
|
||||
arguments={"a": 1, "b": 2},
|
||||
allowed_mcp_servers=[other_server],
|
||||
start_time=datetime.now(),
|
||||
mcp_auth_header="Bearer caller-token",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert worker.fetch_tools.await_count == 0
|
||||
assert worker.dispatched == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator():
|
||||
"""A server with no alias publishes its UUID server_id as the tool prefix.
|
||||
|
|
|
|||
|
|
@ -5556,6 +5556,59 @@ class TestMCPServerManager:
|
|||
)
|
||||
mock_inject.assert_awaited_once()
|
||||
|
||||
def test_server_owning_tool_name_prefix_is_known_before_the_server_is_ever_listed(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="lazy-map-1",
|
||||
name="lazy_map",
|
||||
server_name="lazy_map",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.true_passthrough,
|
||||
)
|
||||
manager.registry = {server.server_id: server}
|
||||
|
||||
assert manager._get_mcp_server_from_tool_name("lazy_map-add") is None
|
||||
assert manager.server_owning_tool_name_prefix("lazy_map-add") is server
|
||||
assert manager.server_owning_tool_name_prefix("someone_else-add") is None
|
||||
|
||||
manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server)
|
||||
|
||||
assert manager._get_mcp_server_from_tool_name("lazy_map-add") is server
|
||||
|
||||
def test_server_exposes_tool_is_per_tool_not_per_server(self):
|
||||
"""A tool listed for the server does not make its unlisted siblings look exposed."""
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="lazy-map-2",
|
||||
name="lazy_map",
|
||||
server_name="lazy_map",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.true_passthrough,
|
||||
)
|
||||
manager.registry = {server.server_id: server}
|
||||
|
||||
assert manager.server_exposes_tool(server, "add") is False
|
||||
|
||||
manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server)
|
||||
|
||||
assert manager.server_exposes_tool(server, "add") is True
|
||||
assert manager.server_exposes_tool(server, "lazy_map-add") is True
|
||||
assert manager.server_exposes_tool(server, "multiply") is False
|
||||
|
||||
def test_known_prefix_to_server_keeps_the_first_registered_owner_of_a_shared_prefix(self):
|
||||
manager = MCPServerManager()
|
||||
first = MCPServer(server_id="first-id", name="first", server_name="first", transport=MCPTransport.http)
|
||||
second = MCPServer(
|
||||
server_id="second-id", name="second", server_name="second", alias="first", transport=MCPTransport.http
|
||||
)
|
||||
manager.registry = {"first-id": first, "second-id": second}
|
||||
|
||||
prefix_to_server = manager._known_prefix_to_server()
|
||||
|
||||
assert prefix_to_server["first"] is first
|
||||
assert prefix_to_server["second"] is second
|
||||
assert manager.server_owning_tool_name_prefix("first-add") is first
|
||||
|
||||
def test_resolve_mcp_server_for_tool_call_via_prefixed_name(self):
|
||||
"""Resolution succeeds when the prefixed tool name is in the mapping."""
|
||||
manager = MCPServerManager()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue