From 16b78df82ed1d4434b2e2055c6e7d8eb205909cd Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Tue, 1 Sep 2026 14:33:44 -0700 Subject: [PATCH] fix(vector-stores): infer Milvus search fields --- .../vector_stores/grpc_transformation.py | 21 +++++----- .../test_milvus_vector_store.py | 42 ++++++++++++++++++- 2 files changed, 51 insertions(+), 12 deletions(-) diff --git a/litellm/llms/milvus/vector_stores/grpc_transformation.py b/litellm/llms/milvus/vector_stores/grpc_transformation.py index 75206bfb0e6..5a271ebe5f6 100644 --- a/litellm/llms/milvus/vector_stores/grpc_transformation.py +++ b/litellm/llms/milvus/vector_stores/grpc_transformation.py @@ -25,7 +25,6 @@ from .transformation import MILVUS_OPTIONAL_PARAMS if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -DEFAULT_ANNS_FIELD: Final = "book_intro_vector" DEFAULT_LIMIT: Final = 10 DEFAULT_TEXT_FIELD: Final = "text" _EMPTY_EMBEDDING_CONFIG: Final[Mapping[str, object]] = MappingProxyType({}) @@ -41,7 +40,7 @@ class _SyncMilvusClient(Protocol): self, collection_name: str, data: list[list[float]], # mutable-ok: PyMilvus requires nested list search data - anns_field: str, + anns_field: str | None, limit: int, filter: str, offset: int | None, @@ -61,7 +60,7 @@ class _AsyncMilvusClient(Protocol): self, collection_name: str, data: list[list[float]], # mutable-ok: PyMilvus requires nested list search data - anns_field: str, + anns_field: str | None, limit: int, filter: str, offset: int | None, @@ -168,7 +167,7 @@ class _MilvusSearchParams(BaseModel): class _MilvusSearchOptions(BaseModel): model_config = ConfigDict(frozen=True, extra="forbid", populate_by_name=True) - anns_field: str = Field(default=DEFAULT_ANNS_FIELD, alias="annsField") + anns_field: str | None = Field(default=None, alias="annsField") limit: int = Field(default=DEFAULT_LIMIT, ge=1, le=50) max_num_results: int | None = Field(default=None, ge=1, le=50) filters: Mapping[str, object] | None = None @@ -185,6 +184,12 @@ class _MilvusSearchOptions(BaseModel): def result_limit(self) -> int: return self.max_num_results or self.limit + def output_fields_with_text(self, text_field: str) -> list[str]: + output_fields: Final = self.output_fields or () + if "*" in output_fields or text_field in output_fields: + return list(output_fields) + return [*output_fields, text_field] + class _EmbeddingItem(BaseModel): model_config = ConfigDict(frozen=True) @@ -339,9 +344,7 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig): filter=options.filter_expression, offset=options.offset, group_by_field=options.grouping_field, - output_fields=list(options.output_fields) # mutable-ok: PyMilvus requires list output fields - if options.output_fields is not None - else None, + 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, @@ -369,9 +372,7 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig): filter=options.filter_expression, offset=options.offset, group_by_field=options.grouping_field, - output_fields=list(options.output_fields) # mutable-ok: PyMilvus requires list output fields - if options.output_fields is not None - else None, + 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, diff --git a/tests/vector_store_tests/test_milvus_vector_store.py b/tests/vector_store_tests/test_milvus_vector_store.py index 9da643144c3..4ad0ea2b37c 100644 --- a/tests/vector_store_tests/test_milvus_vector_store.py +++ b/tests/vector_store_tests/test_milvus_vector_store.py @@ -493,7 +493,7 @@ class TestMilvusVectorStore: ] @pytest.mark.asyncio - async def test_grpc_search_uses_async_pymilvus_client(self): + async def test_async_grpc_search_infers_vector_field_and_requests_text_by_default(self): mock_client = MagicMock() mock_client.search = AsyncMock( return_value=[ @@ -515,7 +515,6 @@ class TestMilvusVectorStore: vector_store_search_optional_params=cast( VectorStoreSearchOptionalRequestParams, { - "annsField": "book_intro_vector", "max_num_results": 2, }, ), @@ -533,8 +532,47 @@ class TestMilvusVectorStore: {}, ) assert mock_client.search.await_args.kwargs["limit"] == 2 + 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" + def test_grpc_search_always_requests_configured_text_field(self): + mock_client = MagicMock() + mock_client.search.return_value = [ + [ + { + "id": 9, + "distance": 0.87, + "entity": { + "body": "result text", + "category": "reference", + }, + } + ] + ] + config = MilvusGRPCVectorStoreConfig( + sync_client=mock_client, + embedding_fn=MagicMock(return_value=MOCK_EMBEDDING_RESPONSE), + ) + + response = config.execute_search_vector_store_request( + query="what is machine learning?", + vector_store_id="documents", + vector_store_search_optional_params=cast( + VectorStoreSearchOptionalRequestParams, + {"outputFields": ["category"]}, + ), + litellm_logging_obj=MagicMock(), + litellm_params={ + "api_base": "http://localhost:19530", + "litellm_embedding_model": "text-embedding-3-large", + "milvus_text_field": "body", + }, + ) + + assert mock_client.search.call_args.kwargs["output_fields"] == ["category", "body"] + assert response["data"][0]["content"][0]["text"] == "result text" + @pytest.mark.parametrize( "optional_params", [