mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-06 16:24:46 +00:00
Temporary MCP OAuth sessions were kept in process-local memory, so on
multi-instance/LB proxy deployments a session created on instance A could
not be found when the follow-up /server/oauth/{server_id}/... request
landed on instance B.
Persist temporary session records to Redis (encrypted with the existing
proxy encryption helpers) as a best-effort L2 cache alongside the current
in-memory L1. Convert get_cached_temporary_mcp_server to async and await
it from the authorize/token/register OAuth endpoints.
Made-with: Cursor
This commit is contained in:
@@ -52,12 +52,17 @@ from litellm.proxy._experimental.mcp_server.utils import (
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
validate_and_normalize_mcp_server_payload as _base_validate_and_normalize_mcp_server_payload,
|
||||
)
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
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(
|
||||
@@ -329,13 +334,115 @@ if MCP_AVAILABLE:
|
||||
)
|
||||
return server
|
||||
|
||||
def get_cached_temporary_mcp_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(
|
||||
server_id: str,
|
||||
) -> Optional[MCPServer]:
|
||||
"""
|
||||
Best-effort read from Redis shared cache. Returns None on miss/errors.
|
||||
|
||||
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"):
|
||||
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",
|
||||
)
|
||||
if decrypted_json 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:
|
||||
return 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
|
||||
|
||||
def _redact_mcp_credentials(
|
||||
@@ -1325,6 +1432,10 @@ if MCP_AVAILABLE:
|
||||
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,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
f"Error caching temporary mcp server: {str(e)}"
|
||||
@@ -1336,10 +1447,10 @@ if MCP_AVAILABLE:
|
||||
|
||||
return _redact_mcp_credentials(temp_record)
|
||||
|
||||
def _get_cached_temporary_mcp_server_or_404(
|
||||
async def _get_cached_temporary_mcp_server_or_404(
|
||||
server_id: str, request: Optional[Request] = None
|
||||
) -> MCPServer:
|
||||
server = get_cached_temporary_mcp_server(server_id)
|
||||
server = await get_cached_temporary_mcp_server(server_id)
|
||||
if server is None:
|
||||
# Fall back to real DB/config server (e.g. for the user-side OAuth flow
|
||||
# which calls these endpoints with a real server_id, not a temp session id).
|
||||
@@ -1378,7 +1489,9 @@ if MCP_AVAILABLE:
|
||||
response_type: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
):
|
||||
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(
|
||||
server_id, request=request
|
||||
)
|
||||
# Use the server's stored client_id when the caller doesn't supply one
|
||||
resolved_client_id = mcp_server.client_id or client_id or ""
|
||||
if not resolved_client_id:
|
||||
@@ -1422,7 +1535,9 @@ if MCP_AVAILABLE:
|
||||
refresh_token: Optional[str] = Form(None),
|
||||
scope: Optional[str] = Form(None),
|
||||
):
|
||||
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(
|
||||
server_id, request=request
|
||||
)
|
||||
resolved_client_id = mcp_server.client_id or client_id or ""
|
||||
if not resolved_client_id:
|
||||
raise HTTPException(
|
||||
@@ -1458,7 +1573,9 @@ if MCP_AVAILABLE:
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id, request=request)
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(
|
||||
server_id, request=request
|
||||
)
|
||||
request_data = await _read_request_body(request=request)
|
||||
data: dict = {**request_data}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
import json
|
||||
from datetime import datetime, timedelta
|
||||
from types import SimpleNamespace
|
||||
from typing import List, Optional
|
||||
@@ -1311,7 +1312,8 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
assert cache["temp-cache"].server is server
|
||||
assert cache["temp-cache"].expires_at > datetime.utcnow()
|
||||
|
||||
def test_get_cached_temporary_mcp_server_prunes_expired_entries(self):
|
||||
@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,
|
||||
@@ -1327,12 +1329,13 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._temporary_mcp_servers",
|
||||
cache,
|
||||
):
|
||||
result = get_cached_temporary_mcp_server("expired")
|
||||
result = await get_cached_temporary_mcp_server("expired")
|
||||
|
||||
assert result is None
|
||||
assert "expired" not in cache
|
||||
|
||||
def test_get_cached_temporary_mcp_server_or_404(self):
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_temporary_mcp_server_or_404(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_get_cached_temporary_mcp_server_or_404,
|
||||
)
|
||||
@@ -1343,17 +1346,17 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
|
||||
return_value=server,
|
||||
) as get_cached:
|
||||
result = _get_cached_temporary_mcp_server_or_404("cached")
|
||||
result = await _get_cached_temporary_mcp_server_or_404("cached")
|
||||
|
||||
assert result is server
|
||||
get_cached.assert_called_once_with("cached")
|
||||
get_cached.assert_awaited_once_with("cached")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
|
||||
return_value=None,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_get_cached_temporary_mcp_server_or_404("missing")
|
||||
await _get_cached_temporary_mcp_server_or_404("missing")
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
@@ -1403,6 +1406,10 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._cache_temporary_mcp_server",
|
||||
MagicMock(),
|
||||
) as cache_mock,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._cache_temporary_mcp_server_in_redis",
|
||||
AsyncMock(),
|
||||
) as redis_cache_mock,
|
||||
):
|
||||
response = await add_session_mcp_server(
|
||||
payload=payload,
|
||||
@@ -1414,6 +1421,9 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
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]
|
||||
@@ -1486,7 +1496,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
)
|
||||
|
||||
assert result is authorize_response
|
||||
get_server.assert_called_once_with("server-1", request=request)
|
||||
get_server.assert_awaited_once_with("server-1", request=request)
|
||||
authorize_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
mcp_server=server,
|
||||
@@ -1533,7 +1543,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
)
|
||||
|
||||
assert result is exchange_response
|
||||
get_server.assert_called_once_with("server-1", request=request)
|
||||
get_server.assert_awaited_once_with("server-1", request=request)
|
||||
exchange_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
mcp_server=server,
|
||||
@@ -1581,7 +1591,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
)
|
||||
|
||||
assert result is exchange_response
|
||||
get_server.assert_called_once_with("server-1", request=request)
|
||||
get_server.assert_awaited_once_with("server-1", request=request)
|
||||
exchange_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
mcp_server=server,
|
||||
@@ -1628,7 +1638,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
result = await mcp_register(request=request, server_id="server-1")
|
||||
|
||||
assert result is register_response
|
||||
get_server.assert_called_once_with("server-1", request=request)
|
||||
get_server.assert_awaited_once_with("server-1", request=request)
|
||||
read_body.assert_awaited_once_with(request=request)
|
||||
register_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
@@ -1640,6 +1650,218 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
fallback_client_id="server-1",
|
||||
)
|
||||
|
||||
@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"""
|
||||
|
||||
Reference in New Issue
Block a user