perf(mcp): cache per-user env var values off the hot tool-call path

This commit is contained in:
mateo-berri 2026-06-04 04:31:24 +00:00 • committed by Claude
parent 21c7c38331
commit 4b31a6d711
No known key found for this signature in database
3 changed files with 134 additions and 4 deletions

View file

@ -124,6 +124,28 @@ _AZURE_ENTRA_HOSTS = {
"login.chinacloudapi.cn", # China
}
# Short-lived in-memory cache for per-user MCP env var values, mirroring the
# BYOK credential cache. Keyed by (user_id, server_id); value is
# (values_dict, monotonic_timestamp). Keeps the tool-call and tool-listing
# paths off the DB on every request within the TTL window.
_user_env_vars_cache: Dict[Tuple[str, str], Tuple[Dict[str, str], float]] = {}
_USER_ENV_VARS_CACHE_TTL = 60 # seconds
_USER_ENV_VARS_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth
def invalidate_user_env_vars_cache(user_id: str, server_id: str) -> None:
"""Drop a cached entry after the user stores or clears their env var values
so the next request reads the fresh value instead of a stale one."""
_user_env_vars_cache.pop((user_id, server_id), None)
def _write_user_env_vars_cache(
user_id: str, server_id: str, values: Dict[str, str]
) -> None:
if len(_user_env_vars_cache) >= _USER_ENV_VARS_CACHE_MAX_SIZE:
_user_env_vars_cache.clear()
_user_env_vars_cache[(user_id, server_id)] = (values, time.monotonic())
def _should_strip_caller_authorization(
mcp_server: MCPServer,
@ -1626,15 +1648,26 @@ class MCPServerManager:
) -> Dict[str, str]:
"""Look up the calling user's env var values for ``server``.
Returns an empty dict when no user is available. DB errors propagate
so the caller can decide between failing the request (tool-call path)
and staying best-effort (listing path).
Returns an empty dict when no user is available. Results are cached in a
short-lived in-memory map keyed by (user_id, server_id) so the tool-call
and tool-listing paths avoid a DB round-trip per request within the TTL
window; the cache is invalidated when the user stores or clears values.
DB errors propagate so the caller can decide between failing the request
(tool-call path) and staying best-effort (listing path).
"""
if user_api_key_auth is None:
return {}
user_id = getattr(user_api_key_auth, "user_id", None)
if not user_id:
return {}
cache_key = (user_id, server.server_id)
cached = _user_env_vars_cache.get(cache_key)
if cached is not None:
values, ts = cached
if time.monotonic() - ts < _USER_ENV_VARS_CACHE_TTL:
return values
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415
if prisma_client is None:
@ -1643,7 +1676,9 @@ class MCPServerManager:
get_user_env_vars,
)
return await get_user_env_vars(prisma_client, user_id, server.server_id)
values = await get_user_env_vars(prisma_client, user_id, server.server_id)
_write_user_env_vars_cache(user_id, server.server_id, values)
return values
async def _create_mcp_client(
self,

View file

@ -2307,6 +2307,11 @@ if MCP_AVAILABLE:
k: v for k, v in payload.values.items() if k in allowed_names and v != ""
}
await store_user_env_vars(prisma_client, user_id, server_id, filtered)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
invalidate_user_env_vars_cache,
)
invalidate_user_env_vars_cache(user_id, server_id)
return _compute_user_env_var_status(server=server, stored_values=filtered)
@router.delete(
@ -2337,6 +2342,11 @@ if MCP_AVAILABLE:
detail={"error": f"MCP Server {server_id} not found"},
)
await delete_user_env_vars(prisma_client, user_id, server_id)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
invalidate_user_env_vars_cache,
)
invalidate_user_env_vars_cache(user_id, server_id)
return _compute_user_env_var_status(server=server, stored_values={})
@router.get(

View file

@ -459,6 +459,91 @@ async def test_load_user_env_vars_returns_empty_when_db_unavailable(monkeypatch)
assert await manager._load_user_env_vars(server, fake_auth) == {}
@pytest.mark.asyncio
async def test_load_user_env_vars_caches_within_ttl(env_vars_salt_key, monkeypatch):
"""A second load within the TTL window is served from the in-memory cache,
keeping the hot tool-call/tool-listing path off the DB."""
from unittest.mock import MagicMock
from litellm.proxy._experimental.mcp_server import mcp_server_manager as mgr_mod
from litellm.proxy._experimental.mcp_server.db import store_user_env_vars
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
mgr_mod._user_env_vars_cache.clear()
blob_prisma = _mock_env_vars_prisma()
await store_user_env_vars(blob_prisma, "alice", "srv-1", {"TOKEN": "t0p"})
row = MagicMock()
row.values_b64 = _captured_values_blob(blob_prisma)
prisma = _mock_env_vars_prisma(row=row)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
manager = MCPServerManager()
server = MCPServer(
server_id="srv-1", name="s", transport="http", url="https://example.com"
)
fake_auth = MagicMock()
fake_auth.user_id = "alice"
first = await manager._load_user_env_vars(server, fake_auth)
second = await manager._load_user_env_vars(server, fake_auth)
assert first == {"TOKEN": "t0p"} == second
assert prisma.db.litellm_mcpuserenvvars.find_unique.await_count == 1
mgr_mod._user_env_vars_cache.clear()
@pytest.mark.asyncio
async def test_load_user_env_vars_invalidation_forces_refetch(
env_vars_salt_key, monkeypatch
):
"""After invalidation (store/clear) the next load reads fresh from the DB
instead of serving the stale cached value."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._experimental.mcp_server import mcp_server_manager as mgr_mod
from litellm.proxy._experimental.mcp_server.db import store_user_env_vars
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
invalidate_user_env_vars_cache,
)
from litellm.types.mcp_server.mcp_server_manager import MCPServer
mgr_mod._user_env_vars_cache.clear()
blob_prisma = _mock_env_vars_prisma()
await store_user_env_vars(blob_prisma, "alice", "srv-1", {"TOKEN": "old"})
old_row = MagicMock()
old_row.values_b64 = _captured_values_blob(blob_prisma)
await store_user_env_vars(blob_prisma, "alice", "srv-1", {"TOKEN": "new"})
new_row = MagicMock()
new_row.values_b64 = _captured_values_blob(blob_prisma)
prisma = _mock_env_vars_prisma()
prisma.db.litellm_mcpuserenvvars.find_unique = AsyncMock(
side_effect=[old_row, new_row]
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
manager = MCPServerManager()
server = MCPServer(
server_id="srv-1", name="s", transport="http", url="https://example.com"
)
fake_auth = MagicMock()
fake_auth.user_id = "alice"
assert await manager._load_user_env_vars(server, fake_auth) == {"TOKEN": "old"}
invalidate_user_env_vars_cache("alice", "srv-1")
assert await manager._load_user_env_vars(server, fake_auth) == {"TOKEN": "new"}
assert prisma.db.litellm_mcpuserenvvars.find_unique.await_count == 2
mgr_mod._user_env_vars_cache.clear()
# ── DB helpers: per-user env vars ─────────────────────────────────────────
_SALT_KEY = "test-salt-key-for-env-vars-tests-1234"