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
This commit is contained in:
mateo-berri 2026-09-05 13:37:21 -07:00
parent 7ffa8eb1d0
commit b87186a717
3 changed files with 41 additions and 7 deletions

View file

@ -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()
)

View file

@ -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},

View file

@ -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 = {