mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vector-stores): infer Milvus search fields
This commit is contained in:
parent
ecd63e315a
commit
16b78df82e
2 changed files with 51 additions and 12 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue