fix(vector-stores): infer Milvus search fields

This commit is contained in:
Yujong Lee 2026-09-01 14:33:44 -07:00
parent ecd63e315a
commit 16b78df82e
2 changed files with 51 additions and 12 deletions

View file

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

View file

@ -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",
[