refactor(vector-stores): simplify Milvus gRPC integration

This commit is contained in:
Yujong Lee 2026-09-04 17:53:23 -07:00
parent 79c32ad0a7
commit 123406c0fe
15 changed files with 131 additions and 140 deletions

View file

@ -57,7 +57,7 @@
"limit": 5601
},
"reportMissingTypeArgument": {
"limit": 15284
"limit": 15281
},
"reportMissingTypeStubs": {
"limit": 40
@ -99,13 +99,13 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44149
"limit": 43934
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38309
"limit": 38305
},
"reportUnknownParameterType": {
"limit": 19622

View file

@ -140,18 +140,6 @@ class VectorStorePreCallHook(CustomLogger):
request_metadata = (
request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {}
)
if llm_router is not None or prisma_client is not None:
from litellm.proxy.vector_store_endpoints.utils import (
assert_proxy_admin_for_user_supplied_vector_store_connection,
)
assert_proxy_admin_for_user_supplied_vector_store_connection(
custom_llm_provider=litellm_params_for_vector_store.get(
"custom_llm_provider", custom_llm_provider
),
litellm_params=litellm_params_for_vector_store,
managed=True,
)
if llm_router is not None:
search_function = cast( # cast-ok: normalize router search callable
Callable[..., Awaitable[VectorStoreSearchResponse]],
@ -163,6 +151,18 @@ class VectorStorePreCallHook(CustomLogger):
litellm.vector_stores.asearch,
)
try:
if llm_router is not None or prisma_client is not None:
from litellm.proxy.vector_store_endpoints.utils import (
assert_proxy_admin_for_user_supplied_vector_store_connection,
)
assert_proxy_admin_for_user_supplied_vector_store_connection(
custom_llm_provider=litellm_params_for_vector_store.get(
"custom_llm_provider", custom_llm_provider
),
litellm_params=litellm_params_for_vector_store,
managed=True,
)
search_response = await search_function(
**{
"vector_store_id": vector_store_id,

View file

@ -62,7 +62,7 @@ class AzurePassthroughConfig(BasePassthroughConfig):
) -> dict:
return BaseAzureLLM._base_validate_azure_environment(
headers=headers,
litellm_params=GenericLiteLLMParams(**{**litellm_params, "api_key": api_key}),
litellm_params=GenericLiteLLMParams.model_validate({**litellm_params, "api_key": api_key}),
)
@staticmethod

View file

@ -1,5 +1,5 @@
import typing
from collections.abc import Awaitable, Mapping, Sequence
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
@ -10,6 +10,7 @@ from typing_extensions import Protocol
import litellm
from litellm.llms.base_llm.vector_store.transformation import (
BaseDirectVectorStoreConfig,
LiteLLMVectorStoreEmbeddingExecutor,
VectorStoreEmbeddingExecutor,
)
from litellm.secret_managers.main import get_secret_str
@ -75,30 +76,6 @@ class _AsyncMilvusClient(Protocol):
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
@ -215,14 +192,10 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
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,
@ -398,18 +371,11 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
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,
)
executor: Final = embedding_executor or LiteLLMVectorStoreEmbeddingExecutor()
embedding_response: Final = executor.embed(
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)
@ -451,18 +417,11 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
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,
)
executor: Final = embedding_executor or LiteLLMVectorStoreEmbeddingExecutor()
embedding_response: Final = await executor.aembed(
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)

View file

