refactor(milvus): share search arguments and preserve managed inputs

This commit is contained in:
Yujong Lee 2026-09-05 08:35:53 -07:00
parent 2079b2e29c
commit bbe9982069
11 changed files with 171 additions and 141 deletions

View file

@ -1,9 +1,9 @@
{
"reportAny": {
"limit": 13429
"limit": 13428
},
"reportArgumentType": {
"limit": 2178
"limit": 2168
},
"reportAssignmentType": {
"limit": 319
@ -24,7 +24,7 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 3367
"limit": 3366
},
"reportFunctionMemberAccess": {
"limit": 7
@ -57,7 +57,7 @@
"limit": 5570
},
"reportMissingTypeArgument": {
"limit": 15270
"limit": 15264
},
"reportMissingTypeStubs": {
"limit": 40
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 43439
"limit": 43218
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38255
"limit": 38239
},
"reportUnknownParameterType": {
"limit": 19584
"limit": 19582
},
"reportUnknownVariableType": {
"limit": 29827
"limit": 29823
},
"reportUnnecessaryCast": {
"limit": 110
@ -123,7 +123,7 @@
"limit": 4
},
"reportUnnecessaryIsInstance": {
"limit": 814
"limit": 813
},
"reportUntypedBaseClass": {
"limit": 0

View file

@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Final
import httpx
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
from typing_extensions import Protocol
from typing_extensions import Protocol, ReadOnly, TypedDict
import litellm
from litellm.llms.base_llm.vector_store.transformation import (
@ -36,6 +36,21 @@ _MILVUS_ENTITY_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_STRING_KEYS_ADAPTER: Final = TypeAdapter(tuple[str, ...])
class _MilvusSearchArguments(TypedDict):
collection_name: ReadOnly[str]
data: ReadOnly[list[list[float]]] # mutable-ok: PyMilvus requires nested list search data
anns_field: ReadOnly[str | None]
limit: ReadOnly[int]
filter: ReadOnly[str]
offset: ReadOnly[int | None]
group_by_field: ReadOnly[str | None]
output_fields: ReadOnly[list[str]] # mutable-ok: PyMilvus requires list output fields
search_params: ReadOnly[dict[str, object] | None] # mutable-ok: PyMilvus requires dict search params
consistency_level: ReadOnly[str | None]
partition_names: ReadOnly[list[str] | None] # mutable-ok: PyMilvus requires list partition names
timeout: ReadOnly[float | None]
class _SyncMilvusClient(Protocol):
def search(
self,
@ -311,43 +326,14 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
)
@staticmethod
def _sync_search(
client: _SyncMilvusClient,
def _search_arguments(
vector_store_id: str,
query_vector: Sequence[float],
options: _MilvusSearchOptions,
params: _MilvusSearchParams,
timeout: float | None,
) -> object:
return client.search(
collection_name=vector_store_id,
data=[list(query_vector)], # mutable-ok: PyMilvus requires nested list search data
anns_field=options.anns_field,
limit=options.result_limit,
filter=options.filter_expression,
offset=options.offset,
group_by_field=options.grouping_field,
output_fields=options.output_fields_with_text(params.text_field),
search_params=dict(options.search_params) # mutable-ok: PyMilvus requires dict search params
if options.search_params is not None
else None,
consistency_level=options.consistency_level,
partition_names=list(params.milvus_partition_names) # mutable-ok: PyMilvus requires list partition names
if params.milvus_partition_names is not None
else None,
timeout=timeout,
)
@staticmethod
async def _async_search(
client: _AsyncMilvusClient,
vector_store_id: str,
query_vector: Sequence[float],
options: _MilvusSearchOptions,
params: _MilvusSearchParams,
timeout: float | None,
) -> object:
return await client.search(
) -> _MilvusSearchArguments:
return _MilvusSearchArguments(
collection_name=vector_store_id,
data=[list(query_vector)], # mutable-ok: PyMilvus requires nested list search data
anns_field=options.anns_field,
@ -387,30 +373,18 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
)
query_vector: Final = _EmbeddingPayload.model_validate(embedding_response).vector()
connection_timeout, search_timeout = self._timeouts(timeout)
if self.sync_client is not None:
raw_result: Final = self._sync_search(
self.sync_client,
vector_store_id,
query_vector,
options,
params,
search_timeout,
)
return self._to_response(raw_result, query_text, params.text_field)
client: Final = _new_sync_client(params.uri, params.token, params.db_name, connection_timeout)
arguments: Final = self._search_arguments(vector_store_id, query_vector, options, params, search_timeout)
client: Final = (
self.sync_client
if self.sync_client is not None
else _new_sync_client(params.uri, params.token, params.db_name, connection_timeout)
)
try:
result: Final = self._sync_search(
client,
vector_store_id,
query_vector,
options,
params,
search_timeout,
)
result: Final = client.search(**arguments)
return self._to_response(result, query_text, params.text_field)
finally:
client.close()
if self.sync_client is None:
client.close()
async def aexecute_search_vector_store_request(
self,
@ -433,27 +407,15 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
)
query_vector: Final = _EmbeddingPayload.model_validate(embedding_response).vector()
connection_timeout, search_timeout = self._timeouts(timeout)
if self.async_client is not None:
raw_result: Final = await self._async_search(
self.async_client,
vector_store_id,
query_vector,
options,
params,
search_timeout,
)
return self._to_response(raw_result, query_text, params.text_field)
client: Final = _new_async_client(params.uri, params.token, params.db_name, connection_timeout)
arguments: Final = self._search_arguments(vector_store_id, query_vector, options, params, search_timeout)
client: Final = (
self.async_client
if self.async_client is not None
else _new_async_client(params.uri, params.token, params.db_name, connection_timeout)
)
try:
result: Final = await self._async_search(
client,
vector_store_id,
query_vector,
options,
params,
search_timeout,
)
result: Final = await client.search(**arguments)
return self._to_response(result, query_text, params.text_field)
finally:
await client.close()
if self.async_client is None:
await client.close()

View file

@ -73,10 +73,10 @@ def build_request_data_from_managed_vector_store(
async def _update_request_data_with_litellm_managed_vector_store_registry(
data: dict,
data: Mapping[str, object],
vector_store_id: str,
user_api_key_dict: UserAPIKeyAuth | None = None,
) -> dict:
) -> dict[str, object]:
"""
Update the request data with the litellm managed vector store registry.
@ -92,27 +92,31 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
vector_store_id=vector_store_id
)
if vector_store_to_run is None:
data.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None)
caller_data: Final = {key: value for key, value in data.items() if key != MILVUS_ADMIN_CONFIGURED_CONNECTION}
if user_api_key_dict is not None:
assert_proxy_admin_for_user_supplied_vector_store_connection(
custom_llm_provider=data.get("custom_llm_provider"),
litellm_params=data,
custom_llm_provider=caller_data.get("custom_llm_provider"),
litellm_params=caller_data,
user_api_key_dict=user_api_key_dict,
)
return data
return caller_data
if user_api_key_dict is not None:
await assert_user_can_access_vector_store(
vector_store=vector_store_to_run,
user_api_key_dict=user_api_key_dict,
)
if normalize_vector_store_provider(vector_store_to_run.get("custom_llm_provider")) == "milvus":
for field in MILVUS_MANAGED_CONFIGURATION_FIELDS:
data.pop(field, None)
data.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None)
data.pop("custom_llm_provider", None)
data.pop("litellm_credential_name", None)
blocked_fields: Final = frozenset(
(MILVUS_ADMIN_CONFIGURED_CONNECTION, "custom_llm_provider", "litellm_credential_name")
) | (
MILVUS_MANAGED_CONFIGURATION_FIELDS
if normalize_vector_store_provider(vector_store_to_run.get("custom_llm_provider")) == "milvus"
else frozenset()
)
managed_data: Final = build_request_data_from_managed_vector_store(vector_store_to_run)
request_data: Final = {**data, **managed_data} # mutable-ok: request processing requires a mutable payload
request_data: Final = {
**{key: value for key, value in data.items() if key not in blocked_fields},
**managed_data,
}
if user_api_key_dict is not None:
assert_proxy_admin_for_user_supplied_vector_store_connection(
custom_llm_provider=request_data.get("custom_llm_provider"),

View file

@ -336,17 +336,10 @@ async def new_vector_store(
user_id=user_api_key_dict.user_id,
)
# Apply the same litellm_params redaction the list / info / update
# endpoints already use, so a caller-supplied credential or a
# cleartext value persisted by an earlier proxy version doesn't
# come back in the response.
response_vs: Final = LiteLLM_ManagedVectorStore(**new_vector_store)
response_vs["litellm_params"] = _redact_sensitive_litellm_params(new_vector_store.get("litellm_params"))
return {
"status": "success",
"message": f"Vector store {vector_store.get('vector_store_id')} created successfully",
"vector_store": response_vs,
"vector_store": _redact_vector_store(new_vector_store),
}
except HTTPException:
raise
@ -398,9 +391,7 @@ def _synchronize_vector_store_registry(
def _redact_vector_store(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_ManagedVectorStore:
redacted: Final = LiteLLM_ManagedVectorStore(**vector_store)
redacted["litellm_params"] = _redact_sensitive_litellm_params(vector_store.get("litellm_params"))
return redacted
return {**vector_store, "litellm_params": _redact_sensitive_litellm_params(vector_store.get("litellm_params"))}
@router.get(
@ -585,10 +576,11 @@ async def get_vector_store_info(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
vector_store_dict: Final = dict(vector_store_typed)
if "litellm_params" in vector_store_dict:
vector_store_dict["litellm_params"] = _redact_sensitive_litellm_params(vector_store_dict["litellm_params"])
return {"vector_store": vector_store_dict}
return {
"vector_store": _redact_vector_store(vector_store_typed)
if "litellm_params" in vector_store_typed
else dict(vector_store_typed)
}
except HTTPException:
# Preserve 403/404 from the access-control / not-found checks above;
# the catch-all below would otherwise rewrite them as 500.
@ -683,16 +675,10 @@ async def update_vector_store(
"Updated vector store %s in both database and in-memory registry", vector_store_id
)
# The DB row is returned in full, so the response would otherwise
# echo the persisted ``litellm_params`` (including provider
# credentials) back to the caller — even when the caller only
# changed unrelated fields like ``vector_store_description``.
response_vs: Final = LiteLLM_ManagedVectorStore(**updated_vs)
response_vs["litellm_params"] = _redact_sensitive_litellm_params(updated_vs.get("litellm_params"))
return {
"status": "success",
"message": f"Vector store {vector_store_id} updated successfully",
"vector_store": response_vs,
"vector_store": _redact_vector_store(updated_vs),
}
except HTTPException:
# Preserve 403/404 responses from the access-control / not-found

View file

@ -1,4 +1,3 @@
import json
import re
from collections.abc import Iterable, Mapping
from types import MappingProxyType
@ -25,6 +24,7 @@ from litellm.types.vector_stores import (
VectorStoreIndexEndpoints,
)
from litellm.utils import ProviderConfigManager
from litellm.vector_stores.vector_store_registry import deserialize_litellm_params
MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = frozenset(
{
@ -48,13 +48,9 @@ def _normalize_litellm_params(
) -> LiteLLM_ManagedVectorStore:
litellm_params: Final = vector_store.get("litellm_params")
if isinstance(litellm_params, str):
normalized: Final = _MANAGED_VECTOR_STORE_ADAPTER.validate_python(vector_store)
try:
parsed: Final = json.loads(litellm_params)
normalized["litellm_params"] = parsed if isinstance(parsed, dict) else {}
except (TypeError, ValueError):
normalized["litellm_params"] = {}
return normalized
return _MANAGED_VECTOR_STORE_ADAPTER.validate_python(
{**vector_store, "litellm_params": deserialize_litellm_params(litellm_params)}
)
return vector_store

View file

@ -246,7 +246,7 @@
"limit": 109
},
"TRY300": {
"limit": 852
"limit": 851
},
"UP028": {
"limit": 2

View file

@ -21,6 +21,6 @@
"limit": 117
},
"TQ008": {
"limit": 10984
"limit": 10981
}
}

View file

@ -1,4 +1,5 @@
import json
from copy import deepcopy
from datetime import datetime, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
@ -726,19 +727,23 @@ async def test_managed_milvus_uses_only_persisted_connection_for_non_admin():
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store
payload: Final = {
"query": "safe",
"custom_llm_provider": "milvus/probe",
"milvus_transport": "grpc",
"api_base": "http://attacker:19530",
"api_key": "attacker-token",
"litellm_embedding_model": "openai/attacker-model",
"litellm_embedding_config": {"api_base": "http://attacker-embedding"},
}
original_payload: Final = deepcopy(payload)
original_store: Final = deepcopy(managed_vector_store)
with patch.object( # test-quality-ok: the helper reads the process-wide registry directly
litellm, "vector_store_registry", mock_registry
):
result = await _update_request_data_with_litellm_managed_vector_store_registry(
data={
"query": "safe",
"custom_llm_provider": "milvus/probe",
"milvus_transport": "grpc",
"api_base": "http://attacker:19530",
"api_key": "attacker-token",
"litellm_embedding_model": "openai/attacker-model",
"litellm_embedding_config": {"api_base": "http://attacker-embedding"},
},
data=payload,
vector_store_id="managed",
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
)
@ -749,6 +754,9 @@ async def test_managed_milvus_uses_only_persisted_connection_for_non_admin():
assert result["litellm_embedding_model"] == "team-embedding-alias"
assert "litellm_embedding_config" not in result
assert payload == original_payload
assert managed_vector_store == original_store
@pytest.mark.asyncio
async def test_unmarked_managed_milvus_connection_requires_admin_resave():

View file

@ -1,5 +1,6 @@
import json
from datetime import datetime, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -7,7 +8,28 @@ import pytest
import litellm
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
from litellm.vector_stores.main import search
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry, deserialize_litellm_params
@pytest.mark.parametrize(
("value", "expected"),
[('{"milvus_transport":"grpc"}', {"milvus_transport": "grpc"}),
({"milvus_transport": "rest"}, {"milvus_transport": "rest"}),
("invalid-json", {}), ("[]", {}), ("null", {}), (None, {})],
)
def test_deserialize_litellm_params(value: object, expected: dict[str, object]) -> None:
assert deserialize_litellm_params(value) == expected
@pytest.mark.parametrize("serialized_params", ['{"milvus_transport":"grpc"}', "invalid-json", "[]", "null"])
def test_normalizing_stored_params_preserves_source(serialized_params: str) -> None:
from litellm.proxy.vector_store_endpoints.utils import _normalize_litellm_params
store: Final = {"vector_store_id": "documents", "custom_llm_provider": "milvus", "litellm_params": serialized_params}
result: Final = _normalize_litellm_params(store)
assert result["litellm_params"] == deserialize_litellm_params(serialized_params)
assert store["litellm_params"] == serialized_params
@pytest.mark.asyncio

View file

@ -2,6 +2,7 @@
Tests for Milvus Vector Store
"""
import asyncio
import json
from typing import Final, cast
from unittest.mock import AsyncMock, MagicMock, patch
@ -116,6 +117,7 @@ class TestMilvusVectorStore:
({"max_num_results": 1}, 1),
({"max_num_results": 50}, 50),
({"max_num_results": 2, "limit": 7}, 2),
({"max_num_results": None, "limit": 75}, 75),
],
)
@pytest.mark.parametrize("async_mode", [False, True])
@ -127,6 +129,7 @@ class TestMilvusVectorStore:
executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
config: Final = MilvusVectorStoreConfig()
original_params: Final = optional_params.copy()
kwargs: Final = {
"vector_store_id": "documents",
"query": "limit probe",
@ -145,6 +148,7 @@ class TestMilvusVectorStore:
assert body.get("limit") == expected_limit
assert "max_num_results" not in body
assert body["collectionName"] == "documents"
assert optional_params == original_params
@pytest.mark.parametrize("max_num_results", [0, 51])
@pytest.mark.parametrize("async_mode", [False, True])
@ -561,6 +565,7 @@ class TestMilvusVectorStore:
"attributes": {"category": "reference"},
}
]
mock_client.close.assert_not_called()
@pytest.mark.asyncio
async def test_async_grpc_search_infers_vector_field_and_requests_text_by_default(self):
@ -587,6 +592,7 @@ class TestMilvusVectorStore:
VectorStoreSearchOptionalRequestParams,
{
"max_num_results": 2,
"limit": 7,
},
),
litellm_logging_obj=MagicMock(),
@ -607,6 +613,52 @@ class TestMilvusVectorStore:
assert mock_client.search.await_args.kwargs["anns_field"] is None
assert mock_client.search.await_args.kwargs["output_fields"] == ["book_intro_text"]
assert response["data"][0]["content"][0]["text"] == "async result"
mock_client.close.assert_not_called()
@pytest.mark.parametrize("injected", [False, True])
@pytest.mark.parametrize("async_mode", [False, True])
@pytest.mark.parametrize("failure", ["search", "response", "cancellation"])
@pytest.mark.asyncio
async def test_grpc_client_ownership_after_failure(
self, injected: bool, async_mode: bool, failure: str
) -> None:
client: Final = MagicMock()
executor: Final = MagicMock()
executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
error: Final = asyncio.CancelledError if failure == "cancellation" else RuntimeError
search: Final = AsyncMock() if async_mode else MagicMock()
search.side_effect = None if failure == "response" else error("search interrupted")
search.return_value = "invalid result"
client.search = search
client.close = AsyncMock() if async_mode else MagicMock()
config: Final = MilvusGRPCVectorStoreConfig(
sync_client=client if injected and not async_mode else None,
async_client=client if injected and async_mode else None,
)
kwargs: Final = {
"vector_store_id": "documents",
"query": "cleanup probe",
"vector_store_search_optional_params": {},
"litellm_logging_obj": MagicMock(),
"litellm_params": {
"api_base": "http://milvus:19530",
"litellm_embedding_model": "embedding-alias",
},
"embedding_executor": executor,
}
with (
patch("pymilvus.AsyncMilvusClient" if async_mode else "pymilvus.MilvusClient", return_value=client),
pytest.raises(TypeError if failure == "response" else error),
):
(
await config.aexecute_search_vector_store_request(**kwargs)
if async_mode else config.execute_search_vector_store_request(**kwargs)
)
assert client.close.call_count == (0 if injected else 1)
if async_mode:
assert client.close.await_count == (0 if injected else 1)
def test_grpc_search_always_requests_configured_text_field(self):
mock_client = MagicMock()

View file

@ -1,6 +1,6 @@
{
"LIT001": {
"limit": 22165
"limit": 22156
},
"LIT002": {
"limit": 26733
@ -27,7 +27,7 @@
"limit": 0
},
"LIT010": {
"limit": 16442
"limit": 16431
},
"LIT011": {
"limit": 5506