fix(milvus): honor REST search result limits

This commit is contained in:
Yujong Lee 2026-09-05 08:28:02 -07:00
parent 6b6d7a1e24
commit 685a24b444
2 changed files with 93 additions and 4 deletions

View file

@ -1,10 +1,12 @@
from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
from typing import TYPE_CHECKING, Annotated, Any, Final
import httpx
from pydantic import Field, TypeAdapter, ValidationError
from litellm.exceptions import BadRequestError
from litellm.llms.base_llm.vector_store.transformation import (
BaseQueryEmbeddingVectorStoreConfig,
VectorStoreEmbeddingExecutor,
@ -39,6 +41,7 @@ MILVUS_OPTIONAL_PARAMS: Final = {
"searchParams",
"consistencyLevel",
}
_RESULT_LIMIT_ADAPTER: Final = TypeAdapter(Annotated[int, Field(ge=1, le=50)])
class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
@ -130,13 +133,14 @@ class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
extra_body: Mapping[str, object] | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
search_options: Final = self._search_options(vector_store_search_optional_params)
query_text: Final = self.query_text(query)
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor)
return self._search_request(
vector_store_id,
query_text,
query_vector,
vector_store_search_optional_params,
search_options,
api_base,
litellm_logging_obj,
litellm_params,
@ -153,24 +157,41 @@ class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
extra_body: Mapping[str, object] | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
search_options: Final = self._search_options(vector_store_search_optional_params)
query_text: Final = self.query_text(query)
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor)
return self._search_request(
vector_store_id,
query_text,
query_vector,
vector_store_search_optional_params,
search_options,
api_base,
litellm_logging_obj,
litellm_params,
)
@staticmethod
def _search_options(optional_params: VectorStoreSearchOptionalRequestParams) -> Mapping[str, object]:
native_options: Final = {key: value for key, value in optional_params.items() if key != "max_num_results"}
max_num_results: Final = optional_params.get("max_num_results")
if max_num_results is None:
return native_options
try:
limit: Final = _RESULT_LIMIT_ADAPTER.validate_python(max_num_results)
except ValidationError as exc:
raise BadRequestError(
message="Milvus max_num_results must be an integer between 1 and 50",
model="milvus",
llm_provider="milvus",
) from exc
return {**native_options, "limit": limit}
@staticmethod
def _search_request(
vector_store_id: str,
query_text: str,
query_vector: Sequence[float],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
vector_store_search_optional_params: Mapping[str, object],
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],

View file

@ -108,6 +108,74 @@ class MockPyMilvusHit(dict[str, object]):
class TestMilvusVectorStore:
"""Test Milvus Vector Store with mocked responses"""
@pytest.mark.parametrize(
("optional_params", "expected_limit"),
[
({}, None),
({"limit": 75}, 75),
({"max_num_results": 1}, 1),
({"max_num_results": 50}, 50),
({"max_num_results": 2, "limit": 7}, 2),
],
)
@pytest.mark.parametrize("async_mode", [False, True])
@pytest.mark.asyncio
async def test_rest_result_limit_contract(
self, optional_params: VectorStoreSearchOptionalRequestParams, expected_limit: int | None, async_mode: bool
) -> None:
executor: Final = MagicMock()
executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
config: Final = MilvusVectorStoreConfig()
kwargs: Final = {
"vector_store_id": "documents",
"query": "limit probe",
"vector_store_search_optional_params": optional_params,
"api_base": "http://milvus:19530",
"litellm_logging_obj": MagicMock(),
"litellm_params": {"litellm_embedding_model": "embedding-alias"},
"embedding_executor": executor,
}
_, body = (
await config.atransform_search_vector_store_request(**kwargs)
if async_mode
else config.transform_search_vector_store_request(**kwargs)
)
assert body.get("limit") == expected_limit
assert "max_num_results" not in body
assert body["collectionName"] == "documents"
@pytest.mark.parametrize("max_num_results", [0, 51])
@pytest.mark.parametrize("async_mode", [False, True])
@pytest.mark.asyncio
async def test_rest_invalid_result_limit_is_rejected_before_embedding(
self, max_num_results: int, async_mode: bool
) -> None:
executor: Final = MagicMock()
executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
config: Final = MilvusVectorStoreConfig()
kwargs: Final = {
"vector_store_id": "documents",
"query": "limit probe",
"vector_store_search_optional_params": {"max_num_results": max_num_results},
"api_base": "http://milvus:19530",
"litellm_logging_obj": MagicMock(),
"litellm_params": {"litellm_embedding_model": "embedding-alias"},
"embedding_executor": executor,
}
with pytest.raises(litellm.BadRequestError) as exc_info:
(
await config.atransform_search_vector_store_request(**kwargs)
if async_mode
else config.transform_search_vector_store_request(**kwargs)
)
assert exc_info.value.status_code == 400
executor.embed.assert_not_called()
executor.aembed.assert_not_called()
@pytest.mark.asyncio
async def test_basic_search_with_mock_async(self):
"""Test basic vector search with mocked backend response (async)"""