Replace Redis+in-memory temp MCP server caching with DB-backed draft servers

- Add 'draft' value to MCPApprovalStatus enum
- Add DB functions: create_draft_mcp_server, get_draft_mcp_server,
  _delete_draft_mcp_server, delete_expired_draft_mcp_servers
- Replace in-memory dict and Redis caching in OAuth session flow with
  DB writes using approval_status='draft'
- Add periodic cleanup job to prune expired draft rows (5-min TTL)
- Remove all Redis/in-memory temporary server caching code
- Update and add tests for the new DB-backed flow

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-07-06 18:49:46 +00:00
parent 29035c4a99
commit 11e738a66e
6 changed files with 403 additions and 478 deletions

View file

@ -581,6 +581,76 @@ async def create_mcp_server(
return new_mcp_server
async def create_draft_mcp_server(
prisma_client: PrismaClient,
data: NewMCPServerRequest,
touched_by: str,
) -> LiteLLM_MCPServerTable:
data.approval_status = MCPApprovalStatus.draft
if data.server_id is None:
data.server_id = str(uuid.uuid4())
await _delete_draft_mcp_server(prisma_client, data.server_id)
return await create_mcp_server(prisma_client, data, touched_by)
async def get_draft_mcp_server(
prisma_client: PrismaClient,
server_id: str,
ttl_seconds: int,
) -> Optional[LiteLLM_MCPServerTable]:
row = await MCPServerRepository(prisma_client).table.find_first(
where={
"server_id": server_id,
"approval_status": MCPApprovalStatus.draft,
}
)
if row is None:
return None
created = row.created_at
if created is not None:
if created.tzinfo is None:
created = created.replace(tzinfo=timezone.utc)
if datetime.now(timezone.utc) - created > timedelta(seconds=ttl_seconds):
return None
_decrypt_env_vars_on_returned_row(row)
return LiteLLM_MCPServerTable(**row.model_dump())
async def _delete_draft_mcp_server(
prisma_client: PrismaClient,
server_id: str,
) -> None:
try:
await MCPServerRepository(prisma_client).table.delete_many(
where={
"server_id": server_id,
"approval_status": MCPApprovalStatus.draft,
}
)
except Exception:
pass
async def delete_expired_draft_mcp_servers(
prisma_client: PrismaClient,
ttl_seconds: int,
) -> int:
cutoff = datetime.now(timezone.utc) - timedelta(seconds=ttl_seconds)
result = await MCPServerRepository(prisma_client).table.delete_many(
where={
"approval_status": MCPApprovalStatus.draft,
"created_at": {"lt": cutoff},
}
)
if result > 0:
verbose_proxy_logger.info("Cleaned up %d expired draft MCP server(s)", result)
return result
async def update_mcp_server(
prisma_client: PrismaClient,
data: UpdateMCPServerRequest,

View file

@ -1216,6 +1216,7 @@ class MCPApprovalStatus(str, enum.Enum):
pending_review = "pending_review"
active = "active"
rejected = "rejected"
draft = "draft"
from litellm.models.mcp_server import ( # noqa: E402

View file

@ -19,8 +19,7 @@ import functools
import importlib
import json
import os
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, List, Literal, Optional, Set
from fastapi import (
@ -72,7 +71,6 @@ router = APIRouter(prefix="/v1/mcp", tags=["mcp"])
MCP_AVAILABLE: bool = True
TEMPORARY_MCP_SERVER_TTL_SECONDS = 300
TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX = "litellm:mcp:temporary_server"
def does_mcp_server_exist(mcp_server_records: Iterable[Any], mcp_server_id: str) -> bool:
@ -113,11 +111,14 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.db import (
approve_mcp_server,
create_draft_mcp_server,
create_mcp_server,
delete_expired_draft_mcp_servers,
delete_mcp_server,
delete_user_credential,
delete_user_env_vars,
get_all_mcp_servers_for_user,
get_draft_mcp_server,
get_mcp_server,
get_mcp_servers,
get_mcp_submissions,
@ -179,11 +180,6 @@ if MCP_AVAILABLE:
from litellm.types.mcp import MCPAuth, MCPCredentials
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@dataclass
class _TemporaryMCPServerEntry:
server: MCPServer
expires_at: datetime
def _validate_mcp_server_name_fields(payload: Any) -> None:
candidates: List[tuple[str, Optional[str]]] = []
@ -322,123 +318,21 @@ if MCP_AVAILABLE:
],
}
_temporary_mcp_servers: Dict[str, _TemporaryMCPServerEntry] = {}
def _prune_expired_temporary_mcp_servers() -> None:
if not _temporary_mcp_servers:
return
now = datetime.utcnow()
expired_ids = [server_id for server_id, entry in _temporary_mcp_servers.items() if entry.expires_at <= now]
for server_id in expired_ids:
_temporary_mcp_servers.pop(server_id, None)
def _cache_temporary_mcp_server(server: MCPServer, ttl_seconds: int) -> MCPServer:
ttl_seconds = max(1, ttl_seconds)
_prune_expired_temporary_mcp_servers()
expires_at = datetime.utcnow() + timedelta(seconds=ttl_seconds)
_temporary_mcp_servers[server.server_id] = _TemporaryMCPServerEntry(
server=server,
expires_at=expires_at,
)
return server
async def _cache_temporary_mcp_server_in_redis(server: MCPServer, ttl_seconds: int) -> None:
"""
Best-effort write-through to Redis so temporary MCP OAuth sessions are
shared across proxy instances. Keep local in-memory cache as fallback.
"""
if litellm.cache is None or not hasattr(litellm.cache, "cache"):
return
cache_backend = getattr(litellm.cache, "cache", None)
if cache_backend is None or not hasattr(cache_backend, "async_set_cache"):
return
payload: Dict[str, Any] = server.model_dump(mode="json")
payload_json = json.dumps(payload)
try:
encrypted_payload = encrypt_value_helper(payload_json)
except Exception as e:
verbose_proxy_logger.debug(f"Failed to encrypt temporary MCP server payload for Redis cache: {str(e)}")
return
if not isinstance(encrypted_payload, str):
verbose_proxy_logger.debug("Encrypted temporary MCP payload is not a string; skipping Redis cache write")
return
try:
await cache_backend.async_set_cache(
key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server.server_id}",
value=encrypted_payload,
ttl=max(1, ttl_seconds),
)
except Exception as e:
verbose_proxy_logger.debug(f"Failed to write temporary MCP server to Redis cache: {str(e)}")
async def _get_temporary_mcp_server_from_redis(
async def _get_draft_mcp_server_as_mcp_server(
server_id: str,
) -> Optional[MCPServer]:
"""
Best-effort read from Redis shared cache. Returns None on miss/errors.
from litellm.proxy.proxy_server import prisma_client as _prisma_client # noqa: PLC0415
Values must be encrypted strings (same contract as _cache_temporary_mcp_server_in_redis);
legacy plaintext dict payloads are rejected.
"""
if litellm.cache is None or not hasattr(litellm.cache, "cache"):
if _prisma_client is None:
return None
cache_backend = getattr(litellm.cache, "cache", None)
if cache_backend is None or not hasattr(cache_backend, "async_get_cache"):
return None
try:
cached_server = await cache_backend.async_get_cache(
key=f"{TEMPORARY_MCP_SERVER_REDIS_KEY_PREFIX}:{server_id}"
)
except Exception as e:
verbose_proxy_logger.debug(f"Failed reading temporary MCP server from Redis cache: {str(e)}")
return None
if not isinstance(cached_server, str):
verbose_proxy_logger.debug(
"Temporary MCP Redis cache value must be an encrypted string; rejecting non-string payload"
)
return None
decrypted_json = decrypt_value_helper(
value=cached_server,
key="temporary_mcp_server",
exception_type="debug",
draft = await get_draft_mcp_server(
_prisma_client, server_id, ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS
)
if decrypted_json is None:
if draft is None:
return None
try:
loaded = json.loads(decrypted_json)
except Exception as e:
verbose_proxy_logger.debug(f"Invalid decrypted temporary MCP payload in Redis cache: {str(e)}")
return None
if not isinstance(loaded, dict):
return None
payload_dict: Dict[str, Any] = loaded
try:
return MCPServer(**payload_dict)
except Exception as e:
verbose_proxy_logger.debug(f"Invalid temporary MCP server payload in Redis cache: {str(e)}")
return None
async def get_cached_temporary_mcp_server(
server_id: str,
) -> Optional[MCPServer]:
_prune_expired_temporary_mcp_servers()
entry = _temporary_mcp_servers.get(server_id)
if entry is None:
redis_server = await _get_temporary_mcp_server_from_redis(server_id)
if redis_server is None:
return None
# Intentionally avoid repopulating local cache from Redis to prevent
# extending effective lifetime beyond the remaining Redis TTL.
return redis_server
return entry.server
return await global_mcp_server_manager.build_mcp_server_from_table(
draft, credentials_are_encrypted=True
)
def _redact_mcp_credentials(
mcp_server: LiteLLM_MCPServerTable,
@ -619,44 +513,6 @@ if MCP_AVAILABLE:
payload_dict["credentials"] = inherited_credentials
return NewMCPServerRequest(**payload_dict)
def _build_temporary_mcp_server_record(
payload: NewMCPServerRequest,
created_by: Optional[str],
) -> LiteLLM_MCPServerTable:
now = datetime.utcnow()
server_id = payload.server_id or str(uuid.uuid4())
server_name = payload.server_name or payload.alias or server_id
return LiteLLM_MCPServerTable(
server_id=server_id,
server_name=server_name,
alias=payload.alias,
description=payload.description,
url=payload.url,
transport=payload.transport,
auth_type=payload.auth_type,
credentials=payload.credentials,
created_at=now,
updated_at=now,
created_by=created_by,
updated_by=created_by,
teams=[],
mcp_access_groups=payload.mcp_access_groups,
allowed_tools=payload.allowed_tools or [],
extra_headers=payload.extra_headers or [],
mcp_info=payload.mcp_info,
static_headers=payload.static_headers,
command=payload.command,
args=payload.args,
env=payload.env,
authorization_url=payload.authorization_url,
token_url=payload.token_url,
registration_url=payload.registration_url,
allow_all_keys=payload.allow_all_keys,
available_on_public_internet=payload.available_on_public_internet,
timeout=payload.timeout,
max_concurrent_requests=payload.max_concurrent_requests,
)
def get_prisma_client_or_throw(message: str):
from litellm.proxy.proxy_server import prisma_client
@ -1343,7 +1199,12 @@ if MCP_AVAILABLE:
)
if payload.server_id is not None:
# fail if the mcp server with id already exists
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415
_delete_draft_mcp_server,
)
await _delete_draft_mcp_server(prisma_client, payload.server_id)
mcp_server = await get_mcp_server(prisma_client, payload.server_id)
if mcp_server is not None:
raise HTTPException(
@ -1391,7 +1252,7 @@ if MCP_AVAILABLE:
@router.post(
"/server/oauth/session",
description="Temporarily cache an MCP server in memory without writing to the database",
description="Persist a draft MCP server in the database for the OAuth session flow",
dependencies=[Depends(user_api_key_auth)],
status_code=status.HTTP_200_OK,
)
@ -1405,16 +1266,14 @@ if MCP_AVAILABLE:
),
):
"""
Cache MCP server info in memory for a short duration (~5 minutes).
This endpoint does not write to the database. If the same server_id is provided
again while the cache entry is active, it will refresh the cached data + TTL.
Persist a draft MCP server row in the database for the duration of the
OAuth authorization flow (~5 minutes). The draft is promoted to a real
server when the user finalizes via POST /v1/mcp/server, or cleaned up
by a periodic job if the user abandons the flow.
"""
# Validate and normalize payload fields (alias/server name rules)
validate_and_normalize_mcp_server_payload(payload)
# Restrict to proxy admins similar to the persistent create endpoint
if LitellmUserRoles.PROXY_ADMIN != user_api_key_dict.user_role:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
@ -1423,34 +1282,22 @@ if MCP_AVAILABLE:
},
)
prisma_client = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy")
created_by = user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME
payload_with_credentials = _inherit_credentials_from_existing_server(payload)
temp_record = _build_temporary_mcp_server_record(
payload_with_credentials,
created_by,
)
try:
temporary_server = await global_mcp_server_manager.build_mcp_server_from_table(
temp_record,
credentials_are_encrypted=False,
)
_cache_temporary_mcp_server(
temporary_server,
ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS,
)
await _cache_temporary_mcp_server_in_redis(
temporary_server,
ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS,
draft_record = await create_draft_mcp_server(
prisma_client, payload_with_credentials, touched_by=created_by
)
except Exception as e:
verbose_proxy_logger.exception(f"Error caching temporary mcp server: {str(e)}")
verbose_proxy_logger.exception(f"Error creating draft mcp server: {str(e)}")
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Error caching temporary mcp server: {str(e)}"},
detail={"error": f"Error creating draft mcp server: {str(e)}"},
)
return _redact_mcp_credentials(temp_record)
return _redact_mcp_credentials(draft_record)
async def _mcp_oauth_user_api_key_auth(request: Request) -> UserAPIKeyAuth:
"""
@ -1557,7 +1404,7 @@ if MCP_AVAILABLE:
user_api_key_dict: UserAPIKeyAuth,
request: Optional[Request] = None,
) -> MCPServer:
server = await get_cached_temporary_mcp_server(server_id)
server = await _get_draft_mcp_server_as_mcp_server(server_id)
resolved_from_temp_cache = server is not None
if server is None:
# Fall back to real DB/config server (e.g. for the user-side OAuth flow

View file

@ -7549,6 +7549,24 @@ class ProxyStartupEvent:
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
### CLEANUP EXPIRED DRAFT MCP SERVERS ###
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415
delete_expired_draft_mcp_servers,
)
from litellm.proxy.management_endpoints.mcp_management_endpoints import ( # noqa: PLC0415
TEMPORARY_MCP_SERVER_TTL_SECONDS,
)
scheduler.add_job(
delete_expired_draft_mcp_servers,
"interval",
seconds=300,
args=[prisma_client, TEMPORARY_MCP_SERVER_TTL_SECONDS],
id="cleanup_draft_mcp_servers_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
### UPDATE SPEND ###
scheduler.add_job(
update_spend,

View file

@ -0,0 +1,267 @@
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._experimental.mcp_server.db import (
create_draft_mcp_server,
delete_expired_draft_mcp_servers,
get_draft_mcp_server,
_delete_draft_mcp_server,
)
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
MCPApprovalStatus,
MCPTransport,
NewMCPServerRequest,
)
def _make_prisma_row(
server_id: str = "draft-1",
approval_status: str = "draft",
created_at: datetime | None = None,
) -> MagicMock:
now = created_at or datetime.now(timezone.utc)
row = MagicMock()
row.server_id = server_id
row.approval_status = approval_status
row.created_at = now
row.updated_at = now
row.env_vars = None
row.model_dump = MagicMock(
return_value={
"server_id": server_id,
"server_name": server_id,
"alias": server_id,
"url": "https://example.com",
"transport": MCPTransport.http,
"approval_status": approval_status,
"created_at": now,
"updated_at": now,
"created_by": "admin",
"updated_by": "admin",
"teams": [],
"mcp_access_groups": [],
"allowed_tools": [],
"extra_headers": [],
"env_vars": None,
"credentials": None,
"auth_type": None,
"description": None,
"mcp_info": None,
"static_headers": None,
"command": None,
"args": [],
"env": {},
"authorization_url": None,
"token_url": None,
"registration_url": None,
"allow_all_keys": False,
"available_on_public_internet": True,
"timeout": None,
"max_concurrent_requests": None,
"is_byok": False,
"byok_description": [],
"byok_api_key_help_url": None,
"source_url": None,
"instructions": None,
"submitted_by": None,
"submitted_at": None,
"delegate_auth_to_upstream": False,
"oauth_passthrough": False,
"oauth2_flow": None,
"spec_path": None,
"tool_name_to_display_name": None,
"tool_name_to_description": None,
},
)
return row
def _mock_repo(find_first_return=None, delete_many_return=0):
mock_table = MagicMock()
mock_table.find_first = AsyncMock(return_value=find_first_return)
mock_table.delete_many = AsyncMock(return_value=delete_many_return)
mock_repo_instance = MagicMock()
mock_repo_instance.table = mock_table
return mock_repo_instance, mock_table
class TestCreateDraftMcpServer:
@pytest.mark.asyncio
async def test_sets_approval_status_to_draft(self):
prisma = MagicMock()
payload = NewMCPServerRequest(
server_id="draft-new",
alias="Draft",
url="https://example.com",
transport=MCPTransport.http,
)
draft_row = _make_prisma_row(server_id="draft-new")
with (
patch(
"litellm.proxy._experimental.mcp_server.db._delete_draft_mcp_server",
AsyncMock(),
) as del_mock,
patch(
"litellm.proxy._experimental.mcp_server.db.create_mcp_server",
AsyncMock(return_value=draft_row),
) as create_mock,
):
result = await create_draft_mcp_server(prisma, payload, touched_by="admin")
assert payload.approval_status == MCPApprovalStatus.draft
del_mock.assert_awaited_once_with(prisma, "draft-new")
create_mock.assert_awaited_once_with(prisma, payload, "admin")
assert result is draft_row
@pytest.mark.asyncio
async def test_generates_server_id_when_none(self):
prisma = MagicMock()
payload = NewMCPServerRequest(
alias="NoId",
url="https://example.com",
transport=MCPTransport.http,
)
assert payload.server_id is None
draft_row = _make_prisma_row()
with (
patch(
"litellm.proxy._experimental.mcp_server.db._delete_draft_mcp_server",
AsyncMock(),
),
patch(
"litellm.proxy._experimental.mcp_server.db.create_mcp_server",
AsyncMock(return_value=draft_row),
),
):
await create_draft_mcp_server(prisma, payload, touched_by="admin")
assert payload.server_id is not None
class TestGetDraftMcpServer:
@pytest.mark.asyncio
async def test_returns_none_when_no_row(self):
repo_instance, mock_table = _mock_repo(find_first_return=None)
prisma = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
):
result = await get_draft_mcp_server(prisma, "missing", ttl_seconds=300)
assert result is None
@pytest.mark.asyncio
async def test_returns_none_when_expired(self):
old_time = datetime.now(timezone.utc) - timedelta(seconds=600)
row = _make_prisma_row(created_at=old_time)
repo_instance, _ = _mock_repo(find_first_return=row)
prisma = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
):
result = await get_draft_mcp_server(prisma, "draft-1", ttl_seconds=300)
assert result is None
@pytest.mark.asyncio
async def test_returns_table_when_valid(self):
recent_time = datetime.now(timezone.utc) - timedelta(seconds=10)
row = _make_prisma_row(created_at=recent_time)
repo_instance, _ = _mock_repo(find_first_return=row)
prisma = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
):
result = await get_draft_mcp_server(prisma, "draft-1", ttl_seconds=300)
assert result is not None
assert isinstance(result, LiteLLM_MCPServerTable)
assert result.server_id == "draft-1"
@pytest.mark.asyncio
async def test_handles_naive_created_at(self):
naive_time = datetime.utcnow() - timedelta(seconds=10)
row = _make_prisma_row(created_at=naive_time)
repo_instance, _ = _mock_repo(find_first_return=row)
prisma = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
):
result = await get_draft_mcp_server(prisma, "draft-1", ttl_seconds=300)
assert result is not None
class TestDeleteExpiredDraftMcpServers:
@pytest.mark.asyncio
async def test_deletes_old_drafts(self):
repo_instance, mock_table = _mock_repo(delete_many_return=3)
prisma = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
):
count = await delete_expired_draft_mcp_servers(prisma, ttl_seconds=300)
assert count == 3
mock_table.delete_many.assert_awaited_once()
where = mock_table.delete_many.call_args.kwargs["where"]
assert where["approval_status"] == MCPApprovalStatus.draft
assert "lt" in where["created_at"]
@pytest.mark.asyncio
async def test_returns_zero_when_none_expired(self):
repo_instance, _ = _mock_repo(delete_many_return=0)
prisma = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
):
count = await delete_expired_draft_mcp_servers(prisma, ttl_seconds=300)
assert count == 0
class TestDeleteDraftMcpServer:
@pytest.mark.asyncio
async def test_deletes_by_server_id_and_status(self):
repo_instance, mock_table = _mock_repo(delete_many_return=1)
prisma = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
):
await _delete_draft_mcp_server(prisma, "target-id")
mock_table.delete_many.assert_awaited_once()
where = mock_table.delete_many.call_args.kwargs["where"]
assert where["server_id"] == "target-id"
assert where["approval_status"] == MCPApprovalStatus.draft
@pytest.mark.asyncio
async def test_swallows_exceptions(self):
repo_instance, mock_table = _mock_repo()
mock_table.delete_many = AsyncMock(side_effect=Exception("db error"))
prisma = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.db.MCPServerRepository",
return_value=repo_instance,
):
await _delete_draft_mcp_server(prisma, "target-id")

View file

@ -1547,45 +1547,6 @@ class TestTemporaryMCPSessionEndpoints:
}
mock_manager.get_mcp_server_by_id.assert_called_once_with("server-123")
def test_cache_temporary_mcp_server_stores_entry_with_ttl(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server,
)
server = generate_mock_mcp_server_config_record(server_id="temp-cache")
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._temporary_mcp_servers",
{},
) as cache:
cached_server = _cache_temporary_mcp_server(server, ttl_seconds=2)
assert cached_server is server
assert "temp-cache" in cache
assert cache["temp-cache"].server is server
assert cache["temp-cache"].expires_at > datetime.utcnow()
@pytest.mark.asyncio
async def test_get_cached_temporary_mcp_server_prunes_expired_entries(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_TemporaryMCPServerEntry,
get_cached_temporary_mcp_server,
)
server = generate_mock_mcp_server_config_record(server_id="expired")
expired_entry = _TemporaryMCPServerEntry(
server=server,
expires_at=datetime.utcnow() - timedelta(seconds=30),
)
cache = {"expired": expired_entry}
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._temporary_mcp_servers",
cache,
):
result = await get_cached_temporary_mcp_server("expired")
assert result is None
assert "expired" not in cache
@pytest.mark.asyncio
async def test_get_cached_temporary_mcp_server_or_404(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
@ -1598,7 +1559,7 @@ class TestTemporaryMCPSessionEndpoints:
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_draft_mcp_server_as_mcp_server",
return_value=server,
) as get_cached:
result = await _get_cached_temporary_mcp_server_or_404("cached", admin_auth)
@ -1607,7 +1568,7 @@ class TestTemporaryMCPSessionEndpoints:
get_cached.assert_awaited_once_with("cached")
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_draft_mcp_server_as_mcp_server",
return_value=None,
):
with pytest.raises(HTTPException) as exc_info:
@ -1633,7 +1594,7 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_draft_mcp_server_as_mcp_server",
return_value=None,
),
patch(
@ -1669,7 +1630,7 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_draft_mcp_server_as_mcp_server",
return_value=None,
),
patch(
@ -1720,7 +1681,7 @@ class TestTemporaryMCPSessionEndpoints:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_draft_mcp_server_as_mcp_server",
return_value=None,
),
patch(
@ -1752,7 +1713,7 @@ class TestTemporaryMCPSessionEndpoints:
)
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_draft_mcp_server_as_mcp_server",
return_value=temp_server,
):
with pytest.raises(HTTPException) as exc_info:
@ -1761,9 +1722,8 @@ class TestTemporaryMCPSessionEndpoints:
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_add_session_mcp_server_caches_and_redacts_credentials(self):
async def test_add_session_mcp_server_writes_draft_to_db_and_redacts(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
TEMPORARY_MCP_SERVER_TTL_SECONDS,
add_session_mcp_server,
)
@ -1788,10 +1748,10 @@ class TestTemporaryMCPSessionEndpoints:
aws_region_name=None,
aws_service_name=None,
)
built_server = generate_mock_mcp_server_config_record(server_id="temp-server")
draft_record = generate_mock_mcp_server_db_record(server_id="temp-server")
mock_manager = MagicMock()
mock_manager.get_mcp_server_by_id.return_value = inherited_server
mock_manager.build_mcp_server_from_table = AsyncMock(return_value=built_server)
mock_prisma = MagicMock()
with (
patch(
@ -1803,13 +1763,13 @@ class TestTemporaryMCPSessionEndpoints:
mock_manager,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._cache_temporary_mcp_server",
MagicMock(),
) as cache_mock,
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
return_value=mock_prisma,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._cache_temporary_mcp_server_in_redis",
AsyncMock(),
) as redis_cache_mock,
"litellm.proxy.management_endpoints.mcp_management_endpoints.create_draft_mcp_server",
AsyncMock(return_value=draft_record),
) as create_draft_mock,
):
response = await add_session_mcp_server(
payload=payload,
@ -1817,22 +1777,7 @@ class TestTemporaryMCPSessionEndpoints:
)
validate_mock.assert_called_once_with(payload)
mock_manager.build_mcp_server_from_table.assert_awaited_once()
cache_mock.assert_called_once_with(
built_server, ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS
)
redis_cache_mock.assert_awaited_once_with(
built_server, ttl_seconds=TEMPORARY_MCP_SERVER_TTL_SECONDS
)
args, _ = mock_manager.build_mcp_server_from_table.call_args
temp_record = args[0]
assert temp_record.credentials == {
"auth_value": "token-abc",
"client_id": "client-id",
"client_secret": "client-secret",
"scopes": ["scope1"],
}
create_draft_mock.assert_awaited_once()
assert response.credentials is None
@pytest.mark.asyncio
@ -2430,229 +2375,6 @@ class TestTemporaryMCPSessionEndpoints:
assert result is register_response
assert register_mock.await_args.kwargs["persist_credentials"] is False
@pytest.mark.asyncio
async def test_get_cached_temporary_mcp_server_falls_back_to_redis(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
get_cached_temporary_mcp_server,
)
server = generate_mock_mcp_server_config_record(server_id="from-redis")
serialized = json.dumps(server.model_dump(mode="json"))
mock_cache_backend = SimpleNamespace(
async_get_cache=AsyncMock(return_value="encrypted-payload")
)
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._temporary_mcp_servers",
{},
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper",
return_value=serialized,
),
):
result = await get_cached_temporary_mcp_server("from-redis")
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is not None
assert result.server_id == "from-redis"
mock_cache_backend.async_get_cache.assert_awaited_once_with(
key="litellm:mcp:temporary_server:from-redis"
)
@pytest.mark.asyncio
async def test_cache_temporary_mcp_server_in_redis_uses_ttl_and_key(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
)
server = generate_mock_mcp_server_config_record(server_id="to-redis")
mock_cache_backend = SimpleNamespace(async_set_cache=AsyncMock())
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.encrypt_value_helper",
return_value="encrypted-payload",
):
await _cache_temporary_mcp_server_in_redis(server, ttl_seconds=123)
finally:
mgmt_endpoints.litellm.cache = original_cache
mock_cache_backend.async_set_cache.assert_awaited_once()
call_kwargs = mock_cache_backend.async_set_cache.await_args.kwargs
assert call_kwargs["key"] == "litellm:mcp:temporary_server:to-redis"
assert call_kwargs["ttl"] == 123
@pytest.mark.asyncio
async def test_cache_temporary_mcp_server_in_redis_encrypts_payload(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
)
server = generate_mock_mcp_server_config_record(server_id="to-redis-encrypted")
mock_cache_backend = SimpleNamespace(async_set_cache=AsyncMock())
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.encrypt_value_helper",
return_value="encrypted-payload",
) as encrypt_mock:
await _cache_temporary_mcp_server_in_redis(server, ttl_seconds=60)
finally:
mgmt_endpoints.litellm.cache = original_cache
encrypt_mock.assert_called_once()
call_kwargs = mock_cache_backend.async_set_cache.await_args.kwargs
assert call_kwargs["value"] == "encrypted-payload"
@pytest.mark.asyncio
async def test_get_temporary_mcp_server_from_redis_decrypts_payload(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_temporary_mcp_server_from_redis,
)
server = generate_mock_mcp_server_config_record(
server_id="from-redis-encrypted"
)
serialized = json.dumps(server.model_dump(mode="json"))
mock_cache_backend = SimpleNamespace(
async_get_cache=AsyncMock(return_value="encrypted-payload")
)
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper",
return_value=serialized,
) as decrypt_mock:
result = await _get_temporary_mcp_server_from_redis(
"from-redis-encrypted"
)
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is not None
assert result.server_id == "from-redis-encrypted"
decrypt_mock.assert_called_once()
@pytest.mark.asyncio
async def test_cache_temporary_mcp_server_in_redis_skips_on_encrypt_failure(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
)
server = generate_mock_mcp_server_config_record(server_id="encrypt-fail")
mock_cache_backend = SimpleNamespace(async_set_cache=AsyncMock())
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.encrypt_value_helper",
side_effect=Exception("boom"),
):
await _cache_temporary_mcp_server_in_redis(server, ttl_seconds=60)
finally:
mgmt_endpoints.litellm.cache = original_cache
mock_cache_backend.async_set_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_cache_temporary_mcp_server_in_redis_skips_non_string_encryption_result(
self,
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_cache_temporary_mcp_server_in_redis,
)
server = generate_mock_mcp_server_config_record(server_id="encrypt-non-string")
mock_cache_backend = SimpleNamespace(async_set_cache=AsyncMock())
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.encrypt_value_helper",
return_value={"not": "a-string"},
):
await _cache_temporary_mcp_server_in_redis(server, ttl_seconds=60)
finally:
mgmt_endpoints.litellm.cache = original_cache
mock_cache_backend.async_set_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_temporary_mcp_server_from_redis_returns_none_on_invalid_decrypt_json(
self,
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_temporary_mcp_server_from_redis,
)
mock_cache_backend = SimpleNamespace(
async_get_cache=AsyncMock(return_value="enc")
)
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper",
return_value="{not json}",
):
result = await _get_temporary_mcp_server_from_redis("bad-json")
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is None
@pytest.mark.asyncio
async def test_get_temporary_mcp_server_from_redis_returns_none_on_decrypt_none(
self,
):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_temporary_mcp_server_from_redis,
)
mock_cache_backend = SimpleNamespace(
async_get_cache=AsyncMock(return_value="enc")
)
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
with patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.decrypt_value_helper",
return_value=None,
):
result = await _get_temporary_mcp_server_from_redis("decrypt-none")
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is None
@pytest.mark.asyncio
async def test_get_temporary_mcp_server_from_redis_rejects_plain_dict_payload(self):
"""Plain dict values in Redis are not accepted (write path is encrypted-only)."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
_get_temporary_mcp_server_from_redis,
)
server = generate_mock_mcp_server_config_record(server_id="legacy-dict")
mock_cache_backend = SimpleNamespace(
async_get_cache=AsyncMock(return_value=server.model_dump(mode="json"))
)
original_cache = mgmt_endpoints.litellm.cache
mgmt_endpoints.litellm.cache = SimpleNamespace(cache=mock_cache_backend)
try:
result = await _get_temporary_mcp_server_from_redis("legacy-dict")
finally:
mgmt_endpoints.litellm.cache = original_cache
assert result is None
class TestUpdateMCPServer:
"""Test suite for update MCP server functionality"""