mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): keep discovery and OAuth metadata caches across an OpenAPI spec re-read
add_server and update_server ran the full server-definition invalidation a second time after the awaited OpenAPI spec fetch, which also dropped the prompts/resources/templates discovery entries and the OAuth protected-resource metadata filled under the already-published definition, so the next request went upstream again. Only the listed-tool catalog recorded during the fetch holds pre-save entries, so the post-fetch pass now drops just that catalog and bumps its generation via the new _drop_listed_tools helper, which the full invalidation also calls.
This commit is contained in:
parent
71be87b893
commit
46639297ad
2 changed files with 64 additions and 3 deletions
|
|
@ -3352,7 +3352,7 @@ class MCPServerManager:
|
|||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Added MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -3391,7 +3391,7 @@ class MCPServerManager:
|
|||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
if new_server.spec_path:
|
||||
self._invalidate_server_definition_caches(mcp_server.server_id)
|
||||
self._drop_listed_tools(mcp_server.server_id)
|
||||
self.prime_oauth_metadata_discovery(new_server)
|
||||
verbose_logger.debug("Updated MCP Server: %s", new_server.name)
|
||||
|
||||
|
|
@ -4675,9 +4675,12 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._drop_listed_tools(server_id)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _drop_listed_tools(self, server_id: str) -> None:
|
||||
self._listed_tools_by_server_id.pop(server_id, None)
|
||||
self._listed_tools_generations[server_id] = self._listed_tools_generations.get(server_id, 0) + 1
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
|
||||
"""Key the listed-tool cache by every request input that can change the served catalog.
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import json
|
|||
import logging
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
|
@ -40,6 +41,7 @@ from mcp.types import Tool as MCPTool
|
|||
from pydantic import AnyUrl, TypeAdapter
|
||||
|
||||
from litellm.constants import MCP_METADATA_TIMEOUT
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
ListedToolsCaller,
|
||||
|
|
@ -7267,6 +7269,62 @@ class TestMCPServerManager:
|
|||
assert manager.registry["srv"] is new
|
||||
assert manager.get_listed_tool(new, "search", caller) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("already_registered", [False, True], ids=["add_server", "update_server"])
|
||||
async def test_openapi_spec_re_read_keeps_discovery_and_oauth_metadata_filled_during_the_fetch(
|
||||
self, already_registered: bool
|
||||
):
|
||||
"""The listed-tool catalog recorded during the spec fetch holds pre-save entries, but a prompts
|
||||
discovery or OAuth protected-resource fetch answered in that window already saw the published
|
||||
definition; dropping those too sends the next request upstream again."""
|
||||
manager = MCPServerManager()
|
||||
if already_registered:
|
||||
manager.registry["srv"] = MCPServer(
|
||||
server_id="srv", name="srv", transport=MCPTransport.http, url="http://old", spec_path="/old.json"
|
||||
)
|
||||
new = MCPServer(
|
||||
server_id="srv", name="srv", transport=MCPTransport.http, url="http://new", spec_path="/new.json"
|
||||
)
|
||||
caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister"))
|
||||
metadata_key: Final = (new.server_id, new.url)
|
||||
prompt_fetches = 0
|
||||
|
||||
async def fetch_prompts() -> list[Prompt]:
|
||||
nonlocal prompt_fetches
|
||||
prompt_fetches += 1
|
||||
return [Prompt(name="greet")]
|
||||
|
||||
async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None:
|
||||
manager._record_listed_tools(
|
||||
server,
|
||||
[MCPTool(name="search", description="pre-save", inputSchema={})],
|
||||
caller,
|
||||
manager._listed_tools_generations.get(server.server_id, 0),
|
||||
)
|
||||
await manager._prompt_discovery_cache.get((server.server_id, None), fetch_prompts)
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE[metadata_key] = (time.time() + 300, {"resource": new.url})
|
||||
|
||||
manager.build_mcp_server_from_table = AsyncMock(return_value=new)
|
||||
manager._maybe_register_openapi_tools = register_while_discovery_fills
|
||||
manager.prime_oauth_metadata_discovery = MagicMock()
|
||||
record = LiteLLM_MCPServerTable(
|
||||
server_id="srv", server_name="srv", url="http://new", transport=MCPTransport.http
|
||||
)
|
||||
save = manager.update_server if already_registered else manager.add_server
|
||||
|
||||
try:
|
||||
await save(record)
|
||||
|
||||
assert manager.registry["srv"] is new
|
||||
assert manager.get_listed_tool(new, "search", caller) is None
|
||||
prompts = await manager._prompt_discovery_cache.get((new.server_id, None), fetch_prompts)
|
||||
assert [prompt.name for prompt in prompts] == ["greet"]
|
||||
assert prompt_fetches == 1, "the prompts list filled after the save was published went upstream again"
|
||||
cached_metadata = discoverable_endpoints._OAUTH_METADATA_CACHE.get(metadata_key)
|
||||
assert cached_metadata is not None and cached_metadata[1] == {"resource": new.url}
|
||||
finally:
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(metadata_key, None)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_oauth_refresh_keeps_listed_tools(self):
|
||||
manager = MCPServerManager()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue