diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index 663c3fac80..8b88ef9482 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -21,6 +21,7 @@ class SensitiveDataMasker: "auth", "authorization", "credential", + "credentials", "access", "private", "certificate", diff --git a/litellm/proxy/common_utils/openai_endpoint_utils.py b/litellm/proxy/common_utils/openai_endpoint_utils.py index 7e5c83500a..c4bfe11aec 100644 --- a/litellm/proxy/common_utils/openai_endpoint_utils.py +++ b/litellm/proxy/common_utils/openai_endpoint_utils.py @@ -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) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index a68c8ca9fa..fca08e591f 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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} diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index 6a4eb89d0f..1317c557d6 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -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: diff --git a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py b/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py index bb5c71ceb6..6808c4821c 100644 --- a/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py +++ b/tests/test_litellm/litellm_core_utils/test_sensitive_data_masker.py @@ -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"}) diff --git a/tests/test_litellm/proxy/common_utils/test_openai_endpoint_utils.py b/tests/test_litellm/proxy/common_utils/test_openai_endpoint_utils.py index f2ca41a2c9..44288b027a 100644 --- a/tests/test_litellm/proxy/common_utils/test_openai_endpoint_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_openai_endpoint_utils.py @@ -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"] 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 442265d3af..821e200290 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 @@ -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: