diff --git a/litellm/llms/milvus/vector_stores/grpc_transformation.py b/litellm/llms/milvus/vector_stores/grpc_transformation.py new file mode 100644 index 00000000000..75206bfb0e6 --- /dev/null +++ b/litellm/llms/milvus/vector_stores/grpc_transformation.py @@ -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() diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 1feda0b0bb5..b301c7bc132 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -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( diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index fe4732c6492..ec266cb87ca 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -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( diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 6e94a5a88ac..997c6bae49c 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -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) diff --git a/litellm/types/router.py b/litellm/types/router.py index 7ebd50f1328..4135917cf12 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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: diff --git a/litellm/utils.py b/litellm/utils.py index 8b1b32ea328..d6284e0975f 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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, ) diff --git a/litellm/vector_stores/main.py b/litellm/vector_stores/main.py index 976e6dead76..75ea3fae2d6 100644 --- a/litellm/vector_stores/main.py +++ b/litellm/vector_stores/main.py @@ -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: diff --git a/pyproject.toml b/pyproject.toml index 5567fb5d6e2..27f1f970d8b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", 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 1abbbe91e97..a5a74c174b3 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 @@ -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 diff --git a/tests/vector_store_tests/test_milvus_vector_store.py b/tests/vector_store_tests/test_milvus_vector_store.py index 2ba9168b49f..9da643144c3 100644 --- a/tests/vector_store_tests/test_milvus_vector_store.py +++ b/tests/vector_store_tests/test_milvus_vector_store.py @@ -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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index d94b5425ba9..e668645a4ec 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */ diff --git a/uv.lock b/uv.lock index 99d694848f5..fa50250ee86 100644 --- a/uv.lock +++ b/uv.lock @@ -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"