From bbe9982069c030c7cc6ab04f0f04ab7f9c5ab35d Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 08:35:53 -0700 Subject: [PATCH] refactor(milvus): share search arguments and preserve managed inputs --- basedpyright-code-budget.json | 18 +-- .../vector_stores/grpc_transformation.py | 112 ++++++------------ .../proxy/vector_store_endpoints/endpoints.py | 30 +++-- .../management_endpoints.py | 30 ++--- litellm/proxy/vector_store_endpoints/utils.py | 12 +- ruff-strict-budget.json | 2 +- test-quality-budget.json | 2 +- .../test_vector_store_endpoints.py | 26 ++-- .../test_vector_store_registry.py | 24 +++- .../test_milvus_vector_store.py | 52 ++++++++ type-discipline-budget.json | 4 +- 11 files changed, 171 insertions(+), 141 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 420bc920eea..b1158c63ef6 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,9 +1,9 @@ { "reportAny": { - "limit": 13429 + "limit": 13428 }, "reportArgumentType": { - "limit": 2178 + "limit": 2168 }, "reportAssignmentType": { "limit": 319 @@ -24,7 +24,7 @@ "limit": 19 }, "reportExplicitAny": { - "limit": 3367 + "limit": 3366 }, "reportFunctionMemberAccess": { "limit": 7 @@ -57,7 +57,7 @@ "limit": 5570 }, "reportMissingTypeArgument": { - "limit": 15270 + "limit": 15264 }, "reportMissingTypeStubs": { "limit": 40 @@ -99,19 +99,19 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 43439 + "limit": 43218 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 38255 + "limit": 38239 }, "reportUnknownParameterType": { - "limit": 19584 + "limit": 19582 }, "reportUnknownVariableType": { - "limit": 29827 + "limit": 29823 }, "reportUnnecessaryCast": { "limit": 110 @@ -123,7 +123,7 @@ "limit": 4 }, "reportUnnecessaryIsInstance": { - "limit": 814 + "limit": 813 }, "reportUntypedBaseClass": { "limit": 0 diff --git a/litellm/llms/milvus/vector_stores/grpc_transformation.py b/litellm/llms/milvus/vector_stores/grpc_transformation.py index 74d9b2b6ba1..a902a013ea7 100644 --- a/litellm/llms/milvus/vector_stores/grpc_transformation.py +++ b/litellm/llms/milvus/vector_stores/grpc_transformation.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Final import httpx from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError -from typing_extensions import Protocol +from typing_extensions import Protocol, ReadOnly, TypedDict import litellm from litellm.llms.base_llm.vector_store.transformation import ( @@ -36,6 +36,21 @@ _MILVUS_ENTITY_ADAPTER: Final = TypeAdapter(Mapping[str, object]) _STRING_KEYS_ADAPTER: Final = TypeAdapter(tuple[str, ...]) +class _MilvusSearchArguments(TypedDict): + collection_name: ReadOnly[str] + data: ReadOnly[list[list[float]]] # mutable-ok: PyMilvus requires nested list search data + anns_field: ReadOnly[str | None] + limit: ReadOnly[int] + filter: ReadOnly[str] + offset: ReadOnly[int | None] + group_by_field: ReadOnly[str | None] + output_fields: ReadOnly[list[str]] # mutable-ok: PyMilvus requires list output fields + search_params: ReadOnly[dict[str, object] | None] # mutable-ok: PyMilvus requires dict search params + consistency_level: ReadOnly[str | None] + partition_names: ReadOnly[list[str] | None] # mutable-ok: PyMilvus requires list partition names + timeout: ReadOnly[float | None] + + class _SyncMilvusClient(Protocol): def search( self, @@ -311,43 +326,14 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig): ) @staticmethod - def _sync_search( - client: _SyncMilvusClient, + def _search_arguments( vector_store_id: str, query_vector: Sequence[float], options: _MilvusSearchOptions, params: _MilvusSearchParams, timeout: float | None, - ) -> object: - return client.search( - collection_name=vector_store_id, - data=[list(query_vector)], # mutable-ok: PyMilvus requires nested list search data - anns_field=options.anns_field, - limit=options.result_limit, - filter=options.filter_expression, - offset=options.offset, - group_by_field=options.grouping_field, - output_fields=options.output_fields_with_text(params.text_field), - search_params=dict(options.search_params) # mutable-ok: PyMilvus requires dict search params - if options.search_params is not None - else None, - consistency_level=options.consistency_level, - partition_names=list(params.milvus_partition_names) # mutable-ok: PyMilvus requires list partition names - if params.milvus_partition_names is not None - else None, - timeout=timeout, - ) - - @staticmethod - async def _async_search( - client: _AsyncMilvusClient, - vector_store_id: str, - query_vector: Sequence[float], - options: _MilvusSearchOptions, - params: _MilvusSearchParams, - timeout: float | None, - ) -> object: - return await client.search( + ) -> _MilvusSearchArguments: + return _MilvusSearchArguments( collection_name=vector_store_id, data=[list(query_vector)], # mutable-ok: PyMilvus requires nested list search data anns_field=options.anns_field, @@ -387,30 +373,18 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig): ) query_vector: Final = _EmbeddingPayload.model_validate(embedding_response).vector() connection_timeout, search_timeout = self._timeouts(timeout) - if self.sync_client is not None: - raw_result: Final = self._sync_search( - self.sync_client, - vector_store_id, - query_vector, - options, - params, - search_timeout, - ) - return self._to_response(raw_result, query_text, params.text_field) - - client: Final = _new_sync_client(params.uri, params.token, params.db_name, connection_timeout) + arguments: Final = self._search_arguments(vector_store_id, query_vector, options, params, search_timeout) + client: Final = ( + self.sync_client + if self.sync_client is not None + else _new_sync_client(params.uri, params.token, params.db_name, connection_timeout) + ) try: - result: Final = self._sync_search( - client, - vector_store_id, - query_vector, - options, - params, - search_timeout, - ) + result: Final = client.search(**arguments) return self._to_response(result, query_text, params.text_field) finally: - client.close() + if self.sync_client is None: + client.close() async def aexecute_search_vector_store_request( self, @@ -433,27 +407,15 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig): ) query_vector: Final = _EmbeddingPayload.model_validate(embedding_response).vector() connection_timeout, search_timeout = self._timeouts(timeout) - if self.async_client is not None: - raw_result: Final = await self._async_search( - self.async_client, - vector_store_id, - query_vector, - options, - params, - search_timeout, - ) - return self._to_response(raw_result, query_text, params.text_field) - - client: Final = _new_async_client(params.uri, params.token, params.db_name, connection_timeout) + arguments: Final = self._search_arguments(vector_store_id, query_vector, options, params, search_timeout) + client: Final = ( + self.async_client + if self.async_client is not None + else _new_async_client(params.uri, params.token, params.db_name, connection_timeout) + ) try: - result: Final = await self._async_search( - client, - vector_store_id, - query_vector, - options, - params, - search_timeout, - ) + result: Final = await client.search(**arguments) return self._to_response(result, query_text, params.text_field) finally: - await client.close() + if self.async_client is None: + await client.close() diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 70055201d68..724322f3dcd 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -73,10 +73,10 @@ def build_request_data_from_managed_vector_store( async def _update_request_data_with_litellm_managed_vector_store_registry( - data: dict, + data: Mapping[str, object], vector_store_id: str, user_api_key_dict: UserAPIKeyAuth | None = None, -) -> dict: +) -> dict[str, object]: """ Update the request data with the litellm managed vector store registry. @@ -92,27 +92,31 @@ async def _update_request_data_with_litellm_managed_vector_store_registry( vector_store_id=vector_store_id ) if vector_store_to_run is None: - data.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None) + caller_data: Final = {key: value for key, value in data.items() if key != MILVUS_ADMIN_CONFIGURED_CONNECTION} if user_api_key_dict is not None: assert_proxy_admin_for_user_supplied_vector_store_connection( - custom_llm_provider=data.get("custom_llm_provider"), - litellm_params=data, + custom_llm_provider=caller_data.get("custom_llm_provider"), + litellm_params=caller_data, user_api_key_dict=user_api_key_dict, ) - return data + return caller_data if user_api_key_dict is not None: await assert_user_can_access_vector_store( vector_store=vector_store_to_run, user_api_key_dict=user_api_key_dict, ) - if normalize_vector_store_provider(vector_store_to_run.get("custom_llm_provider")) == "milvus": - for field in MILVUS_MANAGED_CONFIGURATION_FIELDS: - data.pop(field, None) - data.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None) - data.pop("custom_llm_provider", None) - data.pop("litellm_credential_name", None) + blocked_fields: Final = frozenset( + (MILVUS_ADMIN_CONFIGURED_CONNECTION, "custom_llm_provider", "litellm_credential_name") + ) | ( + MILVUS_MANAGED_CONFIGURATION_FIELDS + if normalize_vector_store_provider(vector_store_to_run.get("custom_llm_provider")) == "milvus" + else frozenset() + ) managed_data: Final = build_request_data_from_managed_vector_store(vector_store_to_run) - request_data: Final = {**data, **managed_data} # mutable-ok: request processing requires a mutable payload + request_data: Final = { + **{key: value for key, value in data.items() if key not in blocked_fields}, + **managed_data, + } if user_api_key_dict is not None: assert_proxy_admin_for_user_supplied_vector_store_connection( custom_llm_provider=request_data.get("custom_llm_provider"), diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 9c43da3a75a..2fe2fbb9009 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -336,17 +336,10 @@ async def new_vector_store( user_id=user_api_key_dict.user_id, ) - # Apply the same litellm_params redaction the list / info / update - # endpoints already use, so a caller-supplied credential or a - # cleartext value persisted by an earlier proxy version doesn't - # come back in the response. - response_vs: Final = LiteLLM_ManagedVectorStore(**new_vector_store) - response_vs["litellm_params"] = _redact_sensitive_litellm_params(new_vector_store.get("litellm_params")) - return { "status": "success", "message": f"Vector store {vector_store.get('vector_store_id')} created successfully", - "vector_store": response_vs, + "vector_store": _redact_vector_store(new_vector_store), } except HTTPException: raise @@ -398,9 +391,7 @@ def _synchronize_vector_store_registry( def _redact_vector_store(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_ManagedVectorStore: - redacted: Final = LiteLLM_ManagedVectorStore(**vector_store) - redacted["litellm_params"] = _redact_sensitive_litellm_params(vector_store.get("litellm_params")) - return redacted + return {**vector_store, "litellm_params": _redact_sensitive_litellm_params(vector_store.get("litellm_params"))} @router.get( @@ -585,10 +576,11 @@ async def get_vector_store_info( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, ) - vector_store_dict: Final = dict(vector_store_typed) - if "litellm_params" in vector_store_dict: - vector_store_dict["litellm_params"] = _redact_sensitive_litellm_params(vector_store_dict["litellm_params"]) - return {"vector_store": vector_store_dict} + return { + "vector_store": _redact_vector_store(vector_store_typed) + if "litellm_params" in vector_store_typed + else dict(vector_store_typed) + } except HTTPException: # Preserve 403/404 from the access-control / not-found checks above; # the catch-all below would otherwise rewrite them as 500. @@ -683,16 +675,10 @@ async def update_vector_store( "Updated vector store %s in both database and in-memory registry", vector_store_id ) - # The DB row is returned in full, so the response would otherwise - # echo the persisted ``litellm_params`` (including provider - # credentials) back to the caller — even when the caller only - # changed unrelated fields like ``vector_store_description``. - response_vs: Final = LiteLLM_ManagedVectorStore(**updated_vs) - response_vs["litellm_params"] = _redact_sensitive_litellm_params(updated_vs.get("litellm_params")) return { "status": "success", "message": f"Vector store {vector_store_id} updated successfully", - "vector_store": response_vs, + "vector_store": _redact_vector_store(updated_vs), } except HTTPException: # Preserve 403/404 responses from the access-control / not-found diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 1eca9b66a68..c0b1ac5e722 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -1,4 +1,3 @@ -import json import re from collections.abc import Iterable, Mapping from types import MappingProxyType @@ -25,6 +24,7 @@ from litellm.types.vector_stores import ( VectorStoreIndexEndpoints, ) from litellm.utils import ProviderConfigManager +from litellm.vector_stores.vector_store_registry import deserialize_litellm_params MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = frozenset( { @@ -48,13 +48,9 @@ def _normalize_litellm_params( ) -> LiteLLM_ManagedVectorStore: litellm_params: Final = vector_store.get("litellm_params") if isinstance(litellm_params, str): - normalized: Final = _MANAGED_VECTOR_STORE_ADAPTER.validate_python(vector_store) - try: - parsed: Final = json.loads(litellm_params) - normalized["litellm_params"] = parsed if isinstance(parsed, dict) else {} - except (TypeError, ValueError): - normalized["litellm_params"] = {} - return normalized + return _MANAGED_VECTOR_STORE_ADAPTER.validate_python( + {**vector_store, "litellm_params": deserialize_litellm_params(litellm_params)} + ) return vector_store diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index fd7b30bc314..89d75dd1a02 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -246,7 +246,7 @@ "limit": 109 }, "TRY300": { - "limit": 852 + "limit": 851 }, "UP028": { "limit": 2 diff --git a/test-quality-budget.json b/test-quality-budget.json index 8d70fd12d5b..9b26a1765ea 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -21,6 +21,6 @@ "limit": 117 }, "TQ008": { - "limit": 10984 + "limit": 10981 } } 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 ff4094f9c83..9820052c37c 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 @@ -1,4 +1,5 @@ import json +from copy import deepcopy from datetime import datetime, timezone from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -726,19 +727,23 @@ async def test_managed_milvus_uses_only_persisted_connection_for_non_admin(): mock_registry = MagicMock() mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store + payload: Final = { + "query": "safe", + "custom_llm_provider": "milvus/probe", + "milvus_transport": "grpc", + "api_base": "http://attacker:19530", + "api_key": "attacker-token", + "litellm_embedding_model": "openai/attacker-model", + "litellm_embedding_config": {"api_base": "http://attacker-embedding"}, + } + original_payload: Final = deepcopy(payload) + original_store: Final = deepcopy(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", - "custom_llm_provider": "milvus/probe", - "milvus_transport": "grpc", - "api_base": "http://attacker:19530", - "api_key": "attacker-token", - "litellm_embedding_model": "openai/attacker-model", - "litellm_embedding_config": {"api_base": "http://attacker-embedding"}, - }, + data=payload, vector_store_id="managed", user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), ) @@ -749,6 +754,9 @@ async def test_managed_milvus_uses_only_persisted_connection_for_non_admin(): assert result["litellm_embedding_model"] == "team-embedding-alias" assert "litellm_embedding_config" not in result + assert payload == original_payload + assert managed_vector_store == original_store + @pytest.mark.asyncio async def test_unmarked_managed_milvus_connection_requires_admin_resave(): diff --git a/tests/test_litellm/vector_stores/test_vector_store_registry.py b/tests/test_litellm/vector_stores/test_vector_store_registry.py index 0534ad466a7..00c8f7da5aa 100644 --- a/tests/test_litellm/vector_stores/test_vector_store_registry.py +++ b/tests/test_litellm/vector_stores/test_vector_store_registry.py @@ -1,5 +1,6 @@ import json from datetime import datetime, timezone +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -7,7 +8,28 @@ import pytest import litellm from litellm.types.vector_stores import LiteLLM_ManagedVectorStore from litellm.vector_stores.main import search -from litellm.vector_stores.vector_store_registry import VectorStoreRegistry +from litellm.vector_stores.vector_store_registry import VectorStoreRegistry, deserialize_litellm_params + + +@pytest.mark.parametrize( + ("value", "expected"), + [('{"milvus_transport":"grpc"}', {"milvus_transport": "grpc"}), + ({"milvus_transport": "rest"}, {"milvus_transport": "rest"}), + ("invalid-json", {}), ("[]", {}), ("null", {}), (None, {})], +) +def test_deserialize_litellm_params(value: object, expected: dict[str, object]) -> None: + assert deserialize_litellm_params(value) == expected + + +@pytest.mark.parametrize("serialized_params", ['{"milvus_transport":"grpc"}', "invalid-json", "[]", "null"]) +def test_normalizing_stored_params_preserves_source(serialized_params: str) -> None: + from litellm.proxy.vector_store_endpoints.utils import _normalize_litellm_params + + store: Final = {"vector_store_id": "documents", "custom_llm_provider": "milvus", "litellm_params": serialized_params} + result: Final = _normalize_litellm_params(store) + + assert result["litellm_params"] == deserialize_litellm_params(serialized_params) + assert store["litellm_params"] == serialized_params @pytest.mark.asyncio diff --git a/tests/vector_store_tests/test_milvus_vector_store.py b/tests/vector_store_tests/test_milvus_vector_store.py index e3c09ac64fd..c4162683d20 100644 --- a/tests/vector_store_tests/test_milvus_vector_store.py +++ b/tests/vector_store_tests/test_milvus_vector_store.py @@ -2,6 +2,7 @@ Tests for Milvus Vector Store """ +import asyncio import json from typing import Final, cast from unittest.mock import AsyncMock, MagicMock, patch @@ -116,6 +117,7 @@ class TestMilvusVectorStore: ({"max_num_results": 1}, 1), ({"max_num_results": 50}, 50), ({"max_num_results": 2, "limit": 7}, 2), + ({"max_num_results": None, "limit": 75}, 75), ], ) @pytest.mark.parametrize("async_mode", [False, True]) @@ -127,6 +129,7 @@ class TestMilvusVectorStore: executor.embed.return_value = MOCK_EMBEDDING_RESPONSE executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE) config: Final = MilvusVectorStoreConfig() + original_params: Final = optional_params.copy() kwargs: Final = { "vector_store_id": "documents", "query": "limit probe", @@ -145,6 +148,7 @@ class TestMilvusVectorStore: assert body.get("limit") == expected_limit assert "max_num_results" not in body assert body["collectionName"] == "documents" + assert optional_params == original_params @pytest.mark.parametrize("max_num_results", [0, 51]) @pytest.mark.parametrize("async_mode", [False, True]) @@ -561,6 +565,7 @@ class TestMilvusVectorStore: "attributes": {"category": "reference"}, } ] + mock_client.close.assert_not_called() @pytest.mark.asyncio async def test_async_grpc_search_infers_vector_field_and_requests_text_by_default(self): @@ -587,6 +592,7 @@ class TestMilvusVectorStore: VectorStoreSearchOptionalRequestParams, { "max_num_results": 2, + "limit": 7, }, ), litellm_logging_obj=MagicMock(), @@ -607,6 +613,52 @@ class TestMilvusVectorStore: assert mock_client.search.await_args.kwargs["anns_field"] is None assert mock_client.search.await_args.kwargs["output_fields"] == ["book_intro_text"] assert response["data"][0]["content"][0]["text"] == "async result" + mock_client.close.assert_not_called() + + @pytest.mark.parametrize("injected", [False, True]) + @pytest.mark.parametrize("async_mode", [False, True]) + @pytest.mark.parametrize("failure", ["search", "response", "cancellation"]) + @pytest.mark.asyncio + async def test_grpc_client_ownership_after_failure( + self, injected: bool, async_mode: bool, failure: str + ) -> None: + client: Final = MagicMock() + executor: Final = MagicMock() + executor.embed.return_value = MOCK_EMBEDDING_RESPONSE + executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE) + error: Final = asyncio.CancelledError if failure == "cancellation" else RuntimeError + search: Final = AsyncMock() if async_mode else MagicMock() + search.side_effect = None if failure == "response" else error("search interrupted") + search.return_value = "invalid result" + client.search = search + client.close = AsyncMock() if async_mode else MagicMock() + config: Final = MilvusGRPCVectorStoreConfig( + sync_client=client if injected and not async_mode else None, + async_client=client if injected and async_mode else None, + ) + kwargs: Final = { + "vector_store_id": "documents", + "query": "cleanup probe", + "vector_store_search_optional_params": {}, + "litellm_logging_obj": MagicMock(), + "litellm_params": { + "api_base": "http://milvus:19530", + "litellm_embedding_model": "embedding-alias", + }, + "embedding_executor": executor, + } + with ( + patch("pymilvus.AsyncMilvusClient" if async_mode else "pymilvus.MilvusClient", return_value=client), + pytest.raises(TypeError if failure == "response" else error), + ): + ( + await config.aexecute_search_vector_store_request(**kwargs) + if async_mode else config.execute_search_vector_store_request(**kwargs) + ) + + assert client.close.call_count == (0 if injected else 1) + if async_mode: + assert client.close.await_count == (0 if injected else 1) def test_grpc_search_always_requests_configured_text_field(self): mock_client = MagicMock() diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 7008a9b1aa3..74bdcf77652 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22165 + "limit": 22156 }, "LIT002": { "limit": 26733 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16442 + "limit": 16431 }, "LIT011": { "limit": 5506