diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 422610cef24..9b6725ff802 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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, diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index a229bf592a2..dc9046f7c6e 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py index f17a99a1756..c7865841a08 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_env_vars.py @@ -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"