mirror of
https://github.com/tiennm99/litellm.git
synced 2026-07-20 04:18:42 +00:00
Merge pull request #26522 from BerriAI/litellm_yj_apr25
[Infra] Merge dev branch
This commit is contained in:
@@ -21,6 +21,7 @@ class SensitiveDataMasker:
|
||||
"auth",
|
||||
"authorization",
|
||||
"credential",
|
||||
"credentials",
|
||||
"access",
|
||||
"private",
|
||||
"certificate",
|
||||
|
||||
@@ -29,6 +29,7 @@ def remove_sensitive_info_from_deployment(
|
||||
deployment_dict["litellm_params"].pop("api_key", None)
|
||||
deployment_dict["litellm_params"].pop("client_secret", None)
|
||||
deployment_dict["litellm_params"].pop("vertex_credentials", None)
|
||||
deployment_dict["litellm_params"].pop("vertex_ai_credentials", None)
|
||||
deployment_dict["litellm_params"].pop("aws_access_key_id", None)
|
||||
deployment_dict["litellm_params"].pop("aws_secret_access_key", None)
|
||||
|
||||
|
||||
@@ -1448,15 +1448,15 @@ if MCP_AVAILABLE:
|
||||
return _redact_mcp_credentials(temp_record)
|
||||
|
||||
async def _get_cached_temporary_mcp_server_or_404(
|
||||
server_id: str, request: Optional[Request] = None
|
||||
server_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
request: Optional[Request] = None,
|
||||
) -> MCPServer:
|
||||
server = await get_cached_temporary_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
|
||||
# which calls these endpoints with a real server_id, not a temp session id).
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
|
||||
client_ip = IPAddressUtils.get_mcp_client_ip(request) if request else None
|
||||
@@ -1470,6 +1470,28 @@ if MCP_AVAILABLE:
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail={"error": f"MCP server {server_id} not found"},
|
||||
)
|
||||
|
||||
# Per-server access policy mirrors `fetch_mcp_server`: admin-view
|
||||
# callers are unrestricted; non-admins must have the server in their
|
||||
# allowed-servers set. Temporary cached servers come from the
|
||||
# admin-only `/server/oauth/session` setup flow and are not exposed
|
||||
# to non-admins.
|
||||
if not _user_has_admin_view(user_api_key_dict):
|
||||
if resolved_from_temp_cache:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={"error": f"Access denied to MCP server {server_id}"},
|
||||
)
|
||||
allowed_server_ids = (
|
||||
await global_mcp_server_manager.get_allowed_mcp_servers(
|
||||
user_api_key_dict
|
||||
)
|
||||
)
|
||||
if server.server_id not in allowed_server_ids:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail={"error": f"Access denied to MCP server {server_id}"},
|
||||
)
|
||||
return server
|
||||
|
||||
@router.get(
|
||||
@@ -1490,7 +1512,7 @@ if MCP_AVAILABLE:
|
||||
scope: Optional[str] = None,
|
||||
):
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(
|
||||
server_id, request=request
|
||||
server_id, user_api_key_dict, 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 ""
|
||||
@@ -1536,7 +1558,7 @@ if MCP_AVAILABLE:
|
||||
scope: Optional[str] = Form(None),
|
||||
):
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(
|
||||
server_id, request=request
|
||||
server_id, user_api_key_dict, request=request
|
||||
)
|
||||
resolved_client_id = mcp_server.client_id or client_id or ""
|
||||
if not resolved_client_id:
|
||||
@@ -1574,7 +1596,7 @@ if MCP_AVAILABLE:
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
mcp_server = await _get_cached_temporary_mcp_server_or_404(
|
||||
server_id, request=request
|
||||
server_id, user_api_key_dict, request=request
|
||||
)
|
||||
request_data = await _read_request_body(request=request)
|
||||
data: dict = {**request_data}
|
||||
|
||||
@@ -71,7 +71,10 @@ class BaseRAGIngestion(ABC):
|
||||
Load credentials from litellm_credential_name if provided in vector_store config.
|
||||
|
||||
This allows users to specify a credential name in the vector_store config
|
||||
which will be resolved from litellm.credential_list.
|
||||
which will be resolved from litellm.credential_list. When a stored
|
||||
credential is used, its values take precedence over caller-supplied
|
||||
equivalents so the api_key / api_base pair stays consistent with the
|
||||
credential definition.
|
||||
"""
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
|
||||
@@ -80,10 +83,15 @@ class BaseRAGIngestion(ABC):
|
||||
credential_values = CredentialAccessor.get_credential_values(
|
||||
credential_name
|
||||
)
|
||||
# Merge credentials into vector_store_config (don't overwrite existing values)
|
||||
if not credential_values:
|
||||
return
|
||||
for key, value in credential_values.items():
|
||||
if key not in self.vector_store_config:
|
||||
self.vector_store_config[key] = value
|
||||
self.vector_store_config[key] = value
|
||||
if (
|
||||
"api_base" in self.vector_store_config
|
||||
and "api_base" not in credential_values
|
||||
):
|
||||
del self.vector_store_config["api_base"]
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> str:
|
||||
|
||||
@@ -65,12 +65,13 @@ def test_excluded_keys_exact_match():
|
||||
# Non-sensitive keys should remain unchanged
|
||||
assert masked["port"] == 6379
|
||||
|
||||
# Test case sensitivity - excluded_keys should be exact match
|
||||
# Test case sensitivity - excluded_keys should be exact match. Supplying an
|
||||
# uppercase variant must NOT exclude the lowercase key, so the field falls
|
||||
# back to pattern matching ("credentials" is in the sensitive pattern set)
|
||||
# and is masked.
|
||||
masked = masker.mask_dict(data, excluded_keys={"LITELLM_CREDENTIALS_NAME"})
|
||||
# Should still be masked because case doesn't match (exact match required)
|
||||
assert (
|
||||
masked["litellm_credentials_name"] == "my-credential-name"
|
||||
) # Not masked because it doesn't match patterns anyway
|
||||
assert masked["litellm_credentials_name"] != "my-credential-name"
|
||||
assert "*" in masked["litellm_credentials_name"]
|
||||
|
||||
# Test with api_key in excluded_keys to verify it works for keys that would be masked
|
||||
masked = masker.mask_dict(data, excluded_keys={"api_key"})
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import copy
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.common_utils.openai_endpoint_utils import (
|
||||
@@ -86,7 +88,7 @@ def test_remove_sensitive_info_from_deployment_with_excluded_keys():
|
||||
"""
|
||||
Test that excluded_keys prevents masking of specific keys (exact match).
|
||||
"""
|
||||
model_config = {
|
||||
base_config = {
|
||||
"model_name": "test-model",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4",
|
||||
@@ -98,13 +100,13 @@ def test_remove_sensitive_info_from_deployment_with_excluded_keys():
|
||||
}
|
||||
|
||||
# Without excluded_keys, access_token should be masked (contains "token")
|
||||
sanitized_config = remove_sensitive_info_from_deployment(model_config)
|
||||
sanitized_config = remove_sensitive_info_from_deployment(copy.deepcopy(base_config))
|
||||
assert sanitized_config["litellm_params"]["access_token"] != "token-12345"
|
||||
assert "*" in sanitized_config["litellm_params"]["access_token"]
|
||||
|
||||
# With excluded_keys, litellm_credentials_name should NOT be masked (even if it would match patterns)
|
||||
sanitized_config = remove_sensitive_info_from_deployment(
|
||||
model_config, excluded_keys={"litellm_credentials_name"}
|
||||
copy.deepcopy(base_config), excluded_keys={"litellm_credentials_name"}
|
||||
)
|
||||
assert (
|
||||
sanitized_config["litellm_params"]["litellm_credentials_name"]
|
||||
|
||||
@@ -1341,12 +1341,15 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
)
|
||||
|
||||
server = generate_mock_mcp_server_config_record(server_id="cached")
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
|
||||
return_value=server,
|
||||
) as get_cached:
|
||||
result = await _get_cached_temporary_mcp_server_or_404("cached")
|
||||
result = await _get_cached_temporary_mcp_server_or_404("cached", admin_auth)
|
||||
|
||||
assert result is server
|
||||
get_cached.assert_awaited_once_with("cached")
|
||||
@@ -1356,10 +1359,95 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
return_value=None,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _get_cached_temporary_mcp_server_or_404("missing")
|
||||
await _get_cached_temporary_mcp_server_or_404("missing", admin_auth)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_temporary_mcp_server_non_admin_denied(self):
|
||||
"""Non-admin without access to the server gets 403, not the server."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_get_cached_temporary_mcp_server_or_404,
|
||||
)
|
||||
|
||||
registry_server = generate_mock_mcp_server_config_record(server_id="server-x")
|
||||
non_admin = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = registry_server
|
||||
mock_manager.get_mcp_server_by_name.return_value = None
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=[])
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _get_cached_temporary_mcp_server_or_404("server-x", non_admin)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
mock_manager.get_allowed_mcp_servers.assert_awaited_once_with(non_admin)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_temporary_mcp_server_non_admin_allowed(self):
|
||||
"""Non-admin with the server in their allowed set gets the server."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_get_cached_temporary_mcp_server_or_404,
|
||||
)
|
||||
|
||||
registry_server = generate_mock_mcp_server_config_record(server_id="server-x")
|
||||
non_admin = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
mock_manager = MagicMock()
|
||||
mock_manager.get_mcp_server_by_id.return_value = registry_server
|
||||
mock_manager.get_mcp_server_by_name.return_value = None
|
||||
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server-x"])
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager",
|
||||
mock_manager,
|
||||
),
|
||||
):
|
||||
result = await _get_cached_temporary_mcp_server_or_404(
|
||||
"server-x", non_admin
|
||||
)
|
||||
|
||||
assert result is registry_server
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_cached_temporary_mcp_server_temp_cache_non_admin_denied(self):
|
||||
"""Servers resolved from the admin-only temp cache reject non-admins."""
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
_get_cached_temporary_mcp_server_or_404,
|
||||
)
|
||||
|
||||
temp_server = generate_mock_mcp_server_config_record(server_id="temp-cache")
|
||||
non_admin = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_cached_temporary_mcp_server",
|
||||
return_value=temp_server,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _get_cached_temporary_mcp_server_or_404("temp-cache", non_admin)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_session_mcp_server_caches_and_redacts_credentials(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
@@ -1472,6 +1560,9 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
request = MagicMock()
|
||||
server = generate_mock_mcp_server_config_record(server_id="server-1")
|
||||
authorize_response = MagicMock()
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -1486,6 +1577,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
result = await mcp_authorize(
|
||||
request=request,
|
||||
server_id="server-1",
|
||||
user_api_key_dict=admin_auth,
|
||||
client_id="client-id",
|
||||
redirect_uri="https://example.com/callback",
|
||||
state="state123",
|
||||
@@ -1496,7 +1588,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
)
|
||||
|
||||
assert result is authorize_response
|
||||
get_server.assert_awaited_once_with("server-1", request=request)
|
||||
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
|
||||
authorize_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
mcp_server=server,
|
||||
@@ -1518,6 +1610,9 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
request = MagicMock()
|
||||
server = generate_mock_mcp_server_config_record(server_id="server-1")
|
||||
exchange_response = {"access_token": "token"}
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -1532,6 +1627,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
result = await mcp_token(
|
||||
request=request,
|
||||
server_id="server-1",
|
||||
user_api_key_dict=admin_auth,
|
||||
grant_type="authorization_code",
|
||||
code="code-123",
|
||||
redirect_uri="https://example.com/callback",
|
||||
@@ -1543,7 +1639,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
)
|
||||
|
||||
assert result is exchange_response
|
||||
get_server.assert_awaited_once_with("server-1", request=request)
|
||||
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
|
||||
exchange_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
mcp_server=server,
|
||||
@@ -1566,6 +1662,9 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
request = MagicMock()
|
||||
server = generate_mock_mcp_server_config_record(server_id="server-1")
|
||||
exchange_response = {"access_token": "new-token", "refresh_token": "new-rt"}
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -1580,6 +1679,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
result = await mcp_token(
|
||||
request=request,
|
||||
server_id="server-1",
|
||||
user_api_key_dict=admin_auth,
|
||||
grant_type="refresh_token",
|
||||
code=None,
|
||||
redirect_uri=None,
|
||||
@@ -1591,7 +1691,7 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
)
|
||||
|
||||
assert result is exchange_response
|
||||
get_server.assert_awaited_once_with("server-1", request=request)
|
||||
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
|
||||
exchange_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
mcp_server=server,
|
||||
@@ -1620,6 +1720,9 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
"response_types": ["code"],
|
||||
"token_endpoint_auth_method": "client_secret_basic",
|
||||
}
|
||||
admin_auth = generate_mock_user_api_key_auth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -1635,10 +1738,14 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
AsyncMock(return_value=register_response),
|
||||
) as register_mock,
|
||||
):
|
||||
result = await mcp_register(request=request, server_id="server-1")
|
||||
result = await mcp_register(
|
||||
request=request,
|
||||
server_id="server-1",
|
||||
user_api_key_dict=admin_auth,
|
||||
)
|
||||
|
||||
assert result is register_response
|
||||
get_server.assert_awaited_once_with("server-1", request=request)
|
||||
get_server.assert_awaited_once_with("server-1", admin_auth, request=request)
|
||||
read_body.assert_awaited_once_with(request=request)
|
||||
register_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
@@ -1664,12 +1771,15 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
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,
|
||||
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:
|
||||
@@ -1734,7 +1844,9 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
_get_temporary_mcp_server_from_redis,
|
||||
)
|
||||
|
||||
server = generate_mock_mcp_server_config_record(server_id="from-redis-encrypted")
|
||||
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")
|
||||
@@ -1808,7 +1920,9 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
_get_temporary_mcp_server_from_redis,
|
||||
)
|
||||
|
||||
mock_cache_backend = SimpleNamespace(async_get_cache=AsyncMock(return_value="enc"))
|
||||
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:
|
||||
@@ -1823,12 +1937,16 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_temporary_mcp_server_from_redis_returns_none_on_decrypt_none(self):
|
||||
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"))
|
||||
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:
|
||||
|
||||
Reference in New Issue
Block a user