mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(mcp): drop the OAuth discovery cache when a server definition changes
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
bef557f3ed
commit
3650243b90
3 changed files with 52 additions and 0 deletions
|
|
@ -141,6 +141,17 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None:
|
|||
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
|
||||
|
||||
|
||||
def invalidate_oauth_metadata_cache(server_id: str) -> None:
|
||||
"""Drop cached upstream IdP metadata for a server whose definition changed."""
|
||||
for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]:
|
||||
del _OAUTH_METADATA_CACHE[cache_key]
|
||||
for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]:
|
||||
lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
|
||||
if lock is None or lock.locked():
|
||||
continue
|
||||
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
|
||||
|
||||
|
||||
def encode_state_with_base_url(
|
||||
base_url: str,
|
||||
original_state: str,
|
||||
|
|
|
|||
|
|
@ -4555,8 +4555,13 @@ class MCPServerManager:
|
|||
self._template_discovery_cache.invalidate(server_id)
|
||||
|
||||
def _invalidate_server_definition_caches(self, server_id: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton
|
||||
invalidate_oauth_metadata_cache,
|
||||
)
|
||||
|
||||
self._invalidate_discovery_lists(server_id)
|
||||
self._listed_tools_by_server_id.pop(server_id, None)
|
||||
invalidate_oauth_metadata_cache(server_id)
|
||||
|
||||
def _discovers_per_caller(self, server: MCPServer) -> bool:
|
||||
return (
|
||||
|
|
|
|||
|
|
@ -12611,3 +12611,39 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session(
|
|||
proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called()
|
||||
proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called()
|
||||
proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_server_drops_cached_upstream_oauth_metadata():
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._types import LiteLLM_MCPServerTable
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="oauth-cache-server",
|
||||
name="oauth_cache_server",
|
||||
url="http://old-upstream/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
manager.registry[server.server_id] = server
|
||||
stale_key: Final = (server.server_id, server.url)
|
||||
other_key: Final = ("other-server", "http://other/mcp")
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE[stale_key] = (time.time() + 300, {"iss": "old-idp"})
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE[other_key] = (time.time() + 300, {"iss": "other"})
|
||||
try:
|
||||
await manager.update_server(
|
||||
LiteLLM_MCPServerTable(
|
||||
server_id=server.server_id,
|
||||
server_name=server.name,
|
||||
url="http://new-upstream/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
)
|
||||
assert stale_key not in discoverable_endpoints._OAUTH_METADATA_CACHE
|
||||
assert other_key in discoverable_endpoints._OAUTH_METADATA_CACHE
|
||||
finally:
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None)
|
||||
discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue