fix(mcp): preserve mcp_access_groups in toolset scope, use shared Redis cache for toolset perms

- Remove mcp_access_groups=[] from _apply_toolset_scope (server.py) and the
  responses API toolset path (litellm_proxy_mcp_handler.py). A key's access-group
  grants remain valid even when the request is scoped to a single toolset; clearing
  them silently revoked legitimate entitlements.

- Switch resolve_toolset_tool_permissions and get_toolset_by_name_cached to use
  user_api_key_cache (Redis-backed DualCache in production) instead of per-instance
  in-memory dicts. Cache entries are now shared across workers, eliminating the
  per-worker stale-toolset-permission window flagged as a P1 by Greptile.

- Use union merge (set union of tool names per server) when applying toolset
  permissions in the responses API path so direct-server tool restrictions are not
  overwritten by toolset permissions.
This commit is contained in:
Ishaan Jaffer 2026-03-23 12:09:03 -07:00
parent 57fc329f48
commit 435ace9870
3 changed files with 79 additions and 42 deletions

View file

@ -11,7 +11,6 @@ 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
@ -183,8 +182,10 @@ class MCPServerManager:
}
"""
# Toolset caches are now stored in user_api_key_cache (Redis-backed DualCache
# in production) so entries are shared across workers. These attributes are
# kept as empty stubs so existing callers don't AttributeError during tests.
self._toolset_perm_cache: Dict[str, Tuple[Dict[str, List[str]], float]] = {}
# name → (toolset | None, cached_at) — 60 s TTL
self._toolset_name_cache: Dict[str, Tuple[Optional[Any], float]] = {}
def get_registry(self) -> Dict[str, MCPServer]:
@ -815,20 +816,21 @@ 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. Results are cached for 60 s to avoid per-request DB queries.
the given toolsets. Results are cached via ``user_api_key_cache`` (a
Redis-backed ``DualCache`` in production) so that cache entries are
shared across workers and cold-cache DB hits are minimised.
"""
from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
from litellm.proxy._experimental.mcp_server.toolset_db import list_mcp_toolsets
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
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
cache_key = "toolset_perms:" + ",".join(sorted(toolset_ids))
cached = await user_api_key_cache.async_get_cache(key=cache_key)
if cached is not None:
return cached
try:
toolsets = await list_mcp_toolsets(prisma_client, toolset_ids=toolset_ids)
@ -840,7 +842,11 @@ class MCPServerManager:
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())
await user_api_key_cache.async_set_cache(
key=cache_key,
value=tool_permissions,
ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
)
return tool_permissions
except Exception as e:
verbose_logger.warning(f"Failed to resolve toolset permissions: {str(e)}")
@ -850,47 +856,60 @@ class MCPServerManager:
"""Evict cached toolset permission entries.
Called after create/update/delete of a toolset so stale data is not served.
The in-memory layer of ``user_api_key_cache`` is cleared immediately;
Redis entries expire naturally after the configured TTL.
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()
self._toolset_name_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]
# Also evict any name cache entry that refers to this toolset_id
name_keys_to_remove = [
name
for name, (ts, _) in self._toolset_name_cache.items()
if ts is not None and getattr(ts, "toolset_id", None) == toolset_id
]
for k in name_keys_to_remove:
del self._toolset_name_cache[k]
# Clear the in-memory layer of the shared DualCache for affected keys.
# We can't enumerate Redis keys by pattern, so Redis entries expire via TTL.
try:
from litellm.proxy.proxy_server import user_api_key_cache
in_mem = getattr(user_api_key_cache, "in_memory_cache", None)
if in_mem is None:
return
cache_dict = getattr(in_mem, "cache_dict", {})
if toolset_id is None:
keys_to_remove = [k for k in cache_dict if k.startswith("toolset_")]
else:
keys_to_remove = [
k
for k in cache_dict
if k.startswith("toolset_") and toolset_id in k
]
for k in keys_to_remove:
cache_dict.pop(k, None)
except Exception:
pass
async def get_toolset_by_name_cached(
self,
prisma_client: Any,
toolset_name: str,
) -> Optional[Any]:
"""Return a toolset by name, using a 60 s in-memory TTL cache.
Avoids a DB hit on every MCP request routed through a named toolset.
"""Return a toolset by name, cached in ``user_api_key_cache`` (Redis-backed
``DualCache`` in production) to avoid a DB hit on every routed request.
"""
cached_entry = self._toolset_name_cache.get(toolset_name)
if cached_entry is not None:
toolset, cached_at = cached_entry
if time.time() - cached_at < 60:
return toolset
from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
from litellm.proxy.proxy_server import user_api_key_cache
cache_key = f"toolset_name:{toolset_name}"
cached = await user_api_key_cache.async_get_cache(key=cache_key)
if cached is not None:
# Sentinel value used to cache "not found" so we don't re-query for
# names that don't exist.
return None if cached == "__not_found__" else cached
from litellm.proxy._experimental.mcp_server.toolset_db import (
get_mcp_toolset_by_name,
)
toolset = await get_mcp_toolset_by_name(prisma_client, toolset_name)
self._toolset_name_cache[toolset_name] = (toolset, time.time())
await user_api_key_cache.async_set_cache(
key=cache_key,
value=toolset if toolset is not None else "__not_found__",
ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
)
return toolset
def filter_server_ids_by_ip(

View file

@ -2460,7 +2460,8 @@ if MCP_AVAILABLE:
"mcp_servers": server_ids,
"mcp_tool_permissions": tool_permissions,
"mcp_toolsets": [],
"mcp_access_groups": [],
# mcp_access_groups is preserved: a key's access-group grants
# remain valid even when the request is scoped to a single toolset.
}
)
else:

View file

@ -226,12 +226,21 @@ class LiteLLM_Proxy_MCP_Handler:
)
existing_op = user_api_key_auth.object_permission
if existing_op is not None:
# Merge toolset tool permissions with existing ones (union),
# so direct-server tool restrictions are not overwritten.
merged_tool_perms = dict(existing_op.mcp_tool_permissions or {})
for server_id, tool_names in tool_permissions.items():
existing_tools = merged_tool_perms.get(server_id, [])
merged_tool_perms[server_id] = list(
set(existing_tools) | set(tool_names)
)
updated_op = existing_op.model_copy(
update={
"mcp_servers": all_server_ids,
"mcp_tool_permissions": tool_permissions,
"mcp_tool_permissions": merged_tool_perms,
"mcp_toolsets": [],
"mcp_access_groups": [],
# mcp_access_groups preserved: existing access-group
# grants remain valid alongside toolset grants.
}
)
else:
@ -246,10 +255,18 @@ class LiteLLM_Proxy_MCP_Handler:
except Exception as _e:
verbose_logger.debug(f"Could not apply toolset permissions: {_e}")
# When toolsets were resolved we updated object_permission.mcp_servers to the
# full union (toolset server IDs + direct server names). Passing a name-based
# filter here would exclude those toolset server IDs (which are UUIDs, not
# names), so use None and let the auth object's mcp_servers do the filtering.
effective_server_filter = (
None if resolved_toolset_ids else (resolved_mcp_servers or None)
)
tools = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_servers=resolved_mcp_servers if resolved_mcp_servers else None,
mcp_servers=effective_server_filter,
mcp_server_auth_headers=mcp_server_auth_headers,
log_list_tools_to_spendlogs=True,
list_tools_log_source="responses",
@ -264,7 +281,7 @@ class LiteLLM_Proxy_MCP_Handler:
)
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=resolved_mcp_servers if resolved_mcp_servers else None,
mcp_servers=effective_server_filter,
allowed_mcp_servers=allowed_mcp_servers,
)