mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(vector-stores): simplify Milvus gRPC integration
This commit is contained in:
parent
79c32ad0a7
commit
123406c0fe
15 changed files with 131 additions and 140 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -21,6 +21,6 @@
|
|||
"limit": 117
|
||||
},
|
||||
"TQ008": {
|
||||
"limit": 11000
|
||||
"limit": 10997
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue