From 79c32ad0a7bcbfed20be02b5759cdbd838f63764 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 4 Sep 2026 16:29:56 -0700 Subject: [PATCH] fix(vector-stores): close managed gRPC trust gaps --- .../vector_store_pre_call_hook.py | 12 ++ litellm/proxy/_lazy_openapi_snapshot.json | 23 +++ .../management_endpoints.py | 5 + litellm/proxy/vector_store_endpoints/utils.py | 74 ++++---- litellm/types/vector_stores.py | 2 + .../test_vector_store_pre_call_hook.py | 32 ++++ .../test_vector_store_endpoints.py | 159 ++++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 + 8 files changed, 270 insertions(+), 43 deletions(-) diff --git a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py index 12ff38ce4ba..e789c1d800a 100644 --- a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py +++ b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py @@ -140,6 +140,18 @@ class VectorStorePreCallHook(CustomLogger): request_metadata = ( request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {} ) + if llm_router is not None or prisma_client is not None: + from litellm.proxy.vector_store_endpoints.utils import ( + assert_proxy_admin_for_user_supplied_vector_store_connection, + ) + + assert_proxy_admin_for_user_supplied_vector_store_connection( + custom_llm_provider=litellm_params_for_vector_store.get( + "custom_llm_provider", custom_llm_provider + ), + litellm_params=litellm_params_for_vector_store, + managed=True, + ) if llm_router is not None: search_function = cast( # cast-ok: normalize router search callable Callable[..., Awaitable[VectorStoreSearchResponse]], diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3599b8c8ada..7511b90281b 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -38485,6 +38485,29 @@ ], "title": "Custom Llm Provider" }, + "litellm_credential_name": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Litellm Credential Name" + }, + "litellm_params": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Litellm Params" + }, "vector_store_description": { "anyOf": [ { diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index aaec82b3ea6..da85309530b 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -491,6 +491,8 @@ async def delete_vector_store( raise HTTPException(status_code=500, detail="Database not connected") try: + _reject_config_vector_store_id(data.vector_store_id) + # Check if vector store exists in database or in-memory registry db_vector_store_exists = False memory_vector_store_exists = False @@ -662,6 +664,9 @@ async def update_vector_store( user_api_key_dict=user_api_key_dict, existing_custom_llm_provider=existing_vector_store.get("custom_llm_provider"), existing_litellm_params=existing_litellm_params, + litellm_credential_name=update_data.get("litellm_credential_name"), + existing_litellm_credential_name=existing_vector_store.get("litellm_credential_name"), + litellm_credential_name_supplied="litellm_credential_name" in update_data, ) # Handle metadata serialization diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index a0b74435d57..704fdd901b3 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -24,17 +24,13 @@ from litellm.types.vector_stores import ( ) from litellm.utils import ProviderConfigManager -MILVUS_GRPC_CONNECTION_FIELDS: Final = frozenset( +MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = frozenset( { "api_base", "api_key", "milvus_transport", "milvus_db_name", "milvus_partition_names", - } -) -MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = MILVUS_GRPC_CONNECTION_FIELDS | frozenset( - { "litellm_embedding_config", "litellm_embedding_model", "milvus_text_field", @@ -76,18 +72,6 @@ def normalize_vector_store_provider(custom_llm_provider: object) -> str | None: return custom_llm_provider.split("/", 1)[0] -def strip_client_milvus_trust_marker( - litellm_params: object, -) -> dict[str, Any]: # mutable-ok: caller input is copied before removing server-owned state - sanitized: Final = ( - dict(litellm_params) # mutable-ok: authorization requires an isolated mutable copy - if isinstance(litellm_params, dict) - else {} # mutable-ok: absent parameters normalize to an empty mutable mapping - ) - sanitized.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None) - return sanitized - - def is_milvus_grpc_connection(custom_llm_provider: object, litellm_params: object) -> bool: return ( normalize_vector_store_provider(custom_llm_provider) == "milvus" @@ -113,7 +97,7 @@ def assert_proxy_admin_for_vector_store_index_management( def assert_proxy_admin_for_user_supplied_vector_store_connection( custom_llm_provider: object, litellm_params: object, - user_api_key_dict: UserAPIKeyAuth, + user_api_key_dict: UserAPIKeyAuth | None = None, *, managed: bool = False, ) -> None: @@ -126,7 +110,7 @@ def assert_proxy_admin_for_user_supplied_vector_store_connection( status_code=403, detail="This managed Milvus gRPC connection must be re-saved by a proxy admin before it can be used.", ) - if _is_proxy_admin(user_api_key_dict): + if user_api_key_dict is not None and _is_proxy_admin(user_api_key_dict): return raise HTTPException( status_code=403, @@ -141,41 +125,45 @@ def prepare_milvus_connection_for_persistence( user_api_key_dict: UserAPIKeyAuth, existing_custom_llm_provider: object | None = None, existing_litellm_params: object | None = None, + litellm_credential_name: object | None = None, + existing_litellm_credential_name: object | None = None, + litellm_credential_name_supplied: bool = False, ) -> dict[str, Any]: # mutable-ok: persistence requires a serializable effective-connection dict - supplied: Final = strip_client_milvus_trust_marker(litellm_params) - existing: Final = ( - dict(existing_litellm_params) # mutable-ok: authorization compares an isolated persisted-connection copy - if isinstance(existing_litellm_params, dict) - else {} # mutable-ok: absent persisted parameters normalize to an empty mutable mapping - ) - effective: Final = { # mutable-ok: the server marker is applied to the persisted effective connection - **existing, - **supplied, + 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 } previous_is_grpc: Final = is_milvus_grpc_connection(existing_custom_llm_provider, existing) effective_is_grpc: Final = is_milvus_grpc_connection(custom_llm_provider, effective) is_create: Final = existing_custom_llm_provider is None provider_changed: Final = not is_create and custom_llm_provider != existing_custom_llm_provider - connection_changed: Final = any( - existing.get(field) != effective.get(field) for field in MILVUS_GRPC_CONNECTION_FIELDS + 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 (previous_is_grpc or effective_is_grpc) and ( - is_create or provider_changed or connection_changed or missing_marker + if ( + (previous_is_grpc or effective_is_grpc) + and (is_create or provider_changed or managed_configuration_changed or credential_changed or missing_marker) + and not _is_proxy_admin(user_api_key_dict) ): - if not _is_proxy_admin(user_api_key_dict): - raise HTTPException( - status_code=403, - detail="Only proxy admins can configure vector store connections. Contact your LiteLLM administrator.", - ) + raise HTTPException( + status_code=403, + detail="Only proxy admins can configure vector store connections. Contact your LiteLLM administrator.", + ) - if effective_is_grpc: - if _is_proxy_admin(user_api_key_dict) or existing.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True: - effective[MILVUS_ADMIN_CONFIGURED_CONNECTION] = True - else: - effective.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None) - return effective + return ( + {**effective, MILVUS_ADMIN_CONFIGURED_CONNECTION: True} # mutable-ok: persisted JSON carries the server marker + if effective_is_grpc + else effective + ) def _suffix_after_index_name(request_path: str, index_name: str) -> str | None: diff --git a/litellm/types/vector_stores.py b/litellm/types/vector_stores.py index 8c2101d91d7..deba5b72302 100644 --- a/litellm/types/vector_stores.py +++ b/litellm/types/vector_stores.py @@ -62,6 +62,8 @@ class VectorStoreUpdateRequest(BaseModel): vector_store_name: str | None = None vector_store_description: str | None = None vector_store_metadata: dict | None = None + litellm_credential_name: str | None = None + litellm_params: dict[str, object] | None = None class VectorStoreDeleteRequest(BaseModel): diff --git a/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py b/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py index ae5cffd8ab0..f150df7254e 100644 --- a/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py +++ b/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py @@ -149,6 +149,38 @@ async def test_hook_searches_through_the_injected_router_with_the_request_metada assert messages[0]["content"] == "Context:\n\ncontext from vs-router\n\n" +@pytest.mark.asyncio +async def test_hook_does_not_search_an_untrusted_managed_milvus_grpc_connection( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + litellm, + "vector_store_registry", + VectorStoreRegistry( + vector_stores=[ + LiteLLM_ManagedVectorStore( + vector_store_id="legacy", + custom_llm_provider="milvus", + litellm_params={ + "milvus_transport": "grpc", + "api_base": "http://internal-milvus:19530", + }, + ) + ], + ), + ) + router = RecordingRouter() + + _, messages, _ = await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), + ["legacy"], + FakeLoggingObj({}), + ) + + assert router.calls == [] + assert messages == [{"role": "user", "content": "what is litellm?"}] + + @pytest.mark.asyncio async def test_hook_falls_back_to_the_sdk_when_the_runtime_has_no_router( registry_with: RegisterStores, 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 72c53c9e7b9..9a1097d57c9 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 @@ -943,6 +943,45 @@ async def test_config_vector_store_id_cannot_be_updated_in_database(): prisma_client.db.litellm_managedvectorstorestable.update.assert_not_called() +@pytest.mark.asyncio +async def test_config_vector_store_id_cannot_be_deleted(): + from litellm.proxy.vector_store_endpoints.management_endpoints import delete_vector_store + from litellm.types.vector_stores import VectorStoreDeleteRequest + + registry = VectorStoreRegistry( + vector_stores=[ + LiteLLM_ManagedVectorStore( + vector_store_id="configured", + custom_llm_provider="milvus", + ) + ] + ) + registry.config_vector_store_ids = frozenset(("configured",)) + prisma_client = MagicMock() + prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=None) + + with ( + patch.object( # test-quality-ok: the endpoint reads the process-wide registry directly + litellm, "vector_store_registry", registry + ), + patch( # test-quality-ok: the endpoint reads the proxy database singleton directly + "litellm.proxy.proxy_server.prisma_client", prisma_client + ), + patch( # test-quality-ok: feature entitlement is outside config ownership behavior + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new=AsyncMock(), + ), + pytest.raises(HTTPException, match="defined in proxy configuration") as exc_info, + ): + await delete_vector_store( + data=VectorStoreDeleteRequest(vector_store_id="configured"), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert exc_info.value.status_code == 400 + assert registry.get_litellm_managed_vector_store_from_registry("configured") is not None + + def test_config_vector_store_cannot_be_replaced_or_deleted_from_registry(): configured = LiteLLM_ManagedVectorStore( vector_store_id="configured", @@ -2890,6 +2929,126 @@ class TestUpdateVectorStoreAccessControlAndRedaction: credentials to the caller. Both are fixed at the endpoint level. """ + @pytest.mark.asyncio + async def test_proxy_admin_can_migrate_existing_store_to_milvus_grpc(self): + import json + + 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_owned", + "custom_llm_provider": "milvus", + "litellm_params": { + "api_base": "http://milvus-rest:9091", + "litellm_embedding_model": "embedding-alias", + }, + } + updated_row = MagicMock() + updated_row.model_dump.return_value = { + "vector_store_id": "vs_owned", + "custom_llm_provider": "milvus", + "litellm_credential_name": "milvus-credential", + "litellm_params": { + "api_base": "http://milvus:19530", + "milvus_transport": "grpc", + "litellm_embedding_model": "embedding-alias", + MILVUS_ADMIN_CONFIGURED_CONNECTION: True, + }, + } + prisma_client = MagicMock() + prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=existing_row) + prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock(return_value=updated_row) + + with ( + patch( # test-quality-ok: feature entitlement is outside update persistence behavior + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new=AsyncMock(), + ), + patch( # test-quality-ok: the endpoint reads the proxy database singleton directly + "litellm.proxy.proxy_server.prisma_client", prisma_client + ), + patch.object( # test-quality-ok: registry synchronization is outside persistence behavior + litellm, "vector_store_registry", None + ), + ): + await update_vector_store( + data=VectorStoreUpdateRequest( + vector_store_id="vs_owned", + litellm_credential_name="milvus-credential", + litellm_params={ + "api_base": "http://milvus:19530", + "milvus_transport": "grpc", + }, + ), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + update_data = prisma_client.db.litellm_managedvectorstorestable.update.await_args.kwargs["data"] + persisted_params = json.loads(update_data["litellm_params"]) + assert persisted_params["api_base"] == "http://milvus:19530" + assert persisted_params["milvus_transport"] == "grpc" + assert persisted_params["litellm_embedding_model"] == "embedding-alias" + assert persisted_params[MILVUS_ADMIN_CONFIGURED_CONNECTION] is True + assert update_data["litellm_credential_name"] == "milvus-credential" + + @pytest.mark.parametrize( + "update", + [ + {"litellm_params": {"litellm_embedding_config": {"api_base": "http://attacker-embedding"}}}, + {"litellm_credential_name": "attacker-credential"}, + ], + ) + @pytest.mark.asyncio + async def test_non_admin_cannot_replace_managed_grpc_execution_configuration( + self, update: dict[str, object] + ): + 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_owned", + "custom_llm_provider": "milvus", + "team_id": "team-A", + "litellm_credential_name": "trusted-credential", + "litellm_params": { + "milvus_transport": "grpc", + "api_base": "http://trusted-milvus:19530", + "litellm_embedding_config": {"api_base": "http://trusted-embedding"}, + MILVUS_ADMIN_CONFIGURED_CONNECTION: True, + }, + } + prisma_client = MagicMock() + prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=existing_row) + prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock(return_value=existing_row) + + with ( + patch( # test-quality-ok: feature entitlement is outside connection authorization behavior + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new=AsyncMock(), + ), + patch( # test-quality-ok: the endpoint reads the proxy database singleton directly + "litellm.proxy.proxy_server.prisma_client", prisma_client + ), + pytest.raises(HTTPException) as exc_info, + ): + await update_vector_store( + data=VectorStoreUpdateRequest( + vector_store_id="vs_owned", + **update, + ), + user_api_key_dict=UserAPIKeyAuth( + user_id="owner", + team_id="team-A", + user_role=LitellmUserRoles.INTERNAL_USER, + ), + ) + + assert exc_info.value.status_code == 403 + prisma_client.db.litellm_managedvectorstorestable.update.assert_not_called() + @pytest.mark.asyncio async def test_non_admin_cannot_activate_nested_milvus_grpc_connection(self): from litellm.proxy.vector_store_endpoints.management_endpoints import update_vector_store diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index d2693acc42b..e0184aab8ba 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -39126,6 +39126,12 @@ export interface components { VectorStoreUpdateRequest: { /** Custom Llm Provider */ custom_llm_provider?: string | null; + /** Litellm Credential Name */ + litellm_credential_name?: string | null; + /** Litellm Params */ + litellm_params?: { + [key: string]: unknown; + } | null; /** Vector Store Description */ vector_store_description?: string | null; /** Vector Store Id */