fix(vector-stores): guard nested Milvus settings and validate searches

This commit is contained in:
Yujong Lee 2026-09-05 08:11:23 -07:00
parent 2a13c88789
commit 6b6d7a1e24
4 changed files with 95 additions and 12 deletions

View file

@ -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]:

View file

@ -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"
)

View file

@ -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",

View file

@ -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"),
[