From b87186a717604bffc290647f15ced4c9e8f74d80 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 5 Sep 2026 13:37:21 -0700 Subject: [PATCH] fix(vector-stores): block caller connection fields for nested milvus grpc stores A managed store whose top-level custom_llm_provider is not milvus but whose litellm_params carry custom_llm_provider milvus with the grpc transport let a non-admin caller pass connection fields in the search body. The managed connection field check now looks at both the top-level and nested provider --- .../llms/milvus/vector_stores/connection.py | 16 +++++++---- .../proxy/vector_store_endpoints/endpoints.py | 4 ++- .../test_vector_store_endpoints.py | 28 +++++++++++++++++++ 3 files changed, 41 insertions(+), 7 deletions(-) diff --git a/litellm/llms/milvus/vector_stores/connection.py b/litellm/llms/milvus/vector_stores/connection.py index 73d0dbab7fb..1c7750cda1d 100644 --- a/litellm/llms/milvus/vector_stores/connection.py +++ b/litellm/llms/milvus/vector_stores/connection.py @@ -39,13 +39,17 @@ def _normalize_provider(custom_llm_provider: object) -> str | None: return custom_llm_provider.split("/", 1)[0] +def _targets_milvus(custom_llm_provider: object, litellm_params: object) -> bool: + return _normalize_provider(custom_llm_provider) == "milvus" or ( + isinstance(litellm_params, Mapping) + and _normalize_provider(litellm_params.get("custom_llm_provider")) == "milvus" + ) + + def _is_grpc_connection(custom_llm_provider: object, litellm_params: object) -> bool: return ( isinstance(litellm_params, dict) - and ( - _normalize_provider(custom_llm_provider) == "milvus" - or _normalize_provider(litellm_params.get("custom_llm_provider")) == "milvus" - ) + and _targets_milvus(custom_llm_provider, litellm_params) and litellm_params.get("milvus_transport") == "grpc" ) @@ -116,9 +120,9 @@ def prepare_connection_for_persistence( ) -def managed_connection_fields(custom_llm_provider: object) -> frozenset[str]: +def managed_connection_fields(custom_llm_provider: object, litellm_params: object) -> frozenset[str]: return frozenset((MILVUS_ADMIN_CONFIGURED_CONNECTION, "custom_llm_provider", "litellm_credential_name")) | ( - MILVUS_MANAGED_CONFIGURATION_FIELDS if _normalize_provider(custom_llm_provider) == "milvus" else frozenset() + MILVUS_MANAGED_CONFIGURATION_FIELDS if _targets_milvus(custom_llm_provider, litellm_params) else frozenset() ) diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index b5fe38d2af6..f309c028aef 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -103,7 +103,9 @@ async def _update_request_data_with_litellm_managed_vector_store_registry( vector_store=vector_store_to_run, user_api_key_dict=user_api_key_dict, ) - blocked_fields: Final = managed_connection_fields(vector_store_to_run.get("custom_llm_provider")) + blocked_fields: Final = managed_connection_fields( + vector_store_to_run.get("custom_llm_provider"), vector_store_to_run.get("litellm_params") + ) managed_data: Final = build_request_data_from_managed_vector_store(vector_store_to_run) request_data: Final = { **{key: value for key, value in data.items() if key not in blocked_fields}, 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 4b5057ca47d..53d92e4f713 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 @@ -764,6 +764,34 @@ async def test_managed_milvus_uses_only_persisted_connection_for_non_admin(): assert managed_vector_store == original_store +@pytest.mark.asyncio +@pytest.mark.parametrize("blocked_key", ["api_base", "api_key", "milvus_db_name", "milvus_partition_names"]) +async def test_nested_milvus_grpc_store_drops_caller_connection_fields(blocked_key: str) -> None: + managed_vector_store: LiteLLM_ManagedVectorStore = { + "vector_store_id": "nested", + "custom_llm_provider": "openai", + "litellm_params": { + "custom_llm_provider": "milvus", + "milvus_transport": "grpc", + "litellm_embedding_model": "team-embedding-alias", + MILVUS_ADMIN_CONFIGURED_CONNECTION: True, + }, + } + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store + + with patch.object( # test-quality-ok: the helper reads the process-wide registry directly + litellm, "vector_store_registry", mock_registry + ): + result = await _update_request_data_with_litellm_managed_vector_store_registry( + data={"query": "safe", blocked_key: "attacker-choice"}, + vector_store_id="nested", + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + ) + + assert blocked_key not in result + + @pytest.mark.asyncio async def test_unmarked_managed_milvus_connection_requires_admin_resave(): managed_vector_store: LiteLLM_ManagedVectorStore = {