diff --git a/litellm/llms/milvus/vector_stores/transformation.py b/litellm/llms/milvus/vector_stores/transformation.py index 70d3649debd..ebacbf3377d 100644 --- a/litellm/llms/milvus/vector_stores/transformation.py +++ b/litellm/llms/milvus/vector_stores/transformation.py @@ -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], diff --git a/tests/vector_store_tests/test_milvus_vector_store.py b/tests/vector_store_tests/test_milvus_vector_store.py index f1122d5dfdc..e3c09ac64fd 100644 --- a/tests/vector_store_tests/test_milvus_vector_store.py +++ b/tests/vector_store_tests/test_milvus_vector_store.py @@ -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)"""