diff --git a/litellm/llms/milvus/vector_stores/connection.py b/litellm/llms/milvus/vector_stores/connection.py index 3eed8783362..f84e1651d37 100644 --- a/litellm/llms/milvus/vector_stores/connection.py +++ b/litellm/llms/milvus/vector_stores/connection.py @@ -48,12 +48,20 @@ def _targets_milvus(custom_llm_provider: object, litellm_params: object) -> bool def _is_grpc_connection(custom_llm_provider: object, litellm_params: object) -> bool: return ( - isinstance(litellm_params, dict) + isinstance(litellm_params, Mapping) and _targets_milvus(custom_llm_provider, litellm_params) and litellm_params.get("milvus_transport") == "grpc" ) +def _connection_fields(litellm_params: object) -> Mapping[str, object]: + if not isinstance(litellm_params, Mapping): + return MappingProxyType({}) + return MappingProxyType( + {key: value for key, value in litellm_params.items() if key != MILVUS_ADMIN_CONFIGURED_CONNECTION} + ) + + def connection_rejection( custom_llm_provider: object, litellm_params: object, @@ -64,7 +72,7 @@ def connection_rejection( if not _is_grpc_connection(custom_llm_provider, litellm_params): return None if managed: - if isinstance(litellm_params, dict) and litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True: + if isinstance(litellm_params, Mapping) and litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True: return None return MilvusConnectionRejection.ADMIN_SAVE_REQUIRED if is_proxy_admin: @@ -82,44 +90,32 @@ def prepare_connection_for_persistence( litellm_credential_name: object | None = None, existing_litellm_credential_name: object | None = None, litellm_credential_name_supplied: bool = False, -) -> dict[str, object] | MilvusConnectionRejection: # mutable-ok: persistence requires a serializable connection dict - existing: Final = existing_litellm_params if isinstance(existing_litellm_params, dict) else MappingProxyType({}) - supplied: Final = litellm_params if isinstance(litellm_params, dict) else MappingProxyType({}) - effective: Final = { # mutable-ok: the validated connection must be JSON-serializable for database persistence - key: value - for params in (existing, supplied) - for key, value in params.items() - if key != MILVUS_ADMIN_CONFIGURED_CONNECTION - } +) -> Mapping[str, object] | MilvusConnectionRejection: + existing: Final = _connection_fields(existing_litellm_params) + supplied: Final = _connection_fields(litellm_params) + merged: Final = MappingProxyType({**existing, **supplied}) previous_is_grpc: Final = _is_grpc_connection(existing_custom_llm_provider, existing) - effective_is_grpc: Final = _is_grpc_connection(custom_llm_provider, effective) - if not previous_is_grpc and not effective_is_grpc: - return { # mutable-ok: persistence requires an isolated JSON-serializable dict - key: value - for key, value in (supplied if isinstance(litellm_params, dict) else existing).items() - if key != MILVUS_ADMIN_CONFIGURED_CONNECTION - } - + effective_is_grpc: Final = _is_grpc_connection(custom_llm_provider, merged) + effective: Final = ( + merged + if previous_is_grpc or effective_is_grpc + else supplied + if isinstance(litellm_params, Mapping) + else existing + ) is_create: Final = existing_custom_llm_provider is None provider_changed: Final = not is_create and custom_llm_provider != existing_custom_llm_provider - managed_configuration_changed: Final = any( - existing.get(field) != effective.get(field) for field in MILVUS_MANAGED_CONFIGURATION_FIELDS - ) credential_changed: Final = litellm_credential_name_supplied and ( litellm_credential_name != existing_litellm_credential_name ) - missing_marker: Final = effective_is_grpc and existing.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is not True - - if ( - is_create or provider_changed or managed_configuration_changed or credential_changed or missing_marker - ) and not is_proxy_admin: - return MilvusConnectionRejection.ADMIN_REQUIRED - - return ( - {**effective, MILVUS_ADMIN_CONFIGURED_CONNECTION: True} # mutable-ok: persisted JSON carries the server marker - if effective_is_grpc - else effective + connection_changed: Final = not is_create and (provider_changed or credential_changed or effective != existing) + missing_marker: Final = effective_is_grpc and ( + not isinstance(existing_litellm_params, Mapping) + or existing_litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is not True ) + if (connection_changed or missing_marker) and not is_proxy_admin: + return MilvusConnectionRejection.ADMIN_REQUIRED + return MappingProxyType({**effective, MILVUS_ADMIN_CONFIGURED_CONNECTION: True}) if effective_is_grpc else effective def managed_connection_fields(custom_llm_provider: object, litellm_params: object) -> frozenset[str]: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1cd08fe27c0..53f6f205a58 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7763,7 +7763,10 @@ class ProxyConfig: litellm.vector_store_registry = VectorStoreRegistry(vector_stores=vector_stores) else: for vector_store in vector_stores: - litellm.vector_store_registry.add_vector_store_to_registry(vector_store=vector_store) + if (vector_store_id := vector_store.get("vector_store_id")) is not None: + litellm.vector_store_registry.update_vector_store_in_registry( + vector_store_id=vector_store_id, updated_data=vector_store + ) except Exception as e: verbose_proxy_logger.exception( "litellm.proxy.proxy_server.py::ProxyConfig:_init_vector_stores_in_db - %s", e diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 8814f37d556..7e29e6bcb97 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -92,7 +92,7 @@ def prepare_vector_store_connection_for_persistence( litellm_credential_name: object | None = None, existing_litellm_credential_name: object | None = None, litellm_credential_name_supplied: bool = False, -) -> dict[str, object]: # mutable-ok: persistence requires a serializable effective-connection dict +) -> Mapping[str, object]: result: Final = prepare_connection_for_persistence( custom_llm_provider=custom_llm_provider, litellm_params=litellm_params, diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 9393ec0f8e6..23b5c91786a 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -12781,3 +12781,39 @@ async def test_update_general_settings_keeps_yaml_openai_websocket_passthrough() import litellm.proxy.proxy_server as ps assert ps.general_settings["enable_openai_websocket_passthrough"] is False + + +@pytest.mark.asyncio +async def test_init_vector_stores_in_db_refreshes_a_store_already_in_the_registry(monkeypatch): + from litellm.proxy.proxy_server import ProxyConfig + from litellm.types.vector_stores import MILVUS_ADMIN_CONFIGURED_CONNECTION, LiteLLM_ManagedVectorStore + from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + + stale: Final = LiteLLM_ManagedVectorStore( + vector_store_id="managed-milvus", + custom_llm_provider="milvus", + litellm_params={"api_base": "http://old-milvus:19530", "milvus_transport": "grpc"}, + ) + registry: Final = VectorStoreRegistry(vector_stores=[stale]) + monkeypatch.setattr(litellm, "vector_store_registry", registry) + saved_params: Final = { + "api_base": "http://new-milvus:19530", + "milvus_transport": "grpc", + MILVUS_ADMIN_CONFIGURED_CONNECTION: True, + } + prisma_client: Final = MagicMock() + prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock( + return_value=[ + { + "vector_store_id": "managed-milvus", + "custom_llm_provider": "milvus", + "litellm_params": json.dumps(saved_params), + } + ] + ) + + await ProxyConfig()._init_vector_stores_in_db(prisma_client=prisma_client) + + refreshed: Final = registry.get_litellm_managed_vector_store_from_registry(vector_store_id="managed-milvus") + assert refreshed is not None + assert refreshed["litellm_params"] == saved_params diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index b2a8fed0685..1da2af2a79e 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -1078,11 +1078,11 @@ def test_nested_provider_cannot_bypass_milvus_grpc_registration_authorization( assert exc_info.value.status_code == 403 -def test_non_grpc_connection_update_keeps_replacement_semantics(): +def test_non_grpc_connection_update_keeps_replacement_semantics_for_admins(): params = prepare_vector_store_connection_for_persistence( custom_llm_provider="openai", litellm_params={"api_key": "new-key"}, - user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), existing_custom_llm_provider="openai", existing_litellm_params={"api_key": "old-key", "api_base": "https://old.example"}, ) @@ -1094,7 +1094,7 @@ def test_non_grpc_connection_update_drops_forged_admin_marker(): params = prepare_vector_store_connection_for_persistence( custom_llm_provider="openai", litellm_params={"api_key": "new-key", MILVUS_ADMIN_CONFIGURED_CONNECTION: True}, - user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), existing_custom_llm_provider="openai", existing_litellm_params={"api_key": "old-key"}, ) @@ -1102,6 +1102,59 @@ def test_non_grpc_connection_update_drops_forged_admin_marker(): assert params == {"api_key": "new-key"} +@pytest.mark.parametrize( + ("custom_llm_provider", "litellm_params", "credential"), + ( + ("openai", {"api_key": "sk-real", "api_base": "https://attacker.example"}, None), + ("openai", {"api_key": "sk-real"}, None), + ("openai", {"api_key": "sk-real", "api_base": "https://old.example", "organization": "org-other"}, None), + ("bedrock", None, None), + ("openai", None, "someone-elses-credential"), + ), +) +def test_non_admin_cannot_change_a_non_grpc_connection_on_update( + custom_llm_provider: str, litellm_params: dict[str, object] | None, credential: str | None +) -> None: + with pytest.raises(HTTPException) as exc_info: + prepare_vector_store_connection_for_persistence( + custom_llm_provider=custom_llm_provider, + litellm_params=litellm_params, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + existing_custom_llm_provider="openai", + existing_litellm_params={"api_key": "sk-real", "api_base": "https://old.example"}, + litellm_credential_name=credential, + litellm_credential_name_supplied=credential is not None, + ) + + assert exc_info.value.status_code == 403 + assert "Only proxy admins can configure vector store connections" in exc_info.value.detail + + +@pytest.mark.parametrize("litellm_params", (None, {"api_key": "sk-real", "api_base": "https://old.example"})) +def test_non_admin_update_leaving_the_non_grpc_connection_alone_is_allowed( + litellm_params: dict[str, object] | None, +) -> None: + params = prepare_vector_store_connection_for_persistence( + custom_llm_provider="openai", + litellm_params=litellm_params, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + existing_custom_llm_provider="openai", + existing_litellm_params={"api_key": "sk-real", "api_base": "https://old.example"}, + ) + + assert params == {"api_key": "sk-real", "api_base": "https://old.example"} + + +def test_non_admin_can_still_create_a_non_grpc_store(): + params = prepare_vector_store_connection_for_persistence( + custom_llm_provider="openai", + litellm_params={"api_key": "sk-team-key"}, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + ) + + assert params == {"api_key": "sk-team-key"} + + class TestCheckVectorStorePermission: """Test suite for check_vector_store_permission function.""" @@ -3225,6 +3278,48 @@ class TestUpdateVectorStoreAccessControlAndRedaction: assert exc_info.value.status_code == 403 mock_prisma_client.db.litellm_managedvectorstorestable.update.assert_not_called() + @pytest.mark.asyncio + async def test_non_admin_cannot_move_a_non_grpc_store_by_echoing_the_redacted_key(self): + from litellm.constants import REDACTED_BY_LITELM_STRING + from litellm.proxy.vector_store_endpoints.management_endpoints import update_vector_store + from litellm.types.vector_stores import VectorStoreUpdateRequest + + existing_row = MagicMock() + existing_row.model_dump.return_value = { + "vector_store_id": "vs_shared", + "custom_llm_provider": "openai", + "team_id": None, + "litellm_params": {"api_key": "sk-real-openai-key-123", "api_base": "https://old.example"}, + } + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=existing_row) + + with ( + patch( # test-quality-ok: isolates the connection-authorization behavior from the feature entitlement gate + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: injects the endpoint repository boundary with an admin-created shared store + "litellm.proxy.proxy_server.prisma_client", mock_prisma_client + ), + pytest.raises(HTTPException) as exc_info, + ): + await update_vector_store( + data=VectorStoreUpdateRequest( + vector_store_id="vs_shared", + litellm_params={"api_key": REDACTED_BY_LITELM_STRING, "api_base": "https://attacker.example"}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="member", + team_id="team-A", + user_role=LitellmUserRoles.INTERNAL_USER, + ), + ) + + assert exc_info.value.status_code == 403 + assert "Only proxy admins can configure vector store connections" in exc_info.value.detail + mock_prisma_client.db.litellm_managedvectorstorestable.update.assert_not_called() + @pytest.mark.asyncio async def test_update_denied_when_caller_cannot_access_store(self): from unittest.mock import AsyncMock, MagicMock, patch