mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
perf(mcp): cache per-user env var values off the hot tool-call path
This commit is contained in:
parent
21c7c38331
commit
4b31a6d711
3 changed files with 134 additions and 4 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue