mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(milvus): honor REST search result limits
This commit is contained in:
parent
6b6d7a1e24
commit
685a24b444
2 changed files with 93 additions and 4 deletions
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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)"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue