mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vector-stores): guard nested Milvus settings and validate searches
This commit is contained in:
parent
2a13c88789
commit
6b6d7a1e24
4 changed files with 95 additions and 12 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue