mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
29035c4a99
commit
11e738a66e
6 changed files with 403 additions and 478 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue