feat(vector-stores): support Milvus gRPC search

This commit is contained in:
Yujong Lee 2026-08-31 16:05:45 -07:00
parent d23bec84c4
commit 4751f07014
12 changed files with 1249 additions and 18 deletions

View file

@ -0,0 +1,489 @@
import typing
from collections.abc import Awaitable, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
import httpx
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
from typing_extensions import Protocol
import litellm
from litellm.llms.base_llm.vector_store.transformation import (
BaseDirectVectorStoreConfig,
VectorStoreEmbeddingExecutor,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.vector_stores import (
VectorStoreResultContent,
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
VectorStoreSearchResult,
)
from .transformation import MILVUS_OPTIONAL_PARAMS
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
DEFAULT_ANNS_FIELD: Final = "book_intro_vector"
DEFAULT_LIMIT: Final = 10
DEFAULT_TEXT_FIELD: Final = "text"
_EMPTY_EMBEDDING_CONFIG: Final[Mapping[str, object]] = MappingProxyType({})
_PYMILVUS_INSTALL_HINT: Final = (
"Milvus gRPC transport requires the 'pymilvus' package. Install it with 'pip install litellm[milvus]'."
)
_MILVUS_ENTITY_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_STRING_KEYS_ADAPTER: Final = TypeAdapter(tuple[str, ...])
class _SyncMilvusClient(Protocol):
def search(
self,
collection_name: str,
data: list[list[float]], # mutable-ok: PyMilvus requires nested list search data
anns_field: str,
limit: int,
filter: str,
offset: int | None,
group_by_field: str | None,
output_fields: list[str] | None, # mutable-ok: PyMilvus requires list output fields
search_params: dict[str, object] | None, # mutable-ok: PyMilvus requires dict search params
consistency_level: str | None,
partition_names: list[str] | None, # mutable-ok: PyMilvus requires list partition names
timeout: float | None,
) -> object: ...
def close(self) -> None: ...
class _AsyncMilvusClient(Protocol):
async def search(
self,
collection_name: str,
data: list[list[float]], # mutable-ok: PyMilvus requires nested list search data
anns_field: str,
limit: int,
filter: str,
offset: int | None,
group_by_field: str | None,
output_fields: list[str] | None, # mutable-ok: PyMilvus requires list output fields
search_params: dict[str, object] | None, # mutable-ok: PyMilvus requires dict search params
consistency_level: str | None,
partition_names: list[str] | None, # mutable-ok: PyMilvus requires list partition names
timeout: float | None,
) -> object: ...
async def close(self) -> None: ...
class _EmbeddingFunction(Protocol):
def __call__(self, model: str, query: str, config: Mapping[str, object]) -> object: ...
class _AsyncEmbeddingFunction(Protocol):
def __call__(self, model: str, query: str, config: Mapping[str, object]) -> Awaitable[object]: ...
def _embedding(model: str, query: str, config: Mapping[str, object]) -> object:
return litellm.embedding( # pyright: ignore[reportUnknownMemberType, reportCallIssue, reportUnknownVariableType] # LiteLLM's embedding overload leaves provider-specific settings untyped
model=model,
input=[query], # mutable-ok: litellm.embedding requires list input
**config, # pyright: ignore[reportArgumentType] # kwargs-ok: embedding aliases carry provider-specific settings
)
async def _aembedding(model: str, query: str, config: Mapping[str, object]) -> object:
return await litellm.aembedding( # pyright: ignore[reportUnknownMemberType] # LiteLLM leaves provider-specific settings untyped
model=model,
input=[query], # mutable-ok: litellm.aembedding requires list input
**config, # kwargs-ok: embedding aliases carry provider-specific settings
)
def _new_sync_client(uri: str, token: str, db_name: str, timeout: float | None) -> _SyncMilvusClient:
try:
from pymilvus import ( # pyright: ignore[reportMissingTypeStubs] # pymilvus does not publish typing metadata
MilvusClient,
)
except ImportError as e:
raise ValueError(_PYMILVUS_INSTALL_HINT) from e
return typing.cast( # noqa: TID251 # cast-ok: pymilvus lacks typing metadata; protocol defines the used surface
_SyncMilvusClient,
MilvusClient(uri=uri, token=token, db_name=db_name, timeout=timeout, dedicated=True),
)
def _new_async_client(uri: str, token: str, db_name: str, timeout: float | None) -> _AsyncMilvusClient:
try:
from pymilvus import ( # pyright: ignore[reportMissingTypeStubs] # pymilvus does not publish typing metadata
AsyncMilvusClient,
)
except ImportError as e:
raise ValueError(_PYMILVUS_INSTALL_HINT) from e
return typing.cast( # noqa: TID251 # cast-ok: pymilvus lacks typing metadata; protocol defines the used surface
_AsyncMilvusClient,
AsyncMilvusClient(uri=uri, token=token, db_name=db_name, timeout=timeout, dedicated=True),
)
class _MilvusSearchParams(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
api_base: str | None = None
api_key: str | None = None
litellm_embedding_model: str | None = None
litellm_embedding_config: Mapping[str, object] | None = None
milvus_db_name: str | None = None
milvus_partition_names: tuple[str, ...] | None = None
milvus_text_field: str | None = None
@property
def uri(self) -> str:
uri: Final = self.api_base or get_secret_str("MILVUS_API_BASE")
if not uri:
raise ValueError("Milvus API base URL is required. Set MILVUS_API_BASE or pass api_base in litellm_params.")
return uri.rstrip("/")
@property
def token(self) -> str:
return self.api_key or get_secret_str("MILVUS_API_KEY") or ""
@property
def db_name(self) -> str:
return self.milvus_db_name or ""
@property
def text_field(self) -> str:
return self.milvus_text_field or DEFAULT_TEXT_FIELD
def require_embedding_model(self) -> str:
if not self.litellm_embedding_model:
raise ValueError(
"litellm_embedding_model is required in litellm_params for Milvus. "
"Example: litellm_params['litellm_embedding_model'] = 'openai/text-embedding-3-small'"
)
return self.litellm_embedding_model
class _MilvusSearchOptions(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", populate_by_name=True)
anns_field: str = Field(default=DEFAULT_ANNS_FIELD, alias="annsField")
limit: int = Field(default=DEFAULT_LIMIT, ge=1, le=50)
max_num_results: int | None = Field(default=None, ge=1, le=50)
filters: Mapping[str, object] | None = None
ranking_options: Mapping[str, object] | None = None
rewrite_query: bool | None = None
filter_expression: str = Field(default="", alias="filter")
offset: int | None = None
grouping_field: str | None = Field(default=None, alias="groupingField")
output_fields: tuple[str, ...] | None = Field(default=None, alias="outputFields")
search_params: Mapping[str, object] | None = Field(default=None, alias="searchParams")
consistency_level: str | None = Field(default=None, alias="consistencyLevel")
@property
def result_limit(self) -> int:
return self.max_num_results or self.limit
class _EmbeddingItem(BaseModel):
model_config = ConfigDict(frozen=True)
embedding: tuple[float, ...]
class _EmbeddingPayload(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
data: tuple[_EmbeddingItem, ...]
def vector(self) -> tuple[float, ...]:
if not self.data:
raise ValueError("The embedding response did not contain an embedding")
return self.data[0].embedding
class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
def __init__(
self,
sync_client: _SyncMilvusClient | None = None,
async_client: _AsyncMilvusClient | None = None,
embedding_fn: _EmbeddingFunction | None = None,
aembedding_fn: _AsyncEmbeddingFunction | None = None,
) -> None:
super().__init__()
self.sync_client = sync_client
self.async_client = async_client
self.embedding_fn = embedding_fn or _embedding
self.aembedding_fn = aembedding_fn or _aembedding
def map_openai_params(
self,
non_default_params: dict[str, object], # mutable-ok: BaseVectorStoreConfig requires dict parameters
optional_params: dict[str, object], # mutable-ok: BaseVectorStoreConfig requires dict parameters
drop_params: bool,
) -> dict[str, object]: # mutable-ok: BaseVectorStoreConfig requires a dict result
mapped_params: Final = { # mutable-ok: BaseVectorStoreConfig requires a dict result
key: value for key, value in non_default_params.items() if key in MILVUS_OPTIONAL_PARAMS
}
return {**optional_params, **mapped_params} # mutable-ok: BaseVectorStoreConfig requires a dict result
@staticmethod
def _search_options(
optional_params: VectorStoreSearchOptionalRequestParams,
) -> _MilvusSearchOptions:
for parameter in ("filters", "ranking_options", "rewrite_query"):
if optional_params.get(parameter) is not None:
raise litellm.BadRequestError(
message=f"Milvus gRPC search does not support the {parameter} parameter",
model="milvus",
llm_provider="milvus",
)
return _MilvusSearchOptions.model_validate(optional_params)
@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)
@staticmethod
def _timeouts(timeout: float | httpx.Timeout | None) -> tuple[float | None, float | None]:
if isinstance(timeout, httpx.Timeout):
return timeout.connect, timeout.read
timeout_seconds: Final = float(timeout) if timeout is not None else None
return timeout_seconds, timeout_seconds
@staticmethod
def _is_hit(value: object) -> typing.TypeGuard[Mapping[str, object]]: # noqa: TID251 # guard-ok: validates Mapping and string keys
if not isinstance(value, Mapping):
return False
try:
_STRING_KEYS_ADAPTER.validate_python(
tuple(value.keys()) # pyright: ignore[reportUnknownArgumentType] # runtime Mapping keys are untyped
)
except ValidationError:
return False
return True
@staticmethod
def _to_result(raw_hit: object, text_field: str) -> VectorStoreSearchResult:
if not MilvusGRPCVectorStoreConfig._is_hit(raw_hit):
raise TypeError(f"Milvus returned an invalid search hit: {type(raw_hit).__name__}")
hit: Final = raw_hit
entity_value: Final = hit.get("entity")
entity: Final = _MILVUS_ENTITY_ADAPTER.validate_python(entity_value or _EMPTY_EMBEDDING_CONFIG)
text_value: Final = entity.get(text_field, hit.get(text_field, ""))
attributes: Final[dict[str, object]] = { # mutable-ok: VectorStoreSearchResult requires dict attributes
**{ # mutable-ok: VectorStoreSearchResult requires dict attributes
key: value for key, value in entity.items() if key != text_field
},
**{ # mutable-ok: VectorStoreSearchResult requires dict attributes
key: value
for key, value in hit.items()
if key not in frozenset(("id", "distance", "entity", text_field))
},
}
score_value: Final = hit.get("distance", 0.0)
score: Final = float(score_value) if isinstance(score_value, int | float) else 0.0
content: Final[list[VectorStoreResultContent]] = [ # mutable-ok: VectorStoreSearchResult requires list content
VectorStoreResultContent(text="" if text_value is None else str(text_value), type="text")
]
return VectorStoreSearchResult(
score=score,
content=content,
file_id=None,
filename=None,
attributes=attributes,
)
@staticmethod
def _is_result_sequence(value: object) -> typing.TypeGuard[Sequence[object]]: # noqa: TID251 # guard-ok: validates non-text Sequence
return isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray)
@classmethod
def _result_sequence(cls, value: object) -> Sequence[object]:
if cls._is_result_sequence(value):
return value
raise TypeError(f"Milvus returned an invalid search result: {type(value).__name__}")
@classmethod
def _to_response(cls, raw_result: object, query_text: str, text_field: str) -> VectorStoreSearchResponse:
result_sets: Final = cls._result_sequence(raw_result)
hits: Final = cls._result_sequence(result_sets[0]) if result_sets else ()
data: Final[list[VectorStoreSearchResult]] = [ # mutable-ok: VectorStoreSearchResponse requires list data
cls._to_result(hit, text_field) for hit in hits
]
return VectorStoreSearchResponse(
object="vector_store.search_results.page",
search_query=query_text,
data=data,
)
@staticmethod
def _sync_search(
client: _SyncMilvusClient,
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=list(options.output_fields) # mutable-ok: PyMilvus requires list output fields
if options.output_fields is not None
else None,
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(
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=list(options.output_fields) # mutable-ok: PyMilvus requires list output fields
if options.output_fields is not None
else None,
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,
)
def execute_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
litellm_logging_obj: "LiteLLMLoggingObj",
litellm_params: Mapping[str, object],
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
timeout: float | httpx.Timeout | None = None,
) -> VectorStoreSearchResponse:
params: Final = _MilvusSearchParams.model_validate(litellm_params)
options: Final = self._search_options(vector_store_search_optional_params)
query_text: Final = self._query_text(query)
embedding_response: Final = (
embedding_executor.embed(
params.require_embedding_model(),
query_text,
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
)
if embedding_executor is not None
else self.embedding_fn(
params.require_embedding_model(),
query_text,
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
)
)
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)
try:
result: Final = self._sync_search(
client,
vector_store_id,
query_vector,
options,
params,
search_timeout,
)
return self._to_response(result, query_text, params.text_field)
finally:
client.close()
async def aexecute_search_vector_store_request(
self,
vector_store_id: str,
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
litellm_logging_obj: "LiteLLMLoggingObj",
litellm_params: Mapping[str, object],
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
timeout: float | httpx.Timeout | None = None,
) -> VectorStoreSearchResponse:
params: Final = _MilvusSearchParams.model_validate(litellm_params)
options: Final = self._search_options(vector_store_search_optional_params)
query_text: Final = self._query_text(query)
embedding_response: Final = (
await embedding_executor.aembed(
params.require_embedding_model(),
query_text,
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
)
if embedding_executor is not None
else await self.aembedding_fn(
params.require_embedding_model(),
query_text,
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
)
)
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)
try:
result: Final = await self._async_search(
client,
vector_store_id,
query_vector,
options,
params,
search_timeout,
)
return self._to_response(result, query_text, params.text_field)
finally:
await client.close()

View file

@ -17,9 +17,13 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.utils import jsonify_object
from litellm.proxy.vector_store_endpoints.utils import (
MILVUS_ADMIN_CONFIGURED_CONNECTION,
MILVUS_MANAGED_CONFIGURATION_FIELDS,
assert_proxy_admin_for_user_supplied_vector_store_connection,
assert_proxy_admin_for_vector_store_index_management,
assert_user_can_access_vector_store,
get_litellm_managed_vector_store,
normalize_vector_store_provider,
)
from litellm.repositories.table_repositories import ManagedVectorStoreIndexRepository
from litellm.types.vector_stores import IndexCreateRequest, IndexListResponse
@ -88,13 +92,35 @@ 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)
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,
user_api_key_dict=user_api_key_dict,
)
return 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,
)
return {**data, **build_request_data_from_managed_vector_store(vector_store_to_run)}
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)
managed_data: Final = build_request_data_from_managed_vector_store(vector_store_to_run)
request_data: Final = {**data, **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"),
litellm_params=request_data,
user_api_key_dict=user_api_key_dict,
managed=True,
)
return request_data
@router.post(

View file

@ -31,8 +31,10 @@ from litellm.proxy._types import (
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
from litellm.proxy.vector_store_endpoints.utils import (
MILVUS_ADMIN_CONFIGURED_CONNECTION,
can_user_access_vector_store,
filter_listable_vector_stores,
prepare_milvus_connection_for_persistence,
)
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import ManagedVectorStoresRepository
@ -96,6 +98,8 @@ def _redact_sensitive_litellm_params(litellm_params: Any, _depth: int = 0) -> An
return litellm_params
out: Final[dict[str, Any]] = {}
for k, v in litellm_params.items():
if k == MILVUS_ADMIN_CONFIGURED_CONNECTION:
continue
if _LITELLM_PARAMS_MASKER.is_sensitive_key(k):
out[k] = REDACTED_BY_LITELM_STRING
elif isinstance(v, dict):
@ -105,6 +109,34 @@ def _redact_sensitive_litellm_params(litellm_params: Any, _depth: int = 0) -> An
return out
def _validated_litellm_params(
litellm_params: dict[str, Any],
) -> dict[str, Any]: # mutable-ok: persistence validation returns a serializable parameter dict
from litellm.types.router import GenericLiteLLMParams
trusted: Final = litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True
validated: Final = GenericLiteLLMParams(**litellm_params).model_dump(exclude_none=True)
if trusted:
validated[MILVUS_ADMIN_CONFIGURED_CONNECTION] = True
return validated
def _litellm_params_dict(
litellm_params: object,
) -> dict[str, Any]: # mutable-ok: update authorization merges a copy of persisted parameters
if isinstance(litellm_params, dict):
return dict(litellm_params) # mutable-ok: callers require an isolated copy for effective-connection merging
if isinstance(litellm_params, str):
try:
parsed: Final = json.loads(litellm_params)
if isinstance(parsed, dict):
return dict(parsed) # mutable-ok: parsed persistence data must be copied before merging
return {} # mutable-ok: non-object persistence data normalizes to an empty mutable mapping
except (TypeError, ValueError):
return {} # mutable-ok: invalid persisted parameters normalize to an empty mutable mapping
return {} # mutable-ok: absent persisted parameters normalize to an empty mutable mapping
async def _fetch_and_authorize_vector_store(
vector_store_id: str,
user_api_key_dict: UserAPIKeyAuth,
@ -175,8 +207,6 @@ async def create_vector_store_in_db(
Raises:
HTTPException: If vector store already exists or database error occurs
"""
from litellm.types.router import GenericLiteLLMParams
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
@ -219,7 +249,7 @@ async def create_vector_store_in_db(
# query through the router at request time, so the credentials stay
# on the deployment and never reach the database.
if litellm_params:
litellm_params_dict: Final = GenericLiteLLMParams(**litellm_params).model_dump(exclude_none=True)
litellm_params_dict: Final = _validated_litellm_params(litellm_params)
data_to_create["litellm_params"] = safe_dumps(litellm_params_dict)
else:
# Provide empty dict if no litellm_params provided
@ -275,6 +305,12 @@ async def new_vector_store(
detail="vector_store_id and custom_llm_provider are required",
)
prepared_litellm_params: Final = prepare_milvus_connection_for_persistence(
custom_llm_provider=custom_llm_provider,
litellm_params=vector_store.get("litellm_params"),
user_api_key_dict=user_api_key_dict,
)
# Extract and validate metadata
metadata: Final = vector_store.get("vector_store_metadata")
validated_metadata: dict | None = None
@ -288,7 +324,7 @@ async def new_vector_store(
vector_store_name=vector_store.get("vector_store_name"),
vector_store_description=vector_store.get("vector_store_description"),
vector_store_metadata=validated_metadata,
litellm_params=vector_store.get("litellm_params"),
litellm_params=prepared_litellm_params,
litellm_credential_name=vector_store.get("litellm_credential_name"),
team_id=user_api_key_dict.team_id,
user_id=user_api_key_dict.user_id,
@ -306,6 +342,8 @@ async def new_vector_store(
"message": f"Vector store {vector_store.get('vector_store_id')} created successfully",
"vector_store": response_vs,
}
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception("Error creating vector store: %s", e)
raise HTTPException(status_code=500, detail=str(e))
@ -582,7 +620,6 @@ async def update_vector_store(
await check_feature_access_for_user(user_api_key_dict, "vector_stores")
from litellm.proxy.proxy_server import prisma_client
from litellm.types.router import GenericLiteLLMParams
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
@ -594,25 +631,35 @@ async def update_vector_store(
# Per-store access control: anyone authenticated who passes the
# premium-feature gate could otherwise update *any* vector store —
# including stores belonging to other teams.
await _fetch_and_authorize_vector_store(
existing_vector_store: Final = await _fetch_and_authorize_vector_store(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
existing_litellm_params: Final = _litellm_params_dict(existing_vector_store.get("litellm_params"))
effective_provider: Final = update_data.get("custom_llm_provider") or existing_vector_store.get(
"custom_llm_provider"
)
effective_litellm_params: Final = prepare_milvus_connection_for_persistence(
custom_llm_provider=effective_provider,
litellm_params=update_data.get("litellm_params"),
user_api_key_dict=user_api_key_dict,
existing_custom_llm_provider=existing_vector_store.get("custom_llm_provider"),
existing_litellm_params=existing_litellm_params,
)
# Handle metadata serialization
if update_data.get("vector_store_metadata") is not None:
update_data["vector_store_metadata"] = safe_dumps(update_data["vector_store_metadata"])
# Handle litellm_params if provided. As with the create path, the
# embedding-config auto-resolve previously persisted cleartext
# credentials into the row; each search now embeds the query
# through the router at request time, so this row only ever stores
# the user-supplied ``litellm_embedding_model`` reference.
if "litellm_params" in update_data:
_input_litellm_params: Final[dict] = update_data.get("litellm_params", {}) or {}
litellm_params_dict: Final = GenericLiteLLMParams(**_input_litellm_params).model_dump(exclude_none=True)
update_data["litellm_params"] = safe_dumps(litellm_params_dict)
# credentials into the row; each search embeds the query through
# the router at request time, so this row only ever stores the
# user-supplied ``litellm_embedding_model`` reference.
if "litellm_params" in update_data or effective_litellm_params != existing_litellm_params:
update_data["litellm_params"] = safe_dumps(_validated_litellm_params(effective_litellm_params))
# Update in database
updated: Final = await _vector_store_table(prisma_client).update(

View file

@ -21,6 +21,24 @@ from litellm.types.utils import LlmProviders
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
from litellm.utils import ProviderConfigManager
MILVUS_ADMIN_CONFIGURED_CONNECTION: Final = "_litellm_admin_configured_milvus_grpc"
MILVUS_GRPC_CONNECTION_FIELDS: Final = frozenset(
{
"api_base",
"api_key",
"milvus_transport",
"milvus_db_name",
"milvus_partition_names",
}
)
MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = MILVUS_GRPC_CONNECTION_FIELDS | frozenset(
{
"litellm_embedding_config",
"litellm_embedding_model",
"milvus_text_field",
}
)
def _normalize_litellm_params(
vector_store: LiteLLM_ManagedVectorStore,
@ -44,6 +62,38 @@ def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
)
def normalize_vector_store_provider(custom_llm_provider: object) -> str | None:
if not isinstance(custom_llm_provider, str) or not custom_llm_provider:
return None
if "/" not in custom_llm_provider:
return custom_llm_provider
try:
_, provider, _, _ = litellm.get_llm_provider(model=custom_llm_provider)
return provider
except Exception:
return custom_llm_provider.split("/", 1)[0]
def strip_client_milvus_trust_marker(
litellm_params: object,
) -> dict[str, Any]: # mutable-ok: caller input is copied before removing server-owned state
sanitized: Final = (
dict(litellm_params) # mutable-ok: authorization requires an isolated mutable copy
if isinstance(litellm_params, dict)
else {} # mutable-ok: absent parameters normalize to an empty mutable mapping
)
sanitized.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None)
return sanitized
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)
and litellm_params.get("milvus_transport") == "grpc"
)
def assert_proxy_admin_for_vector_store_index_management(
user_api_key_dict: UserAPIKeyAuth,
*,
@ -58,6 +108,74 @@ def assert_proxy_admin_for_vector_store_index_management(
)
def assert_proxy_admin_for_user_supplied_vector_store_connection(
custom_llm_provider: object,
litellm_params: object,
user_api_key_dict: UserAPIKeyAuth,
*,
managed: bool = False,
) -> None:
if not is_milvus_grpc_connection(custom_llm_provider, litellm_params):
return
if managed:
if isinstance(litellm_params, dict) and litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True:
return
raise HTTPException(
status_code=403,
detail="This managed Milvus gRPC connection must be re-saved by a proxy admin before it can be used.",
)
if _is_proxy_admin(user_api_key_dict):
return
raise HTTPException(
status_code=403,
detail="Only proxy admins can configure vector store connections. Contact your LiteLLM administrator.",
)
def prepare_milvus_connection_for_persistence(
*,
custom_llm_provider: object,
litellm_params: object,
user_api_key_dict: UserAPIKeyAuth,
existing_custom_llm_provider: object | None = None,
existing_litellm_params: object | None = None,
) -> dict[str, Any]: # mutable-ok: persistence requires a serializable effective-connection dict
supplied: Final = strip_client_milvus_trust_marker(litellm_params)
existing: Final = (
dict(existing_litellm_params) # mutable-ok: authorization compares an isolated persisted-connection copy
if isinstance(existing_litellm_params, dict)
else {} # mutable-ok: absent persisted parameters normalize to an empty mutable mapping
)
effective: Final = { # mutable-ok: the server marker is applied to the persisted effective connection
**existing,
**supplied,
}
previous_is_grpc: Final = is_milvus_grpc_connection(existing_custom_llm_provider, existing)
effective_is_grpc: Final = is_milvus_grpc_connection(custom_llm_provider, effective)
is_create: Final = existing_custom_llm_provider is None
provider_changed: Final = not is_create and custom_llm_provider != existing_custom_llm_provider
connection_changed: Final = any(
existing.get(field) != effective.get(field) for field in MILVUS_GRPC_CONNECTION_FIELDS
)
missing_marker: Final = effective_is_grpc and existing.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is not True
if (previous_is_grpc or effective_is_grpc) and (
is_create or provider_changed or connection_changed or missing_marker
):
if not _is_proxy_admin(user_api_key_dict):
raise HTTPException(
status_code=403,
detail="Only proxy admins can configure vector store connections. Contact your LiteLLM administrator.",
)
if effective_is_grpc:
if _is_proxy_admin(user_api_key_dict) or existing.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True:
effective[MILVUS_ADMIN_CONFIGURED_CONNECTION] = True
else:
effective.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None)
return effective
def _suffix_after_index_name(request_path: str, index_name: str) -> str | None:
"""Return the path suffix after ``/indexes/{index_name}``, or None if absent."""
match: Final = re.search(rf"/indexes/{re.escape(index_name)}(?=$|[/?])", request_path)

View file

@ -373,6 +373,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
# Vector Store Params
vector_store_id: str | None = None
milvus_transport: object | None = Field(
default=None,
json_schema_extra={"enum": ["rest", "grpc"]}, # mutable-ok: Pydantic schema metadata requires JSON containers
)
milvus_text_field: str | None = None
milvus_db_name: str | None = None
milvus_partition_names: list[str] | None = None
@ -383,6 +387,13 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
valkey_text_field: str | None = None
valkey_embedding_field: str | None = None
@field_validator("milvus_transport")
@classmethod
def validate_milvus_transport(cls, value: object | None) -> object | None:
if value not in (None, "rest", "grpc"):
raise ValueError("milvus_transport must be 'rest' or 'grpc'")
return value
@model_validator(mode="before")
@classmethod
def preprocess_input_data(cls, data: object) -> object:

View file

@ -8912,6 +8912,7 @@ class ProviderConfigManager:
def get_provider_vector_stores_config(
provider: LlmProviders,
api_type: str | None = None,
transport: object | None = None,
) -> BaseVectorStoreConfig | None:
"""
v2 vector store config, use this for new vector store integrations
@ -8960,6 +8961,12 @@ class ProviderConfigManager:
return AzureAIVectorStoreConfig()
elif litellm.LlmProviders.MILVUS == provider:
if transport == "grpc":
from litellm.llms.milvus.vector_stores.grpc_transformation import (
MilvusGRPCVectorStoreConfig,
)
return MilvusGRPCVectorStoreConfig()
from litellm.llms.milvus.vector_stores.transformation import (
MilvusVectorStoreConfig,
)

View file

@ -436,6 +436,7 @@ def search(
vector_store_provider_config: Final = ProviderConfigManager.get_provider_vector_stores_config(
provider=litellm.LlmProviders(custom_llm_provider),
api_type=api_type,
transport=litellm_params.milvus_transport,
)
if vector_store_provider_config is None:

View file

@ -125,6 +125,7 @@ grpc = [
# Newest non-yanked release older than the 30-day cutoff.
"grpcio==1.78.0",
]
milvus = ["pymilvus>=2.6.17,<3.0"]
stt-nvidia-riva = [
# NVIDIA Riva STT provider (gRPC). These are imported lazily inside the
# provider handler so litellm core remains usable without them.
@ -150,6 +151,7 @@ proxy-runtime = [
"google-genai>=1.37.0,<2.0",
"anthropic[vertex]>=0.84.0,<1.0",
"grpcio==1.78.0",
"pymilvus>=2.6.17,<3.0",
"prometheus-client>=0.20.0,<1.0",
"langfuse>=2.59.7,<3.0",
"opentelemetry-api==1.28.0",

View file

@ -24,9 +24,12 @@ from litellm.proxy.vector_store_endpoints.management_endpoints import (
new_vector_store,
)
from litellm.proxy.vector_store_endpoints.utils import (
MILVUS_ADMIN_CONFIGURED_CONNECTION,
assert_proxy_admin_for_user_supplied_vector_store_connection,
check_vector_store_permission,
is_allowed_to_call_vector_store_endpoint,
is_allowed_to_call_vector_store_files_endpoint,
prepare_milvus_connection_for_persistence,
)
from litellm.proxy.vector_store_files_endpoints.endpoints import (
_update_request_data_with_model_routing_hint,
@ -617,8 +620,12 @@ async def test_update_request_data_with_litellm_managed_vector_store_registry():
# Test with no vector store registry or DB fallback
with (
patch.object(litellm, "vector_store_registry", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
patch.object( # test-quality-ok: simulates an unmanaged id without mutating the process registry
litellm, "vector_store_registry", None
),
patch( # test-quality-ok: prevents the managed-store database fallback for this unmanaged-id test
"litellm.proxy.proxy_server.prisma_client", None
),
):
original_data = {"existing_key": "existing_value"}
result = await _update_request_data_with_litellm_managed_vector_store_registry(
@ -643,7 +650,9 @@ async def test_managed_vector_store_keeps_embedding_reference_and_explicit_confi
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store
with patch.object(litellm, "vector_store_registry", mock_registry):
with patch.object( # test-quality-ok: injects the persisted managed connection under authorization test
litellm, "vector_store_registry", mock_registry
):
result = await _update_request_data_with_litellm_managed_vector_store_registry(
data={},
vector_store_id="test_store",
@ -655,6 +664,125 @@ async def test_managed_vector_store_keeps_embedding_reference_and_explicit_confi
def test_user_supplied_milvus_grpc_connection_requires_proxy_admin():
with pytest.raises(HTTPException) as exc_info:
assert_proxy_admin_for_user_supplied_vector_store_connection(
custom_llm_provider="milvus",
litellm_params={
"milvus_transport": "grpc",
"api_base": "http://internal-milvus:19530",
},
user_api_key_dict=UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER
),
)
assert exc_info.value.status_code == 403
@pytest.mark.parametrize("provider", ["milvus", "milvus/probe"])
@pytest.mark.asyncio
async def test_unmanaged_milvus_grpc_connection_requires_admin_after_provider_normalization(provider):
with (
patch.object(litellm, "vector_store_registry", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
pytest.raises(HTTPException) as exc_info,
):
await _update_request_data_with_litellm_managed_vector_store_registry(
data={
"custom_llm_provider": provider,
"milvus_transport": "grpc",
"api_base": "http://internal-milvus:19530",
MILVUS_ADMIN_CONFIGURED_CONNECTION: True,
},
vector_store_id="unmanaged",
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_managed_milvus_uses_only_persisted_connection_for_non_admin():
managed_vector_store: LiteLLM_ManagedVectorStore = {
"vector_store_id": "managed",
"custom_llm_provider": "milvus",
"litellm_params": {
"milvus_transport": "grpc",
"api_base": "https://managed-milvus:19530",
"api_key": "managed-token",
"litellm_embedding_model": "team-embedding-alias",
MILVUS_ADMIN_CONFIGURED_CONNECTION: True,
},
}
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store
with patch.object(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"},
},
vector_store_id="managed",
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
)
assert result["custom_llm_provider"] == "milvus"
assert result["api_base"] == "https://managed-milvus:19530"
assert result["api_key"] == "managed-token"
assert result["litellm_embedding_model"] == "team-embedding-alias"
assert "litellm_embedding_config" not in result
@pytest.mark.asyncio
async def test_unmarked_managed_milvus_connection_requires_admin_resave():
managed_vector_store: LiteLLM_ManagedVectorStore = {
"vector_store_id": "legacy",
"custom_llm_provider": "milvus",
"litellm_params": {
"milvus_transport": "grpc",
"api_base": "https://legacy-milvus:19530",
},
}
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store
with (
patch.object( # test-quality-ok: injects an unmarked legacy row through the registry boundary
litellm, "vector_store_registry", mock_registry
),
pytest.raises(HTTPException) as exc_info,
):
await _update_request_data_with_litellm_managed_vector_store_registry(
data={"query": "blocked"},
vector_store_id="legacy",
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
)
assert exc_info.value.status_code == 403
assert "re-saved" in str(exc_info.value.detail)
def test_admin_persistence_strips_forged_marker_and_adds_server_marker():
params = prepare_milvus_connection_for_persistence(
custom_llm_provider="milvus/probe",
litellm_params={
"milvus_transport": "grpc",
"api_base": "https://managed-milvus:19530",
MILVUS_ADMIN_CONFIGURED_CONNECTION: "forged",
},
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
assert params[MILVUS_ADMIN_CONFIGURED_CONNECTION] is True
class TestCheckVectorStorePermission:
"""Test suite for check_vector_store_permission function."""
@ -2566,6 +2694,54 @@ class TestUpdateVectorStoreAccessControlAndRedaction:
credentials to the caller. Both are fixed at the endpoint level.
"""
@pytest.mark.asyncio
async def test_non_admin_cannot_activate_nested_milvus_grpc_connection(self):
from litellm.proxy.vector_store_endpoints.management_endpoints import update_vector_store
from litellm.types.vector_stores import VectorStoreUpdateRequest
existing_row = MagicMock()
existing_row.model_dump.return_value = {
"vector_store_id": "vs_owned",
"custom_llm_provider": "openai",
"team_id": "team-A",
"litellm_params": {
"milvus_transport": "grpc",
"api_base": "http://internal-milvus:19530",
},
}
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=existing_row)
with (
patch( # test-quality-ok: isolates the connection-authorization behavior from the feature entitlement gate
"litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user",
new_callable=AsyncMock,
),
patch( # test-quality-ok: fixes ownership as allowed so this test reaches connection authorization
"litellm.proxy.vector_store_endpoints.management_endpoints._check_vector_store_access",
new_callable=AsyncMock,
return_value=True,
),
patch( # test-quality-ok: injects the endpoint repository boundary with an existing nested connection
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
),
pytest.raises(HTTPException) as exc_info,
):
await update_vector_store(
data=VectorStoreUpdateRequest(
vector_store_id="vs_owned",
custom_llm_provider="milvus",
),
user_api_key_dict=UserAPIKeyAuth(
user_id="owner",
team_id="team-A",
user_role=LitellmUserRoles.INTERNAL_USER,
),
)
assert exc_info.value.status_code == 403
mock_prisma_client.db.litellm_managedvectorstorestable.update.assert_not_called()
@pytest.mark.asyncio
async def test_update_denied_when_caller_cannot_access_store(self):
from unittest.mock import AsyncMock, MagicMock, patch

View file

@ -3,6 +3,7 @@ Tests for Milvus Vector Store
"""
import json
from typing import cast
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -11,8 +12,14 @@ import respx
import litellm
from litellm import Router
from litellm.llms.milvus.vector_stores.grpc_transformation import (
MilvusGRPCVectorStoreConfig,
)
from litellm.llms.milvus.vector_stores.transformation import MilvusVectorStoreConfig
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import EmbeddingResponse
from litellm.types.vector_stores import VectorStoreSearchOptionalRequestParams
from litellm.utils import ProviderConfigManager
from litellm.vector_stores import asearch as vector_store_asearch
from litellm.vector_stores import search as vector_store_search
@ -88,6 +95,16 @@ MOCK_EMBEDDING_RESPONSE.data = [
]
class MockPyMilvusHit(dict[str, object]):
def get(self, key: str, default: object = None) -> object:
if key == "entity":
return {
"book_intro_text": "closest result",
"category": "reference",
}
return super().get(key, default)
class TestMilvusVectorStore:
"""Test Milvus Vector Store with mocked responses"""
@ -422,6 +439,309 @@ class TestMilvusVectorStore:
assert request_data["dbName"] == "tenant_a_db"
assert request_data["partitionNames"] == ["tenant_a_partition"]
def test_grpc_search_uses_pymilvus_client(self):
mock_client = MagicMock()
mock_client.search.return_value = [[MockPyMilvusHit(id=7, distance=0.91, entity={})]]
mock_embedding = MagicMock(return_value=MOCK_EMBEDDING_RESPONSE)
config = MilvusGRPCVectorStoreConfig(sync_client=mock_client, embedding_fn=mock_embedding)
response = config.execute_search_vector_store_request(
query="what is machine learning?",
vector_store_id="book_2",
vector_store_search_optional_params=cast(
VectorStoreSearchOptionalRequestParams,
{
"outputFields": ["book_intro_text", "category"],
"annsField": "book_intro_vector",
"limit": 3,
"filter": 'category == "reference"',
},
),
litellm_logging_obj=MagicMock(),
litellm_params={
"api_base": "https://milvus.example.com:19530",
"api_key": "mock_milvus_api_key",
"litellm_embedding_model": "text-embedding-3-large",
"litellm_embedding_config": {"api_key": "mock_openai_api_key"},
"milvus_text_field": "book_intro_text",
"milvus_db_name": "tenant_a_db",
"milvus_partition_names": ["tenant_a_partition"],
},
)
mock_embedding.assert_called_once_with(
"text-embedding-3-large",
"what is machine learning?",
{"api_key": "mock_openai_api_key"},
)
mock_client.search.assert_called_once()
search_kwargs = mock_client.search.call_args.kwargs
assert search_kwargs["collection_name"] == "book_2"
assert search_kwargs["anns_field"] == "book_intro_vector"
assert search_kwargs["limit"] == 3
assert search_kwargs["filter"] == 'category == "reference"'
assert search_kwargs["output_fields"] == ["book_intro_text", "category"]
assert search_kwargs["partition_names"] == ["tenant_a_partition"]
assert response["search_query"] == "what is machine learning?"
assert response["data"] == [
{
"score": 0.91,
"content": [{"text": "closest result", "type": "text"}],
"file_id": None,
"filename": None,
"attributes": {"category": "reference"},
}
]
@pytest.mark.asyncio
async def test_grpc_search_uses_async_pymilvus_client(self):
mock_client = MagicMock()
mock_client.search = AsyncMock(
return_value=[
[
{
"id": 8,
"distance": 0.88,
"entity": {"book_intro_text": "async result"},
}
]
]
)
mock_client.close = AsyncMock()
mock_embedding = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
config = MilvusGRPCVectorStoreConfig(async_client=mock_client, aembedding_fn=mock_embedding)
response = await config.aexecute_search_vector_store_request(
query=["what is", "machine learning?"],
vector_store_id="book_2",
vector_store_search_optional_params=cast(
VectorStoreSearchOptionalRequestParams,
{
"annsField": "book_intro_vector",
"max_num_results": 2,
},
),
litellm_logging_obj=MagicMock(),
litellm_params={
"api_base": "http://localhost:19530",
"litellm_embedding_model": "text-embedding-3-large",
"milvus_text_field": "book_intro_text",
},
)
mock_embedding.assert_awaited_once_with(
"text-embedding-3-large",
"what is machine learning?",
{},
)
assert mock_client.search.await_args.kwargs["limit"] == 2
assert response["data"][0]["content"][0]["text"] == "async result"
@pytest.mark.parametrize(
"optional_params",
[
{"limit": 0},
{"limit": 51},
{"max_num_results": 0},
{"max_num_results": 51},
],
)
def test_grpc_search_rejects_invalid_result_limits(self, optional_params):
mock_client = MagicMock()
mock_embedding = MagicMock(return_value=MOCK_EMBEDDING_RESPONSE)
config = MilvusGRPCVectorStoreConfig(sync_client=mock_client, embedding_fn=mock_embedding)
with pytest.raises(ValueError, match=r"Input should be (greater|less) than or equal"):
config.execute_search_vector_store_request(
query="what is machine learning?",
vector_store_id="book_2",
vector_store_search_optional_params=optional_params,
litellm_logging_obj=MagicMock(),
litellm_params={
"api_base": "https://milvus.example.com:19530",
"litellm_embedding_model": "openai/text-embedding-3-small",
},
)
mock_embedding.assert_not_called()
mock_client.search.assert_not_called()
@pytest.mark.parametrize(
("parameter", "value"),
[
("filters", {"type": "eq", "key": "category", "value": "reference"}),
("ranking_options", {"score_threshold": 0.5}),
("rewrite_query", True),
],
)
def test_grpc_search_rejects_unsupported_openai_params(self, parameter, value):
mock_client = MagicMock()
mock_embedding = MagicMock(return_value=MOCK_EMBEDDING_RESPONSE)
config = MilvusGRPCVectorStoreConfig(sync_client=mock_client, embedding_fn=mock_embedding)
with pytest.raises(litellm.BadRequestError, match=f"does not support the {parameter} parameter") as exc_info:
config.execute_search_vector_store_request(
query="what is machine learning?",
vector_store_id="book_2",
vector_store_search_optional_params=cast(VectorStoreSearchOptionalRequestParams, {parameter: value}),
litellm_logging_obj=MagicMock(),
litellm_params={
"api_base": "https://milvus.example.com:19530",
"litellm_embedding_model": "openai/text-embedding-3-small",
},
)
assert exc_info.value.status_code == 400
mock_embedding.assert_not_called()
mock_client.search.assert_not_called()
def test_grpc_transport_selects_direct_config(self):
config = ProviderConfigManager.get_provider_vector_stores_config(
provider=litellm.LlmProviders.MILVUS,
transport="grpc",
)
assert isinstance(config, MilvusGRPCVectorStoreConfig)
def test_milvus_transport_defaults_to_rest(self):
config = ProviderConfigManager.get_provider_vector_stores_config(
provider=litellm.LlmProviders.MILVUS,
)
assert isinstance(config, MilvusVectorStoreConfig)
def test_public_grpc_search_passes_connection_settings_to_pymilvus(self):
mock_client = MagicMock()
mock_client.search.return_value = [
[
{
"id": 9,
"distance": 1.0,
"entity": {"text": "secured result"},
}
]
]
def embedding_response(request: httpx.Request, *, stream: bool = False) -> httpx.Response:
return httpx.Response(
200,
request=request,
json={
"data": [
{
"embedding": [1.0, 0.0],
"index": 0,
"object": "embedding",
}
],
"model": "test-embedding",
"object": "list",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
},
)
with (
patch("httpx.Client.send", side_effect=embedding_response),
patch("pymilvus.MilvusClient", return_value=mock_client) as client_class,
):
response = vector_store_search(
query="transport probe",
vector_store_id="documents",
custom_llm_provider="milvus",
milvus_transport="grpc",
api_base="https://milvus.example.com:19530",
api_key="root:Milvus",
litellm_embedding_model="openai/test-embedding",
litellm_embedding_config={
"api_base": "https://embeddings.example/v1",
"api_key": "embedding-key",
},
milvus_db_name="tenant_db",
annsField="vector",
outputFields=["text"],
milvus_text_field="text",
timeout=17,
)
client_class.assert_called_once_with(
uri="https://milvus.example.com:19530",
token="root:Milvus",
db_name="tenant_db",
timeout=17.0,
dedicated=True,
)
mock_client.close.assert_called_once_with()
assert response["data"][0]["content"][0]["text"] == "secured result"
@pytest.mark.asyncio
async def test_async_grpc_uses_distinct_timeouts_and_releases_dedicated_client(self):
mock_client = MagicMock()
mock_client.search = AsyncMock(return_value=[[]])
mock_client.close = AsyncMock()
embedding_executor = MagicMock()
embedding_executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
timeout = httpx.Timeout(connect=3, read=11, write=13, pool=17)
with patch("pymilvus.AsyncMilvusClient", return_value=mock_client) as client_class:
response = await MilvusGRPCVectorStoreConfig().aexecute_search_vector_store_request(
query="transport probe",
vector_store_id="documents",
vector_store_search_optional_params={},
litellm_logging_obj=MagicMock(),
litellm_params={
"api_base": "http://milvus.example.com:19530",
"api_key": "root:Milvus",
"litellm_embedding_model": "embedding-alias",
},
embedding_executor=embedding_executor,
timeout=timeout,
)
client_class.assert_called_once_with(
uri="http://milvus.example.com:19530",
token="root:Milvus",
db_name="",
timeout=3,
dedicated=True,
)
assert mock_client.search.await_args.kwargs["timeout"] == 11
assert response["data"] == []
embedding_executor.aembed.assert_awaited_once_with("embedding-alias", "transport probe", {})
mock_client.close.assert_awaited_once_with()
def test_http_and_https_targets_get_distinct_dedicated_clients(self):
clients = [MagicMock(), MagicMock()]
for client in clients:
client.search.return_value = [[]]
embedding_executor = MagicMock()
embedding_executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
responses = []
with patch("pymilvus.MilvusClient", side_effect=clients) as client_class:
for uri in ("http://milvus.example.com:19530", "https://milvus.example.com:19530"):
responses.append(
MilvusGRPCVectorStoreConfig().execute_search_vector_store_request(
query="transport probe",
vector_store_id="documents",
vector_store_search_optional_params={},
litellm_logging_obj=MagicMock(),
litellm_params={
"api_base": uri,
"litellm_embedding_model": "embedding-alias",
},
embedding_executor=embedding_executor,
)
)
assert [response["data"] for response in responses] == [[], []]
assert [call.kwargs["uri"] for call in client_class.call_args_list] == [
"http://milvus.example.com:19530",
"https://milvus.example.com:19530",
]
assert all(call.kwargs["dedicated"] is True for call in client_class.call_args_list)
for client in clients:
client.close.assert_called_once_with()
def test_invalid_milvus_transport_is_rejected(self):
with pytest.raises(ValueError, match="milvus_transport"):
GenericLiteLLMParams.model_validate({"milvus_transport": "http"})
# @pytest.mark.parametrize("sync_mode", [True, False])
# @pytest.mark.asyncio

View file

@ -29399,6 +29399,11 @@ export interface components {
milvus_partition_names?: string[] | null;
/** Milvus Text Field */
milvus_text_field?: string | null;
/**
* Milvus Transport
* @enum {unknown}
*/
milvus_transport?: "rest" | "grpc";
/** Mock Response */
mock_response?: string | components["schemas"]["ModelResponse"] | unknown | null;
/** Model */
@ -39490,6 +39495,11 @@ export interface components {
milvus_partition_names?: string[] | null;
/** Milvus Text Field */
milvus_text_field?: string | null;
/**
* Milvus Transport
* @enum {unknown}
*/
milvus_transport?: "rest" | "grpc";
/** Mock Response */
mock_response?: string | components["schemas"]["ModelResponse"] | unknown | null;
/** Model */

26
uv.lock generated
View file

@ -4412,6 +4412,9 @@ grpc = [
mcp = [
{ name = "mcp" },
]
milvus = [
{ name = "pymilvus" },
]
mlflow = [
{ name = "mlflow" },
]
@ -4465,6 +4468,7 @@ proxy-runtime = [
{ name = "opentelemetry-instrumentation-fastapi" },
{ name = "opentelemetry-sdk" },
{ name = "prometheus-client" },
{ name = "pymilvus" },
{ name = "pypdf" },
{ name = "sentry-sdk" },
]
@ -4646,6 +4650,8 @@ requires-dist = [
{ name = "pydantic", specifier = ">=2.10.0,<3.0.0" },
{ name = "pydantic-settings", specifier = ">=2.14.1,<3.0" },
{ name = "pyjwt", marker = "extra == 'proxy'", specifier = ">=2.13.0,<3.0" },
{ name = "pymilvus", marker = "extra == 'milvus'", specifier = ">=2.6.17,<3.0" },
{ name = "pymilvus", marker = "extra == 'proxy-runtime'", specifier = ">=2.6.17,<3.0" },
{ name = "pynacl", marker = "extra == 'proxy'", specifier = ">=1.6.2,<2.0" },
{ name = "pypdf", marker = "extra == 'proxy-runtime'", specifier = ">=6.16.1,<7.0" },
{ name = "pyroscope-io", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.8.16,<1.0" },
@ -4672,7 +4678,7 @@ requires-dist = [
{ name = "uvloop", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.22.1,<1.0" },
{ name = "websockets", marker = "extra == 'proxy'", specifier = ">=15.0.1,<16.0" },
]
provides-extras = ["proxy", "cli", "extra-proxy", "utils", "caching", "mcp", "saml", "semantic-router", "mlflow", "grpc", "stt-nvidia-riva", "google", "bedrock-realtime", "proxy-runtime"]
provides-extras = ["proxy", "cli", "extra-proxy", "utils", "caching", "mcp", "saml", "semantic-router", "mlflow", "grpc", "milvus", "stt-nvidia-riva", "google", "bedrock-realtime", "proxy-runtime"]
[package.metadata.requires-dev]
ci = [
@ -7616,6 +7622,24 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/d5/6f/9ac2548e290764781f9e7e2aaf0685b086379dabfb29ca38536985471eaf/pylint-4.0.5-py3-none-any.whl", hash = "sha256:00f51c9b14a3b3ae08cff6b2cdd43f28165c78b165b628692e428fb1f8dc2cf2", size = 536694, upload-time = "2026-02-20T09:07:31.028Z" },
]
[[package]]
name = "pymilvus"
version = "2.6.17"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "cachetools" },
{ name = "grpcio" },
{ name = "orjson" },
{ name = "pandas" },
{ name = "protobuf" },
{ name = "python-dotenv" },
{ name = "requests" },
]
sdist = { url = "https://files.pythonhosted.org/packages/61/22/7dad24082f7efa4bf5ff29f0d9241cec6a8c8ebeb08ad126ca7ca4793da7/pymilvus-2.6.17.tar.gz", hash = "sha256:d9f849ef177f242febc4c02ece88feea46a1a9a1d0c67688b600e7cce3b76f43", size = 306908, upload-time = "2026-07-17T06:23:32.592Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/42/a6/9e2824138ffe1b2932c24f836638dabff5c6fbb5a9afcc1ba5ac5cdc39d8/pymilvus-2.6.17-py3-none-any.whl", hash = "sha256:583c680bea60b5944762ae81dd577169efe3b895b579d78e60bab05c087c4f89", size = 342148, upload-time = "2026-07-17T06:23:31.188Z" },
]
[[package]]
name = "pynacl"
version = "1.6.2"