From 6b6d7a1e24a62a4837d64e2255ae6af3bd86e6b6 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 08:11:23 -0700 Subject: [PATCH] fix(vector-stores): guard nested Milvus settings and validate searches --- .../vector_stores/grpc_transformation.py | 22 ++++++--- litellm/proxy/vector_store_endpoints/utils.py | 9 +++- .../test_vector_store_endpoints.py | 30 +++++++++++- .../test_milvus_vector_store.py | 46 ++++++++++++++++++- 4 files changed, 95 insertions(+), 12 deletions(-) diff --git a/litellm/llms/milvus/vector_stores/grpc_transformation.py b/litellm/llms/milvus/vector_stores/grpc_transformation.py index 593c9693589..74d9b2b6ba1 100644 --- a/litellm/llms/milvus/vector_stores/grpc_transformation.py +++ b/litellm/llms/milvus/vector_stores/grpc_transformation.py @@ -217,15 +217,25 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig): model="milvus", llm_provider="milvus", ) - return _MilvusSearchOptions.model_validate(optional_params) + try: + return _MilvusSearchOptions.model_validate(optional_params) + except ValidationError as exc: + raise litellm.BadRequestError( + message=f"Invalid Milvus gRPC search options: {exc}", + model="milvus", + llm_provider="milvus", + ) from exc @staticmethod def _query_text(query: str | Sequence[str]) -> str: - if isinstance(query, str): - return query - if not query: - raise ValueError("query must not be empty") - return " ".join(query) + query_text: Final = query if isinstance(query, str) else " ".join(query) + if not query_text.strip(): + raise litellm.BadRequestError( + message="query must not be empty", + model="milvus", + llm_provider="milvus", + ) + return query_text @staticmethod def _timeouts(timeout: float | httpx.Timeout | None) -> tuple[float | None, float | None]: diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index f1cd880bf57..1eca9b66a68 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -30,6 +30,8 @@ MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = frozenset( { "api_base", "api_key", + "custom_llm_provider", + "litellm_credential_name", "milvus_transport", "milvus_db_name", "milvus_partition_names", @@ -77,8 +79,11 @@ def normalize_vector_store_provider(custom_llm_provider: object) -> str | None: def is_milvus_grpc_connection(custom_llm_provider: object, litellm_params: object) -> bool: return ( - normalize_vector_store_provider(custom_llm_provider) == "milvus" - and isinstance(litellm_params, dict) + isinstance(litellm_params, dict) + and ( + normalize_vector_store_provider(custom_llm_provider) == "milvus" + or normalize_vector_store_provider(litellm_params.get("custom_llm_provider")) == "milvus" + ) and litellm_params.get("milvus_transport") == "grpc" ) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index d2d1b8acf6c..1bd267db510 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -1017,6 +1017,23 @@ def test_admin_persistence_strips_forged_marker_and_adds_server_marker(): assert params[MILVUS_ADMIN_CONFIGURED_CONNECTION] is True +@pytest.mark.parametrize( + ("provider", "nested_provider"), + (("milvus", "openai"), ("openai", "milvus")), +) +def test_nested_provider_cannot_bypass_milvus_grpc_registration_authorization( + provider: str, nested_provider: str +) -> None: + with pytest.raises(HTTPException) as exc_info: + prepare_milvus_connection_for_persistence( + custom_llm_provider=provider, + litellm_params={"custom_llm_provider": nested_provider, "milvus_transport": "grpc"}, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + ) + + assert exc_info.value.status_code == 403 + + def test_non_grpc_connection_update_keeps_replacement_semantics(): params = prepare_milvus_connection_for_persistence( custom_llm_provider="openai", @@ -3044,6 +3061,8 @@ class TestUpdateVectorStoreAccessControlAndRedaction: [ {"litellm_params": {"litellm_embedding_config": {"api_base": "http://attacker-embedding"}}}, {"litellm_credential_name": "attacker-credential"}, + {"litellm_params": {"custom_llm_provider": "openai"}}, + {"litellm_params": {"litellm_credential_name": "attacker-credential"}}, ], ) @pytest.mark.asyncio @@ -3095,8 +3114,15 @@ class TestUpdateVectorStoreAccessControlAndRedaction: assert exc_info.value.status_code == 403 prisma_client.db.litellm_managedvectorstorestable.update.assert_not_called() + @pytest.mark.parametrize( + "update", + ( + {"custom_llm_provider": "milvus"}, + {"litellm_params": {"custom_llm_provider": "milvus"}}, + ), + ) @pytest.mark.asyncio - async def test_non_admin_cannot_activate_nested_milvus_grpc_connection(self): + async def test_non_admin_cannot_activate_nested_milvus_grpc_connection(self, update: dict[str, object]): from litellm.proxy.vector_store_endpoints.management_endpoints import update_vector_store from litellm.types.vector_stores import VectorStoreUpdateRequest @@ -3131,7 +3157,7 @@ class TestUpdateVectorStoreAccessControlAndRedaction: await update_vector_store( data=VectorStoreUpdateRequest( vector_store_id="vs_owned", - custom_llm_provider="milvus", + **update, ), user_api_key_dict=UserAPIKeyAuth( user_id="owner", diff --git a/tests/vector_store_tests/test_milvus_vector_store.py b/tests/vector_store_tests/test_milvus_vector_store.py index 0fc67a16b53..f1122d5dfdc 100644 --- a/tests/vector_store_tests/test_milvus_vector_store.py +++ b/tests/vector_store_tests/test_milvus_vector_store.py @@ -3,7 +3,7 @@ Tests for Milvus Vector Store """ import json -from typing import cast +from typing import Final, cast from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -592,7 +592,7 @@ class TestMilvusVectorStore: embedding_executor.embed.return_value = MOCK_EMBEDDING_RESPONSE config = MilvusGRPCVectorStoreConfig(sync_client=mock_client) - with pytest.raises(ValueError, match=r"Input should be (greater|less) than or equal"): + with pytest.raises(litellm.BadRequestError, match=r"Input should be (greater|less) than or equal") as exc_info: config.execute_search_vector_store_request( query="what is machine learning?", vector_store_id="book_2", @@ -605,9 +605,51 @@ class TestMilvusVectorStore: embedding_executor=embedding_executor, ) + assert exc_info.value.status_code == 400 embedding_executor.embed.assert_not_called() mock_client.search.assert_not_called() + @pytest.mark.parametrize("query", ([], "", " ")) + def test_grpc_search_rejects_empty_query(self, query: str | list[str]) -> None: + client: Final = MagicMock() + embedding_executor: Final = MagicMock() + config: Final = MilvusGRPCVectorStoreConfig(sync_client=client) + + with pytest.raises(litellm.BadRequestError, match="query must not be empty") as exc_info: + config.execute_search_vector_store_request( + query=query, + vector_store_id="book_2", + vector_store_search_optional_params={}, + litellm_logging_obj=MagicMock(), + litellm_params={"api_base": "http://localhost:19530", "litellm_embedding_model": "embedding-alias"}, + embedding_executor=embedding_executor, + ) + + assert exc_info.value.status_code == 400 + assert not embedding_executor.mock_calls + client.search.assert_not_called() + + @pytest.mark.parametrize("query", ([], "", " ")) + @pytest.mark.asyncio + async def test_async_grpc_search_rejects_empty_query(self, query: str | list[str]) -> None: + client: Final = AsyncMock() + embedding_executor: Final = MagicMock() + config: Final = MilvusGRPCVectorStoreConfig(async_client=client) + + with pytest.raises(litellm.BadRequestError, match="query must not be empty") as exc_info: + await config.aexecute_search_vector_store_request( + query=query, + vector_store_id="book_2", + vector_store_search_optional_params={}, + litellm_logging_obj=MagicMock(), + litellm_params={"api_base": "http://localhost:19530", "litellm_embedding_model": "embedding-alias"}, + embedding_executor=embedding_executor, + ) + + assert exc_info.value.status_code == 400 + assert not embedding_executor.mock_calls + client.search.assert_not_called() + @pytest.mark.parametrize( ("parameter", "value"), [