fix(mcp): cache toolset permission lookups to avoid per-request DB calls

This commit is contained in:
Ishaan Jaffer 2026-03-21 20:34:32 -07:00
parent da6e8a0aad
commit e883653326
2 changed files with 58 additions and 11 deletions

View file

@ -11,6 +11,7 @@ import datetime
import hashlib
import json
import re
import time
from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from urllib.parse import urlparse
@ -182,6 +183,8 @@ class MCPServerManager:
}
"""
self._toolset_perm_cache: Dict[str, Tuple[Dict[str, List[str]], float]] = {}
def get_registry(self) -> Dict[str, MCPServer]:
"""
Get the registered MCP Servers from the registry and union with the config MCP Servers
@ -810,8 +813,7 @@ class MCPServerManager:
Resolve a list of toolset IDs into a mcp_tool_permissions dict.
Returns: {server_id: [tool_name, ...]} the union of all tools across
the given toolsets. This is merged (union semantics) into the key's
existing mcp_tool_permissions before access-control filtering runs.
the given toolsets. Results are cached for 60 s to avoid per-request DB queries.
"""
from litellm.proxy._experimental.mcp_server.toolset_db import list_mcp_toolsets
from litellm.proxy.proxy_server import prisma_client
@ -819,24 +821,44 @@ class MCPServerManager:
if not toolset_ids or prisma_client is None:
return {}
cache_key = ",".join(sorted(toolset_ids))
cached_entry = self._toolset_perm_cache.get(cache_key)
if cached_entry is not None:
result, cached_at = cached_entry
if time.time() - cached_at < 60:
return result
try:
toolsets = await list_mcp_toolsets(prisma_client, toolset_ids=toolset_ids)
tool_permissions: Dict[str, List[str]] = {}
for toolset in toolsets:
for tool in toolset.tools:
# Stored tool_names may include the server prefix (e.g.
# "server_alias__tool_name"). filter_tools_by_key_team_permissions
# compares against the unprefixed name, so strip it here.
raw_name = tool["tool_name"]
unprefixed, _ = split_server_prefix_from_name(raw_name)
tool_permissions.setdefault(tool["server_id"], [])
if unprefixed not in tool_permissions[tool["server_id"]]:
tool_permissions[tool["server_id"]].append(unprefixed)
self._toolset_perm_cache[cache_key] = (tool_permissions, time.time())
return tool_permissions
except Exception as e:
verbose_logger.warning(f"Failed to resolve toolset permissions: {str(e)}")
return {}
def invalidate_toolset_cache(self, toolset_id: Optional[str] = None) -> None:
"""Evict cached toolset permission entries.
Called after create/update/delete of a toolset so stale data is not served.
Pass toolset_id to evict only entries containing that ID, or None to clear all.
"""
if toolset_id is None:
self._toolset_perm_cache.clear()
return
keys_to_remove = [
k for k in self._toolset_perm_cache if toolset_id in k.split(",")
]
for k in keys_to_remove:
del self._toolset_perm_cache[k]
def filter_server_ids_by_ip(
self, server_ids: List[str], client_ip: Optional[str]
) -> List[str]:

View file

@ -2072,7 +2072,13 @@ if MCP_AVAILABLE:
touched_by = (
litellm_changed_by or user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME
)
return await create_mcp_toolset(prisma_client, payload, touched_by)
result = await create_mcp_toolset(prisma_client, payload, touched_by)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager.invalidate_toolset_cache()
return result
@router.get(
"/toolset",
@ -2093,10 +2099,16 @@ if MCP_AVAILABLE:
return await list_mcp_toolsets(prisma_client)
op = user_api_key_dict.object_permission
allowed_ids = (getattr(op, "mcp_toolsets", None) or []) if op else []
return await list_mcp_toolsets(
prisma_client, toolset_ids=allowed_ids if allowed_ids else None
)
# Distinguish None (field absent = no restriction) from [] (explicitly empty = zero allowed).
raw_toolsets = getattr(op, "mcp_toolsets", None) if op else None
# raw_toolsets is None → field not set → no restriction, return all
# raw_toolsets is [] → explicitly empty → return nothing
# raw_toolsets is [ids] → return only those
if raw_toolsets is None:
return await list_mcp_toolsets(prisma_client)
if not raw_toolsets:
return []
return await list_mcp_toolsets(prisma_client, toolset_ids=raw_toolsets)
@router.get(
"/toolset/{toolset_id}",
@ -2139,7 +2151,15 @@ if MCP_AVAILABLE:
touched_by = (
litellm_changed_by or user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME
)
return await update_mcp_toolset(prisma_client, payload, touched_by)
result = await update_mcp_toolset(prisma_client, payload, touched_by)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager.invalidate_toolset_cache(
getattr(payload, "toolset_id", None)
)
return result
@router.delete(
"/toolset/{toolset_id}",
@ -2166,4 +2186,9 @@ if MCP_AVAILABLE:
status_code=status.HTTP_404_NOT_FOUND,
detail={"error": f"Toolset '{toolset_id}' not found."},
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
global_mcp_server_manager.invalidate_toolset_cache(toolset_id)
return Response(status_code=status.HTTP_202_ACCEPTED)