@ -2,11 +2,11 @@ import json
import re
from collections.abc import Collection, Mapping
from types import MappingProxyType, UnionType
from typing import Any, Final, Union, get_args, get_origin
from typing import Annotated, Any, Final, Union, get_args, get_origin
import orjson
from fastapi import Request, UploadFile, status
from typing_extensions import ReadOnly
from typing_extensions import NotRequired, ReadOnly, Required
from litellm._logging import verbose_proxy_logger
from litellm.constants import MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB
@ -42,11 +42,17 @@ def _is_json_content_type(content_type: str) -> bool:
return _normalize_media_type(content_type) == "application/json"
def _type_args(annotation: object) -> tuple[object, ...]:
return tuple(get_args(annotation))
def _numeric_form_type(annotation: object) -> type[int] | type[float] | None:
"""The scalar to parse an ``int``/``float``-typed field as, else ``None``."""
unwrapped: Final = get_args(annotation)[0] if get_origin(annotation) is ReadOnly else annotation
unwrapped: object = annotation # rebind-ok: type qualifiers may be nested to arbitrary depth
while get_origin(unwrapped) in (Annotated, NotRequired, ReadOnly, Required):
unwrapped = _type_args(unwrapped)[0] # rebind-ok: peel one qualifier per iteration
candidates: Final = (
tuple(arg for arg in get_args(unwrapped) if arg is not type(None))
tuple(arg for arg in _type_args(unwrapped) if arg is not type(None))
if get_origin(unwrapped) in (Union, UnionType)
else (unwrapped,)
)

View file

@ -8643,12 +8643,8 @@ class Router:
if ptu_error is not None and is_ptu_cost_attribution_enabled():
raise ValueError(ptu_error)
zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(
**(
_litellm_params
if zeroed_pricing is None
else MappingProxyType({**_litellm_params, **zeroed_pricing})
)
litellm_params: Final = LiteLLM_Params.model_validate(
_litellm_params if zeroed_pricing is None else MappingProxyType({**_litellm_params, **zeroed_pricing})
)
warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params)
deployment = Deployment(

View file

@ -373,10 +373,7 @@ 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_transport: Literal["rest", "grpc"] | None = None
milvus_text_field: str | None = None
milvus_db_name: str | None = None
milvus_partition_names: list[str] | None = None
@ -387,13 +384,6 @@ 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

@ -102,9 +102,9 @@ class VectorStoreSearchResponse(TypedDict, total=False):
class VectorStoreSearchOptionalRequestParams(TypedDict, total=False):
"""TypedDict for Optional parameters supported by the vector store search API."""
filters: dict | None
filters: dict[str, object] | None
max_num_results: int | None
ranking_options: dict | None
ranking_options: dict[str, object] | None
rewrite_query: bool | None

View file

@ -10,6 +10,8 @@ from typing import (
get_args,
)
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import remove_items_at_indices
from litellm.repositories.table_repositories import (
@ -31,6 +33,19 @@ if TYPE_CHECKING:
else:
PrismaClient = Any
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(dict[str, object] | None)
def _deserialize_litellm_params(
value: object,
) -> dict[str, object] | None: # mutable-ok: managed vector store rows expose JSON objects as dicts
try:
if isinstance(value, str):
return _LITELLM_PARAMS_ADAPTER.validate_json(value)
return _LITELLM_PARAMS_ADAPTER.validate_python(value)
except ValidationError:
return {}
class VectorStoreIndexRegistry:
def __init__(self, vector_store_indexes: list[LiteLLM_ManagedVectorStoreIndex] = []):
@ -411,7 +426,7 @@ class VectorStoreRegistry:
# cast to VectorStoreConfig
litellm_vector_store_config = LiteLLM_VectorStoreConfig(**vector_store_config)
vector_store_name = litellm_vector_store_config.get("vector_store_name")
vector_store_litellm_params: dict[str, Any] = litellm_vector_store_config.get("litellm_params") or {}
vector_store_litellm_params: dict[str, Any] = dict(litellm_vector_store_config.get("litellm_params") or {})
vector_store_id = vector_store_litellm_params.get("vector_store_id")
if not isinstance(vector_store_id, str) or not vector_store_id:
@ -515,6 +530,9 @@ class VectorStoreRegistry:
)
for vector_store in _vector_stores_from_db:
_dict_vector_store = dict(vector_store)
_dict_vector_store["litellm_params"] = _deserialize_litellm_params(
_dict_vector_store.get("litellm_params")
)
_litellm_managed_vector_store = LiteLLM_ManagedVectorStore(**_dict_vector_store)
vector_stores_from_db.append(_litellm_managed_vector_store)
return vector_stores_from_db

View file

@ -21,6 +21,6 @@
"limit": 117
},
"TQ008": {
"limit": 11000
"limit": 10997
}
}

View file

@ -150,7 +150,7 @@ async def test_hook_searches_through_the_injected_router_with_the_request_metada
@pytest.mark.asyncio
async def test_hook_does_not_search_an_untrusted_managed_milvus_grpc_connection(
async def test_hook_skips_an_untrusted_managed_milvus_grpc_connection(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(
@ -165,7 +165,11 @@ async def test_hook_does_not_search_an_untrusted_managed_milvus_grpc_connection(
"milvus_transport": "grpc",
"api_base": "http://internal-milvus:19530",
},
)
),
LiteLLM_ManagedVectorStore(
vector_store_id="safe",
custom_llm_provider="bedrock",
),
],
),
)
@ -173,12 +177,12 @@ async def test_hook_does_not_search_an_untrusted_managed_milvus_grpc_connection(
_, messages, _ = await _run_hook(
VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)),
["legacy"],
["legacy", "safe"],
FakeLoggingObj({}),
)
assert router.calls == []
assert messages == [{"role": "user", "content": "what is litellm?"}]
assert [call["vector_store_id"] for call in router.calls] == ["safe"]
assert messages[0]["content"] == "Context:\n\ncontext from safe\n\n"
@pytest.mark.asyncio

View file

@ -780,20 +780,17 @@ async def test_unmarked_managed_milvus_connection_requires_admin_resave():
@pytest.mark.asyncio
async def test_config_loaded_milvus_grpc_connection_is_trusted():
registry = VectorStoreRegistry()
registry.load_vector_stores_from_config(
[
{
"vector_store_name": "configured",
"litellm_params": {
"vector_store_id": "configured",
"custom_llm_provider": "milvus",
"milvus_transport": "grpc",
"api_base": "https://configured-milvus:19530",
"litellm_embedding_model": "team-embedding-alias",
},
}
]
)
source = {
"vector_store_name": "configured",
"litellm_params": {
"vector_store_id": "configured",
"custom_llm_provider": "milvus",
"milvus_transport": "grpc",
"api_base": "https://configured-milvus:19530",
"litellm_embedding_model": "team-embedding-alias",
},
}
registry.load_vector_stores_from_config([source])
with patch.object( # test-quality-ok: config trust is established by the process-wide registry
litellm, "vector_store_registry", registry
@ -806,6 +803,7 @@ async def test_config_loaded_milvus_grpc_connection_is_trusted():
assert result["api_base"] == "https://configured-milvus:19530"
assert result[MILVUS_ADMIN_CONFIGURED_CONNECTION] is True
assert MILVUS_ADMIN_CONFIGURED_CONNECTION not in source["litellm_params"]
@pytest.mark.asyncio

View file

@ -8,7 +8,7 @@ from fastapi.testclient import TestClient
from datetime import datetime, timezone
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
import litellm
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
@ -16,6 +16,24 @@ from litellm.vector_stores.main import search
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
@pytest.mark.asyncio
async def test_db_vector_store_litellm_params_are_deserialized():
prisma_client = MagicMock()
prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(
return_value=[
{
"vector_store_id": "managed-milvus",
"custom_llm_provider": "milvus",
"litellm_params": json.dumps({"milvus_transport": "grpc"}),
}
]
)
[vector_store] = await VectorStoreRegistry._get_vector_stores_from_db(prisma_client)
assert vector_store["litellm_params"] == {"milvus_transport": "grpc"}
@pytest.fixture(autouse=True)
def clear_client_cache():
"""

View file

@ -442,8 +442,9 @@ class TestMilvusVectorStore:
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)
embedding_executor = MagicMock()
embedding_executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
config = MilvusGRPCVectorStoreConfig(sync_client=mock_client)
response = config.execute_search_vector_store_request(
query="what is machine learning?",
vector_store_id="book_2",
@ -466,9 +467,10 @@ class TestMilvusVectorStore:
"milvus_db_name": "tenant_a_db",
"milvus_partition_names": ["tenant_a_partition"],
},
embedding_executor=embedding_executor,
)
mock_embedding.assert_called_once_with(
embedding_executor.embed.assert_called_once_with(
"text-embedding-3-large",
"what is machine learning?",
{"api_key": "mock_openai_api_key"},
@ -507,8 +509,9 @@ class TestMilvusVectorStore:
]
)
mock_client.close = AsyncMock()
mock_embedding = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
config = MilvusGRPCVectorStoreConfig(async_client=mock_client, aembedding_fn=mock_embedding)
embedding_executor = MagicMock()
embedding_executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
config = MilvusGRPCVectorStoreConfig(async_client=mock_client)
response = await config.aexecute_search_vector_store_request(
query=["what is", "machine learning?"],
vector_store_id="book_2",
@ -524,9 +527,10 @@ class TestMilvusVectorStore:
"litellm_embedding_model": "text-embedding-3-large",
"milvus_text_field": "book_intro_text",
},
embedding_executor=embedding_executor,
)
mock_embedding.assert_awaited_once_with(
embedding_executor.aembed.assert_awaited_once_with(
"text-embedding-3-large",
"what is machine learning?",
{},
@ -550,10 +554,9 @@ class TestMilvusVectorStore:
}
]
]
config = MilvusGRPCVectorStoreConfig(
sync_client=mock_client,
embedding_fn=MagicMock(return_value=MOCK_EMBEDDING_RESPONSE),
)
embedding_executor = MagicMock()
embedding_executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
config = MilvusGRPCVectorStoreConfig(sync_client=mock_client)
response = config.execute_search_vector_store_request(
query="what is machine learning?",
@ -568,6 +571,7 @@ class TestMilvusVectorStore:
"litellm_embedding_model": "text-embedding-3-large",
"milvus_text_field": "body",
},
embedding_executor=embedding_executor,
)
assert mock_client.search.call_args.kwargs["output_fields"] == ["category", "body"]
@ -584,8 +588,9 @@ class TestMilvusVectorStore:
)
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)
embedding_executor = MagicMock()
embedding_executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
config = MilvusGRPCVectorStoreConfig(sync_client=mock_client)
with pytest.raises(ValueError, match=r"Input should be (greater|less) than or equal"):
config.execute_search_vector_store_request(
@ -597,9 +602,10 @@ class TestMilvusVectorStore:
"api_base": "https://milvus.example.com:19530",
"litellm_embedding_model": "openai/text-embedding-3-small",
},
embedding_executor=embedding_executor,
)
mock_embedding.assert_not_called()
embedding_executor.embed.assert_not_called()
mock_client.search.assert_not_called()
@pytest.mark.parametrize(
@ -612,8 +618,9 @@ class TestMilvusVectorStore:
)
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)
embedding_executor = MagicMock()
embedding_executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
config = MilvusGRPCVectorStoreConfig(sync_client=mock_client)
with pytest.raises(litellm.BadRequestError, match=f"does not support the {parameter} parameter") as exc_info:
config.execute_search_vector_store_request(
@ -625,10 +632,11 @@ class TestMilvusVectorStore:
"api_base": "https://milvus.example.com:19530",
"litellm_embedding_model": "openai/text-embedding-3-small",
},
embedding_executor=embedding_executor,
)
assert exc_info.value.status_code == 400
mock_embedding.assert_not_called()
embedding_executor.embed.assert_not_called()
mock_client.search.assert_not_called()
def test_grpc_transport_selects_direct_config(self):

View file

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