mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(milvus): share search arguments and preserve managed inputs
This commit is contained in:
parent
2079b2e29c
commit
bbe9982069
11 changed files with 171 additions and 141 deletions
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 13429
|
||||
"limit": 13428
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2178
|
||||
"limit": 2168
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 319
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 3367
|
||||
"limit": 3366
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 5570
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15270
|
||||
"limit": 15264
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 43439
|
||||
"limit": 43218
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38255
|
||||
"limit": 38239
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19584
|
||||
"limit": 19582
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29827
|
||||
"limit": 29823
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 110
|
||||
|
|
@ -123,7 +123,7 @@
|
|||
"limit": 4
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"limit": 814
|
||||
"limit": 813
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Final
|
|||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
from typing_extensions import Protocol
|
||||
from typing_extensions import Protocol, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
|
|
@ -36,6 +36,21 @@ _MILVUS_ENTITY_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
|||
_STRING_KEYS_ADAPTER: Final = TypeAdapter(tuple[str, ...])
|
||||
|
||||
|
||||
class _MilvusSearchArguments(TypedDict):
|
||||
collection_name: ReadOnly[str]
|
||||
data: ReadOnly[list[list[float]]] # mutable-ok: PyMilvus requires nested list search data
|
||||
anns_field: ReadOnly[str | None]
|
||||
limit: ReadOnly[int]
|
||||
filter: ReadOnly[str]
|
||||
offset: ReadOnly[int | None]
|
||||
group_by_field: ReadOnly[str | None]
|
||||
output_fields: ReadOnly[list[str]] # mutable-ok: PyMilvus requires list output fields
|
||||
search_params: ReadOnly[dict[str, object] | None] # mutable-ok: PyMilvus requires dict search params
|
||||
consistency_level: ReadOnly[str | None]
|
||||
partition_names: ReadOnly[list[str] | None] # mutable-ok: PyMilvus requires list partition names
|
||||
timeout: ReadOnly[float | None]
|
||||
|
||||
|
||||
class _SyncMilvusClient(Protocol):
|
||||
def search(
|
||||
self,
|
||||
|
|
@ -311,43 +326,14 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _sync_search(
|
||||
client: _SyncMilvusClient,
|
||||
def _search_arguments(
|
||||
vector_store_id: str,
|
||||
query_vector: Sequence[float],
|
||||
options: _MilvusSearchOptions,
|
||||
params: _MilvusSearchParams,
|
||||
timeout: float | None,
|
||||
) -> object:
|
||||
return client.search(
|
||||
collection_name=vector_store_id,
|
||||
data=[list(query_vector)], # mutable-ok: PyMilvus requires nested list search data
|
||||
anns_field=options.anns_field,
|
||||
limit=options.result_limit,
|
||||
filter=options.filter_expression,
|
||||
offset=options.offset,
|
||||
group_by_field=options.grouping_field,
|
||||
output_fields=options.output_fields_with_text(params.text_field),
|
||||
search_params=dict(options.search_params) # mutable-ok: PyMilvus requires dict search params
|
||||
if options.search_params is not None
|
||||
else None,
|
||||
consistency_level=options.consistency_level,
|
||||
partition_names=list(params.milvus_partition_names) # mutable-ok: PyMilvus requires list partition names
|
||||
if params.milvus_partition_names is not None
|
||||
else None,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _async_search(
|
||||
client: _AsyncMilvusClient,
|
||||
vector_store_id: str,
|
||||
query_vector: Sequence[float],
|
||||
options: _MilvusSearchOptions,
|
||||
params: _MilvusSearchParams,
|
||||
timeout: float | None,
|
||||
) -> object:
|
||||
return await client.search(
|
||||
) -> _MilvusSearchArguments:
|
||||
return _MilvusSearchArguments(
|
||||
collection_name=vector_store_id,
|
||||
data=[list(query_vector)], # mutable-ok: PyMilvus requires nested list search data
|
||||
anns_field=options.anns_field,
|
||||
|
|
@ -387,30 +373,18 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
)
|
||||
query_vector: Final = _EmbeddingPayload.model_validate(embedding_response).vector()
|
||||
connection_timeout, search_timeout = self._timeouts(timeout)
|
||||
if self.sync_client is not None:
|
||||
raw_result: Final = self._sync_search(
|
||||
self.sync_client,
|
||||
vector_store_id,
|
||||
query_vector,
|
||||
options,
|
||||
params,
|
||||
search_timeout,
|
||||
)
|
||||
return self._to_response(raw_result, query_text, params.text_field)
|
||||
|
||||
client: Final = _new_sync_client(params.uri, params.token, params.db_name, connection_timeout)
|
||||
arguments: Final = self._search_arguments(vector_store_id, query_vector, options, params, search_timeout)
|
||||
client: Final = (
|
||||
self.sync_client
|
||||
if self.sync_client is not None
|
||||
else _new_sync_client(params.uri, params.token, params.db_name, connection_timeout)
|
||||
)
|
||||
try:
|
||||
result: Final = self._sync_search(
|
||||
client,
|
||||
vector_store_id,
|
||||
query_vector,
|
||||
options,
|
||||
params,
|
||||
search_timeout,
|
||||
)
|
||||
result: Final = client.search(**arguments)
|
||||
return self._to_response(result, query_text, params.text_field)
|
||||
finally:
|
||||
client.close()
|
||||
if self.sync_client is None:
|
||||
client.close()
|
||||
|
||||
async def aexecute_search_vector_store_request(
|
||||
self,
|
||||
|
|
@ -433,27 +407,15 @@ class MilvusGRPCVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
)
|
||||
query_vector: Final = _EmbeddingPayload.model_validate(embedding_response).vector()
|
||||
connection_timeout, search_timeout = self._timeouts(timeout)
|
||||
if self.async_client is not None:
|
||||
raw_result: Final = await self._async_search(
|
||||
self.async_client,
|
||||
vector_store_id,
|
||||
query_vector,
|
||||
options,
|
||||
params,
|
||||
search_timeout,
|
||||
)
|
||||
return self._to_response(raw_result, query_text, params.text_field)
|
||||
|
||||
client: Final = _new_async_client(params.uri, params.token, params.db_name, connection_timeout)
|
||||
arguments: Final = self._search_arguments(vector_store_id, query_vector, options, params, search_timeout)
|
||||
client: Final = (
|
||||
self.async_client
|
||||
if self.async_client is not None
|
||||
else _new_async_client(params.uri, params.token, params.db_name, connection_timeout)
|
||||
)
|
||||
try:
|
||||
result: Final = await self._async_search(
|
||||
client,
|
||||
vector_store_id,
|
||||
query_vector,
|
||||
options,
|
||||
params,
|
||||
search_timeout,
|
||||
)
|
||||
result: Final = await client.search(**arguments)
|
||||
return self._to_response(result, query_text, params.text_field)
|
||||
finally:
|
||||
await client.close()
|
||||
if self.async_client is None:
|
||||
await client.close()
|
||||
|
|
|
|||
|
|
@ -73,10 +73,10 @@ def build_request_data_from_managed_vector_store(
|
|||
|
||||
|
||||
async def _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data: dict,
|
||||
data: Mapping[str, object],
|
||||
vector_store_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
) -> dict:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Update the request data with the litellm managed vector store registry.
|
||||
|
||||
|
|
@ -92,27 +92,31 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
|
|||
vector_store_id=vector_store_id
|
||||
)
|
||||
if vector_store_to_run is None:
|
||||
data.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None)
|
||||
caller_data: Final = {key: value for key, value in data.items() if key != MILVUS_ADMIN_CONFIGURED_CONNECTION}
|
||||
if user_api_key_dict is not None:
|
||||
assert_proxy_admin_for_user_supplied_vector_store_connection(
|
||||
custom_llm_provider=data.get("custom_llm_provider"),
|
||||
litellm_params=data,
|
||||
custom_llm_provider=caller_data.get("custom_llm_provider"),
|
||||
litellm_params=caller_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return data
|
||||
return caller_data
|
||||
if user_api_key_dict is not None:
|
||||
await assert_user_can_access_vector_store(
|
||||
vector_store=vector_store_to_run,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
if normalize_vector_store_provider(vector_store_to_run.get("custom_llm_provider")) == "milvus":
|
||||
for field in MILVUS_MANAGED_CONFIGURATION_FIELDS:
|
||||
data.pop(field, None)
|
||||
data.pop(MILVUS_ADMIN_CONFIGURED_CONNECTION, None)
|
||||
data.pop("custom_llm_provider", None)
|
||||
data.pop("litellm_credential_name", None)
|
||||
blocked_fields: Final = frozenset(
|
||||
(MILVUS_ADMIN_CONFIGURED_CONNECTION, "custom_llm_provider", "litellm_credential_name")
|
||||
) | (
|
||||
MILVUS_MANAGED_CONFIGURATION_FIELDS
|
||||
if normalize_vector_store_provider(vector_store_to_run.get("custom_llm_provider")) == "milvus"
|
||||
else frozenset()
|
||||
)
|
||||
managed_data: Final = build_request_data_from_managed_vector_store(vector_store_to_run)
|
||||
request_data: Final = {**data, **managed_data} # mutable-ok: request processing requires a mutable payload
|
||||
request_data: Final = {
|
||||
**{key: value for key, value in data.items() if key not in blocked_fields},
|
||||
**managed_data,
|
||||
}
|
||||
if user_api_key_dict is not None:
|
||||
assert_proxy_admin_for_user_supplied_vector_store_connection(
|
||||
custom_llm_provider=request_data.get("custom_llm_provider"),
|
||||
|
|
|
|||
|
|
@ -336,17 +336,10 @@ async def new_vector_store(
|
|||
user_id=user_api_key_dict.user_id,
|
||||
)
|
||||
|
||||
# Apply the same litellm_params redaction the list / info / update
|
||||
# endpoints already use, so a caller-supplied credential or a
|
||||
# cleartext value persisted by an earlier proxy version doesn't
|
||||
# come back in the response.
|
||||
response_vs: Final = LiteLLM_ManagedVectorStore(**new_vector_store)
|
||||
response_vs["litellm_params"] = _redact_sensitive_litellm_params(new_vector_store.get("litellm_params"))
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"message": f"Vector store {vector_store.get('vector_store_id')} created successfully",
|
||||
"vector_store": response_vs,
|
||||
"vector_store": _redact_vector_store(new_vector_store),
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
|
|
@ -398,9 +391,7 @@ def _synchronize_vector_store_registry(
|
|||
|
||||
|
||||
def _redact_vector_store(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_ManagedVectorStore:
|
||||
redacted: Final = LiteLLM_ManagedVectorStore(**vector_store)
|
||||
redacted["litellm_params"] = _redact_sensitive_litellm_params(vector_store.get("litellm_params"))
|
||||
return redacted
|
||||
return {**vector_store, "litellm_params": _redact_sensitive_litellm_params(vector_store.get("litellm_params"))}
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -585,10 +576,11 @@ async def get_vector_store_info(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
vector_store_dict: Final = dict(vector_store_typed)
|
||||
if "litellm_params" in vector_store_dict:
|
||||
vector_store_dict["litellm_params"] = _redact_sensitive_litellm_params(vector_store_dict["litellm_params"])
|
||||
return {"vector_store": vector_store_dict}
|
||||
return {
|
||||
"vector_store": _redact_vector_store(vector_store_typed)
|
||||
if "litellm_params" in vector_store_typed
|
||||
else dict(vector_store_typed)
|
||||
}
|
||||
except HTTPException:
|
||||
# Preserve 403/404 from the access-control / not-found checks above;
|
||||
# the catch-all below would otherwise rewrite them as 500.
|
||||
|
|
@ -683,16 +675,10 @@ async def update_vector_store(
|
|||
"Updated vector store %s in both database and in-memory registry", vector_store_id
|
||||
)
|
||||
|
||||
# The DB row is returned in full, so the response would otherwise
|
||||
# echo the persisted ``litellm_params`` (including provider
|
||||
# credentials) back to the caller — even when the caller only
|
||||
# changed unrelated fields like ``vector_store_description``.
|
||||
response_vs: Final = LiteLLM_ManagedVectorStore(**updated_vs)
|
||||
response_vs["litellm_params"] = _redact_sensitive_litellm_params(updated_vs.get("litellm_params"))
|
||||
return {
|
||||
"status": "success",
|
||||
"message": f"Vector store {vector_store_id} updated successfully",
|
||||
"vector_store": response_vs,
|
||||
"vector_store": _redact_vector_store(updated_vs),
|
||||
}
|
||||
except HTTPException:
|
||||
# Preserve 403/404 responses from the access-control / not-found
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
import json
|
||||
import re
|
||||
from collections.abc import Iterable, Mapping
|
||||
from types import MappingProxyType
|
||||
|
|
@ -25,6 +24,7 @@ from litellm.types.vector_stores import (
|
|||
VectorStoreIndexEndpoints,
|
||||
)
|
||||
from litellm.utils import ProviderConfigManager
|
||||
from litellm.vector_stores.vector_store_registry import deserialize_litellm_params
|
||||
|
||||
MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = frozenset(
|
||||
{
|
||||
|
|
@ -48,13 +48,9 @@ def _normalize_litellm_params(
|
|||
) -> LiteLLM_ManagedVectorStore:
|
||||
litellm_params: Final = vector_store.get("litellm_params")
|
||||
if isinstance(litellm_params, str):
|
||||
normalized: Final = _MANAGED_VECTOR_STORE_ADAPTER.validate_python(vector_store)
|
||||
try:
|
||||
parsed: Final = json.loads(litellm_params)
|
||||
normalized["litellm_params"] = parsed if isinstance(parsed, dict) else {}
|
||||
except (TypeError, ValueError):
|
||||
normalized["litellm_params"] = {}
|
||||
return normalized
|
||||
return _MANAGED_VECTOR_STORE_ADAPTER.validate_python(
|
||||
{**vector_store, "litellm_params": deserialize_litellm_params(litellm_params)}
|
||||
)
|
||||
return vector_store
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -246,7 +246,7 @@
|
|||
"limit": 109
|
||||
},
|
||||
"TRY300": {
|
||||
"limit": 852
|
||||
"limit": 851
|
||||
},
|
||||
"UP028": {
|
||||
"limit": 2
|
||||
|
|
|
|||
|
|
@ -21,6 +21,6 @@
|
|||
"limit": 117
|
||||
},
|
||||
"TQ008": {
|
||||
"limit": 10984
|
||||
"limit": 10981
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
from copy import deepcopy
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -726,19 +727,23 @@ async def test_managed_milvus_uses_only_persisted_connection_for_non_admin():
|
|||
mock_registry = MagicMock()
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store
|
||||
|
||||
payload: Final = {
|
||||
"query": "safe",
|
||||
"custom_llm_provider": "milvus/probe",
|
||||
"milvus_transport": "grpc",
|
||||
"api_base": "http://attacker:19530",
|
||||
"api_key": "attacker-token",
|
||||
"litellm_embedding_model": "openai/attacker-model",
|
||||
"litellm_embedding_config": {"api_base": "http://attacker-embedding"},
|
||||
}
|
||||
original_payload: Final = deepcopy(payload)
|
||||
original_store: Final = deepcopy(managed_vector_store)
|
||||
|
||||
with patch.object( # test-quality-ok: the helper reads the process-wide registry directly
|
||||
litellm, "vector_store_registry", mock_registry
|
||||
):
|
||||
result = await _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data={
|
||||
"query": "safe",
|
||||
"custom_llm_provider": "milvus/probe",
|
||||
"milvus_transport": "grpc",
|
||||
"api_base": "http://attacker:19530",
|
||||
"api_key": "attacker-token",
|
||||
"litellm_embedding_model": "openai/attacker-model",
|
||||
"litellm_embedding_config": {"api_base": "http://attacker-embedding"},
|
||||
},
|
||||
data=payload,
|
||||
vector_store_id="managed",
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
|
|
@ -749,6 +754,9 @@ async def test_managed_milvus_uses_only_persisted_connection_for_non_admin():
|
|||
assert result["litellm_embedding_model"] == "team-embedding-alias"
|
||||
assert "litellm_embedding_config" not in result
|
||||
|
||||
assert payload == original_payload
|
||||
assert managed_vector_store == original_store
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unmarked_managed_milvus_connection_requires_admin_resave():
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -7,7 +8,28 @@ import pytest
|
|||
import litellm
|
||||
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
|
||||
from litellm.vector_stores.main import search
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry, deserialize_litellm_params
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("value", "expected"),
|
||||
[('{"milvus_transport":"grpc"}', {"milvus_transport": "grpc"}),
|
||||
({"milvus_transport": "rest"}, {"milvus_transport": "rest"}),
|
||||
("invalid-json", {}), ("[]", {}), ("null", {}), (None, {})],
|
||||
)
|
||||
def test_deserialize_litellm_params(value: object, expected: dict[str, object]) -> None:
|
||||
assert deserialize_litellm_params(value) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize("serialized_params", ['{"milvus_transport":"grpc"}', "invalid-json", "[]", "null"])
|
||||
def test_normalizing_stored_params_preserves_source(serialized_params: str) -> None:
|
||||
from litellm.proxy.vector_store_endpoints.utils import _normalize_litellm_params
|
||||
|
||||
store: Final = {"vector_store_id": "documents", "custom_llm_provider": "milvus", "litellm_params": serialized_params}
|
||||
result: Final = _normalize_litellm_params(store)
|
||||
|
||||
assert result["litellm_params"] == deserialize_litellm_params(serialized_params)
|
||||
assert store["litellm_params"] == serialized_params
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
Tests for Milvus Vector Store
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from typing import Final, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -116,6 +117,7 @@ class TestMilvusVectorStore:
|
|||
({"max_num_results": 1}, 1),
|
||||
({"max_num_results": 50}, 50),
|
||||
({"max_num_results": 2, "limit": 7}, 2),
|
||||
({"max_num_results": None, "limit": 75}, 75),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("async_mode", [False, True])
|
||||
|
|
@ -127,6 +129,7 @@ class TestMilvusVectorStore:
|
|||
executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
|
||||
config: Final = MilvusVectorStoreConfig()
|
||||
original_params: Final = optional_params.copy()
|
||||
kwargs: Final = {
|
||||
"vector_store_id": "documents",
|
||||
"query": "limit probe",
|
||||
|
|
@ -145,6 +148,7 @@ class TestMilvusVectorStore:
|
|||
assert body.get("limit") == expected_limit
|
||||
assert "max_num_results" not in body
|
||||
assert body["collectionName"] == "documents"
|
||||
assert optional_params == original_params
|
||||
|
||||
@pytest.mark.parametrize("max_num_results", [0, 51])
|
||||
@pytest.mark.parametrize("async_mode", [False, True])
|
||||
|
|
@ -561,6 +565,7 @@ class TestMilvusVectorStore:
|
|||
"attributes": {"category": "reference"},
|
||||
}
|
||||
]
|
||||
mock_client.close.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_grpc_search_infers_vector_field_and_requests_text_by_default(self):
|
||||
|
|
@ -587,6 +592,7 @@ class TestMilvusVectorStore:
|
|||
VectorStoreSearchOptionalRequestParams,
|
||||
{
|
||||
"max_num_results": 2,
|
||||
"limit": 7,
|
||||
},
|
||||
),
|
||||
litellm_logging_obj=MagicMock(),
|
||||
|
|
@ -607,6 +613,52 @@ class TestMilvusVectorStore:
|
|||
assert mock_client.search.await_args.kwargs["anns_field"] is None
|
||||
assert mock_client.search.await_args.kwargs["output_fields"] == ["book_intro_text"]
|
||||
assert response["data"][0]["content"][0]["text"] == "async result"
|
||||
mock_client.close.assert_not_called()
|
||||
|
||||
@pytest.mark.parametrize("injected", [False, True])
|
||||
@pytest.mark.parametrize("async_mode", [False, True])
|
||||
@pytest.mark.parametrize("failure", ["search", "response", "cancellation"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_grpc_client_ownership_after_failure(
|
||||
self, injected: bool, async_mode: bool, failure: str
|
||||
) -> None:
|
||||
client: Final = MagicMock()
|
||||
executor: Final = MagicMock()
|
||||
executor.embed.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
executor.aembed = AsyncMock(return_value=MOCK_EMBEDDING_RESPONSE)
|
||||
error: Final = asyncio.CancelledError if failure == "cancellation" else RuntimeError
|
||||
search: Final = AsyncMock() if async_mode else MagicMock()
|
||||
search.side_effect = None if failure == "response" else error("search interrupted")
|
||||
search.return_value = "invalid result"
|
||||
client.search = search
|
||||
client.close = AsyncMock() if async_mode else MagicMock()
|
||||
config: Final = MilvusGRPCVectorStoreConfig(
|
||||
sync_client=client if injected and not async_mode else None,
|
||||
async_client=client if injected and async_mode else None,
|
||||
)
|
||||
kwargs: Final = {
|
||||
"vector_store_id": "documents",
|
||||
"query": "cleanup probe",
|
||||
"vector_store_search_optional_params": {},
|
||||
"litellm_logging_obj": MagicMock(),
|
||||
"litellm_params": {
|
||||
"api_base": "http://milvus:19530",
|
||||
"litellm_embedding_model": "embedding-alias",
|
||||
},
|
||||
"embedding_executor": executor,
|
||||
}
|
||||
with (
|
||||
patch("pymilvus.AsyncMilvusClient" if async_mode else "pymilvus.MilvusClient", return_value=client),
|
||||
pytest.raises(TypeError if failure == "response" else error),
|
||||
):
|
||||
(
|
||||
await config.aexecute_search_vector_store_request(**kwargs)
|
||||
if async_mode else config.execute_search_vector_store_request(**kwargs)
|
||||
)
|
||||
|
||||
assert client.close.call_count == (0 if injected else 1)
|
||||
if async_mode:
|
||||
assert client.close.await_count == (0 if injected else 1)
|
||||
|
||||
def test_grpc_search_always_requests_configured_text_field(self):
|
||||
mock_client = MagicMock()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22165
|
||||
"limit": 22156
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26733
|
||||
|
|
@ -27,7 +27,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16442
|
||||
"limit": 16431
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5506
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue