diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 0db61149249..b1cbc330cfb 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -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 diff --git a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py index e789c1d800a..13494f48cf3 100644 --- a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py +++ b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py @@ -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, diff --git a/litellm/llms/azure/passthrough/transformation.py b/litellm/llms/azure/passthrough/transformation.py index 898852e645f..e3622490579 100644 --- a/litellm/llms/azure/passthrough/transformation.py +++ b/litellm/llms/azure/passthrough/transformation.py @@ -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 diff --git a/litellm/llms/milvus/vector_stores/grpc_transformation.py b/litellm/llms/milvus/vector_stores/grpc_transformation.py index a3c3eecd2ff..519b7ac23ea 100644 --- a/litellm/llms/milvus/vector_stores/grpc_transformation.py +++ b/litellm/llms/milvus/vector_stores/grpc_transformation.py @@ -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) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 552d1ea434f..e2756cd7c0c 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -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,) ) diff --git a/litellm/router.py b/litellm/router.py index 6c7611c6236..4e7555b2aa6 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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( diff --git a/litellm/types/router.py b/litellm/types/router.py index 4135917cf12..585bba42668 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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: diff --git a/litellm/types/vector_stores.py b/litellm/types/vector_stores.py index deba5b72302..2ecf731d05e 100644 --- a/litellm/types/vector_stores.py +++ b/litellm/types/vector_stores.py @@ -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 diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index 34f0b3d471f..ccd25f775cc 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -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 diff --git a/test-quality-budget.json b/test-quality-budget.json index 8f0ebb1ea92..4c4fccc1c86 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -21,6 +21,6 @@ "limit": 117 }, "TQ008": { - "limit": 11000 + "limit": 10997 } } diff --git a/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py b/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py index f150df7254e..26d769b72a2 100644 --- a/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py +++ b/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py @@ -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 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 9a1097d57c9..32e4c4377df 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 @@ -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 diff --git a/tests/test_litellm/vector_stores/test_vector_store_registry.py b/tests/test_litellm/vector_stores/test_vector_store_registry.py index f19c3706845..1b27d27a7bb 100644 --- a/tests/test_litellm/vector_stores/test_vector_store_registry.py +++ b/tests/test_litellm/vector_stores/test_vector_store_registry.py @@ -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(): """ diff --git a/tests/vector_store_tests/test_milvus_vector_store.py b/tests/vector_store_tests/test_milvus_vector_store.py index 4ad0ea2b37c..0fc67a16b53 100644 --- a/tests/vector_store_tests/test_milvus_vector_store.py +++ b/tests/vector_store_tests/test_milvus_vector_store.py @@ -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): diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index e0184aab8ba..304ecfaf963 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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 */