diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 1d62b325dec..56b4e072198 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -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, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index e7b045d66b5..157ebab1d1c 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 3597e75404a..7b1bbd2c355 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1474c15e778..4bcf11a937d 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_draft_mcp_servers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_draft_mcp_servers.py new file mode 100644 index 00000000000..ab5b6791bc9 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_draft_mcp_servers.py @@ -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") diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index e223140b573..679fb784bf7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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"""