diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index 6f6822a975b..35ba0783d64 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -119,6 +119,11 @@ jobs: .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml --extra mongodb uv run --no-sync python -c 'import os, sys; print(sys.version); assert f"{sys.version_info.major}.{sys.version_info.minor}" == os.environ["UV_PYTHON"]' + - name: Install Milvus SDK for provider tests + if: steps.changes.outputs.decision != 'skip' && inputs.artifact-name == 'llm-other-providers' + timeout-minutes: 8 + run: .github/scripts/uv_sync_with_retries.sh --inexact --frozen --extra milvus + - name: Cache Prisma binaries if: steps.changes.outputs.decision != 'skip' timeout-minutes: 3 diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index b1158c63ef6..fe7d6810755 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,9 +1,9 @@ { "reportAny": { - "limit": 13428 + "limit": 13426 }, "reportArgumentType": { - "limit": 2168 + "limit": 2158 }, "reportAssignmentType": { "limit": 319 @@ -24,7 +24,7 @@ "limit": 19 }, "reportExplicitAny": { - "limit": 3366 + "limit": 3365 }, "reportFunctionMemberAccess": { "limit": 7 @@ -57,7 +57,7 @@ "limit": 5570 }, "reportMissingTypeArgument": { - "limit": 15264 + "limit": 15258 }, "reportMissingTypeStubs": { "limit": 40 @@ -99,19 +99,19 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 43218 + "limit": 42997 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 38239 + "limit": 38223 }, "reportUnknownParameterType": { - "limit": 19582 + "limit": 19580 }, "reportUnknownVariableType": { - "limit": 29823 + "limit": 29820 }, "reportUnnecessaryCast": { "limit": 110 @@ -123,7 +123,7 @@ "limit": 4 }, "reportUnnecessaryIsInstance": { - "limit": 813 + "limit": 812 }, "reportUntypedBaseClass": { "limit": 0 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 390e1e4fbd5..bb5dd77f3c1 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 @@ -128,6 +128,23 @@ class VectorStorePreCallHook(CustomLogger): verbose_logger.debug("No query found in messages for vector store search") return model, messages, non_default_params + except Exception as e: # noqa: BLE001 # registry and prompt setup failures retain the existing chat fallback + verbose_logger.exception("Error in VectorStorePreCallHook: %s", e) + return model, messages, non_default_params + + 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, + ) + + for vector_store_to_validate in vector_stores_to_run: + assert_proxy_admin_for_user_supplied_vector_store_connection( + custom_llm_provider=vector_store_to_validate.get("custom_llm_provider"), + litellm_params=vector_store_to_validate.get("litellm_params"), + managed=True, + ) + + try: modified_messages: list[AllMessageValues] = messages.copy() all_search_results: Final[list[VectorStoreSearchResponse]] = [] @@ -151,18 +168,6 @@ 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, @@ -173,8 +178,6 @@ class VectorStorePreCallHook(CustomLogger): }, ) except Exception as search_error: - if getattr(search_error, "status_code", None) in (401, 403): - raise verbose_logger.warning( "Vector store search failed for vector_store_id=%s, continuing without its context: %s", vector_store_id, @@ -204,8 +207,6 @@ class VectorStorePreCallHook(CustomLogger): return model, modified_messages, non_default_params except Exception as e: - if getattr(e, "status_code", None) in (401, 403): - raise verbose_logger.exception("Error in VectorStorePreCallHook: %s", e) # Return original parameters on error return model, messages, non_default_params diff --git a/litellm/llms/milvus/vector_stores/connection.py b/litellm/llms/milvus/vector_stores/connection.py new file mode 100644 index 00000000000..920576c78ac --- /dev/null +++ b/litellm/llms/milvus/vector_stores/connection.py @@ -0,0 +1,130 @@ +from collections.abc import Mapping +from enum import Enum +from types import MappingProxyType +from typing import Final + +import litellm +from litellm.types.vector_stores import MILVUS_ADMIN_CONFIGURED_CONNECTION + +MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = frozenset( + { + "api_base", + "api_key", + "custom_llm_provider", + "litellm_credential_name", + "milvus_transport", + "milvus_db_name", + "milvus_partition_names", + "litellm_embedding_config", + "litellm_embedding_model", + "milvus_text_field", + } +) + + +class MilvusConnectionRejection(Enum): + ADMIN_REQUIRED = "Only proxy admins can configure vector store connections. Contact your LiteLLM administrator." + ADMIN_SAVE_REQUIRED = "This managed Milvus gRPC connection must be re-saved by a proxy admin before it can be used." + + +def _normalize_provider(custom_llm_provider: object) -> str | None: + if not isinstance(custom_llm_provider, str) or not custom_llm_provider: + return None + if "/" not in custom_llm_provider: + return custom_llm_provider + try: + _, provider, _, _ = litellm.get_llm_provider(model=custom_llm_provider) + return provider + except Exception: # noqa: BLE001 # provider parsing failures fall back to the explicit prefix + return custom_llm_provider.split("/", 1)[0] + + +def _is_grpc_connection(custom_llm_provider: object, litellm_params: object) -> bool: + return ( + isinstance(litellm_params, dict) + and ( + _normalize_provider(custom_llm_provider) == "milvus" + or _normalize_provider(litellm_params.get("custom_llm_provider")) == "milvus" + ) + and litellm_params.get("milvus_transport") == "grpc" + ) + + +def connection_rejection( + custom_llm_provider: object, + litellm_params: object, + *, + is_proxy_admin: bool, + managed: bool, +) -> MilvusConnectionRejection | None: + if not _is_grpc_connection(custom_llm_provider, litellm_params): + return None + if managed: + if isinstance(litellm_params, dict) and litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True: + return None + return MilvusConnectionRejection.ADMIN_SAVE_REQUIRED + if is_proxy_admin: + return None + return MilvusConnectionRejection.ADMIN_REQUIRED + + +def prepare_connection_for_persistence( + *, + custom_llm_provider: object, + litellm_params: object, + is_proxy_admin: bool, + existing_custom_llm_provider: object | None = None, + existing_litellm_params: object | None = None, + litellm_credential_name: object | None = None, + existing_litellm_credential_name: object | None = None, + litellm_credential_name_supplied: bool = False, +) -> dict[str, object] | MilvusConnectionRejection: # mutable-ok: persistence requires a serializable connection dict + existing: Final = existing_litellm_params if isinstance(existing_litellm_params, dict) else MappingProxyType({}) + supplied: Final = litellm_params if isinstance(litellm_params, dict) else MappingProxyType({}) + effective: Final = { # mutable-ok: the validated connection must be JSON-serializable for database persistence + key: value + for params in (existing, supplied) + for key, value in params.items() + if key != MILVUS_ADMIN_CONFIGURED_CONNECTION + } + previous_is_grpc: Final = _is_grpc_connection(existing_custom_llm_provider, existing) + effective_is_grpc: Final = _is_grpc_connection(custom_llm_provider, effective) + if not previous_is_grpc and not effective_is_grpc: + return ( # mutable-ok: persistence requires an isolated JSON-serializable dict + dict(supplied) if isinstance(litellm_params, dict) else dict(existing) + ) + + is_create: Final = existing_custom_llm_provider is None + provider_changed: Final = not is_create and custom_llm_provider != existing_custom_llm_provider + managed_configuration_changed: Final = any( + existing.get(field) != effective.get(field) for field in MILVUS_MANAGED_CONFIGURATION_FIELDS + ) + credential_changed: Final = litellm_credential_name_supplied and ( + litellm_credential_name != existing_litellm_credential_name + ) + missing_marker: Final = effective_is_grpc and existing.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is not True + + if ( + is_create or provider_changed or managed_configuration_changed or credential_changed or missing_marker + ) and not is_proxy_admin: + return MilvusConnectionRejection.ADMIN_REQUIRED + + return ( + {**effective, MILVUS_ADMIN_CONFIGURED_CONNECTION: True} # mutable-ok: persisted JSON carries the server marker + if effective_is_grpc + else effective + ) + + +def managed_connection_fields(custom_llm_provider: object) -> frozenset[str]: + return frozenset((MILVUS_ADMIN_CONFIGURED_CONNECTION, "custom_llm_provider", "litellm_credential_name")) | ( + MILVUS_MANAGED_CONFIGURATION_FIELDS if _normalize_provider(custom_llm_provider) == "milvus" else frozenset() + ) + + +def approve_configured_connection( + custom_llm_provider: object, litellm_params: Mapping[str, object] +) -> Mapping[str, object]: + if custom_llm_provider == "milvus" and litellm_params.get("milvus_transport") == "grpc": + return MappingProxyType({**litellm_params, MILVUS_ADMIN_CONFIGURED_CONNECTION: True}) + return litellm_params diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 724322f3dcd..b5fe38d2af6 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -12,21 +12,19 @@ from fastapi import APIRouter, Depends, HTTPException, Request, Response from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( LiteLLM_ManagedVectorStore, ) +from litellm.llms.milvus.vector_stores.connection import managed_connection_fields from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.utils import jsonify_object from litellm.proxy.vector_store_endpoints.utils import ( - MILVUS_ADMIN_CONFIGURED_CONNECTION, - MILVUS_MANAGED_CONFIGURATION_FIELDS, assert_proxy_admin_for_user_supplied_vector_store_connection, assert_proxy_admin_for_vector_store_index_management, assert_user_can_access_vector_store, get_litellm_managed_vector_store, - normalize_vector_store_provider, ) from litellm.repositories.table_repositories import ManagedVectorStoreIndexRepository -from litellm.types.vector_stores import IndexCreateRequest, IndexListResponse +from litellm.types.vector_stores import MILVUS_ADMIN_CONFIGURED_CONNECTION, IndexCreateRequest, IndexListResponse from litellm.vector_stores.vector_store_registry import VectorStoreIndexRegistry router: Final = APIRouter() @@ -105,13 +103,7 @@ async def _update_request_data_with_litellm_managed_vector_store_registry( vector_store=vector_store_to_run, user_api_key_dict=user_api_key_dict, ) - 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() - ) + blocked_fields: Final = managed_connection_fields(vector_store_to_run.get("custom_llm_provider")) managed_data: Final = build_request_data_from_managed_vector_store(vector_store_to_run) request_data: Final = { **{key: value for key, value in data.items() if key not in blocked_fields}, diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 2fe2fbb9009..97ff84adeac 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -31,14 +31,14 @@ from litellm.proxy._types import ( from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user from litellm.proxy.vector_store_endpoints.utils import ( - MILVUS_ADMIN_CONFIGURED_CONNECTION, can_user_access_vector_store, filter_listable_vector_stores, - prepare_milvus_connection_for_persistence, + prepare_vector_store_connection_for_persistence, ) from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ManagedVectorStoresRepository from litellm.types.vector_stores import ( + MILVUS_ADMIN_CONFIGURED_CONNECTION, LiteLLM_ManagedVectorStore, LiteLLM_ManagedVectorStoreListResponse, VectorStoreDeleteRequest, @@ -311,7 +311,7 @@ async def new_vector_store( detail="vector_store_id and custom_llm_provider are required", ) - prepared_litellm_params: Final = prepare_milvus_connection_for_persistence( + prepared_litellm_params: Final = prepare_vector_store_connection_for_persistence( custom_llm_provider=custom_llm_provider, litellm_params=vector_store.get("litellm_params"), user_api_key_dict=user_api_key_dict, @@ -628,7 +628,7 @@ async def update_vector_store( effective_provider: Final = update_data.get("custom_llm_provider") or existing_vector_store.get( "custom_llm_provider" ) - effective_litellm_params: Final = prepare_milvus_connection_for_persistence( + effective_litellm_params: Final = prepare_vector_store_connection_for_persistence( custom_llm_provider=effective_provider, litellm_params=update_data.get("litellm_params"), user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index c0b1ac5e722..8814f37d556 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -8,6 +8,11 @@ from pydantic import TypeAdapter import litellm from litellm._logging import verbose_proxy_logger +from litellm.llms.milvus.vector_stores.connection import ( + MilvusConnectionRejection, + connection_rejection, + prepare_connection_for_persistence, +) from litellm.proxy._experimental.mcp_server.ui_session_utils import ( is_ui_session_credential, resolve_ui_session_team_ids, @@ -19,27 +24,12 @@ from litellm.proxy._types import ( ) from litellm.types.utils import LlmProviders from litellm.types.vector_stores import ( - MILVUS_ADMIN_CONFIGURED_CONNECTION, LiteLLM_ManagedVectorStore, VectorStoreIndexEndpoints, ) from litellm.utils import ProviderConfigManager from litellm.vector_stores.vector_store_registry import deserialize_litellm_params -MILVUS_MANAGED_CONFIGURATION_FIELDS: Final = frozenset( - { - "api_base", - "api_key", - "custom_llm_provider", - "litellm_credential_name", - "milvus_transport", - "milvus_db_name", - "milvus_partition_names", - "litellm_embedding_config", - "litellm_embedding_model", - "milvus_text_field", - } -) _MANAGED_VECTOR_STORE_ADAPTER: Final = TypeAdapter(LiteLLM_ManagedVectorStore) @@ -61,29 +51,6 @@ def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: ) -def normalize_vector_store_provider(custom_llm_provider: object) -> str | None: - if not isinstance(custom_llm_provider, str) or not custom_llm_provider: - return None - if "/" not in custom_llm_provider: - return custom_llm_provider - try: - _, provider, _, _ = litellm.get_llm_provider(model=custom_llm_provider) - return provider - except Exception: # noqa: BLE001 # provider parsing failures fall back to the explicit prefix - return custom_llm_provider.split("/", 1)[0] - - -def is_milvus_grpc_connection(custom_llm_provider: object, litellm_params: object) -> bool: - return ( - isinstance(litellm_params, dict) - and ( - normalize_vector_store_provider(custom_llm_provider) == "milvus" - or normalize_vector_store_provider(litellm_params.get("custom_llm_provider")) == "milvus" - ) - and litellm_params.get("milvus_transport") == "grpc" - ) - - def assert_proxy_admin_for_vector_store_index_management( user_api_key_dict: UserAPIKeyAuth, *, @@ -105,24 +72,17 @@ def assert_proxy_admin_for_user_supplied_vector_store_connection( *, managed: bool = False, ) -> None: - if not is_milvus_grpc_connection(custom_llm_provider, litellm_params): - return - if managed: - if isinstance(litellm_params, dict) and litellm_params.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is True: - return - raise HTTPException( - status_code=403, - detail="This managed Milvus gRPC connection must be re-saved by a proxy admin before it can be used.", - ) - if user_api_key_dict is not None and _is_proxy_admin(user_api_key_dict): - return - raise HTTPException( - status_code=403, - detail="Only proxy admins can configure vector store connections. Contact your LiteLLM administrator.", + rejection: Final = connection_rejection( + custom_llm_provider=custom_llm_provider, + litellm_params=litellm_params, + is_proxy_admin=user_api_key_dict is not None and _is_proxy_admin(user_api_key_dict), + managed=managed, ) + if rejection is not None: + raise HTTPException(status_code=403, detail=rejection.value) -def prepare_milvus_connection_for_persistence( +def prepare_vector_store_connection_for_persistence( *, custom_llm_provider: object, litellm_params: object, @@ -133,44 +93,19 @@ def prepare_milvus_connection_for_persistence( existing_litellm_credential_name: object | None = None, litellm_credential_name_supplied: bool = False, ) -> dict[str, object]: # mutable-ok: persistence requires a serializable effective-connection dict - existing: Final = existing_litellm_params if isinstance(existing_litellm_params, dict) else MappingProxyType({}) - supplied: Final = litellm_params if isinstance(litellm_params, dict) else MappingProxyType({}) - effective: Final = { # mutable-ok: the validated connection must be JSON-serializable for database persistence - key: value - for params in (existing, supplied) - for key, value in params.items() - if key != MILVUS_ADMIN_CONFIGURED_CONNECTION - } - previous_is_grpc: Final = is_milvus_grpc_connection(existing_custom_llm_provider, existing) - effective_is_grpc: Final = is_milvus_grpc_connection(custom_llm_provider, effective) - if not previous_is_grpc and not effective_is_grpc: - return ( # mutable-ok: persistence requires an isolated JSON-serializable dict - dict(supplied) if isinstance(litellm_params, dict) else dict(existing) - ) - - is_create: Final = existing_custom_llm_provider is None - provider_changed: Final = not is_create and custom_llm_provider != existing_custom_llm_provider - managed_configuration_changed: Final = any( - existing.get(field) != effective.get(field) for field in MILVUS_MANAGED_CONFIGURATION_FIELDS - ) - credential_changed: Final = litellm_credential_name_supplied and ( - litellm_credential_name != existing_litellm_credential_name - ) - missing_marker: Final = effective_is_grpc and existing.get(MILVUS_ADMIN_CONFIGURED_CONNECTION) is not True - - if ( - is_create or provider_changed or managed_configuration_changed or credential_changed or missing_marker - ) and not _is_proxy_admin(user_api_key_dict): - raise HTTPException( - status_code=403, - detail="Only proxy admins can configure vector store connections. Contact your LiteLLM administrator.", - ) - - return ( - {**effective, MILVUS_ADMIN_CONFIGURED_CONNECTION: True} # mutable-ok: persisted JSON carries the server marker - if effective_is_grpc - else effective + result: Final = prepare_connection_for_persistence( + custom_llm_provider=custom_llm_provider, + litellm_params=litellm_params, + is_proxy_admin=_is_proxy_admin(user_api_key_dict), + existing_custom_llm_provider=existing_custom_llm_provider, + existing_litellm_params=existing_litellm_params, + litellm_credential_name=litellm_credential_name, + existing_litellm_credential_name=existing_litellm_credential_name, + litellm_credential_name_supplied=litellm_credential_name_supplied, ) + if isinstance(result, MilvusConnectionRejection): + raise HTTPException(status_code=403, detail=result.value) + return result def _suffix_after_index_name(request_path: str, index_name: str) -> str | None: diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index 4476999ee83..20c72f97f4e 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -14,12 +14,12 @@ 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.llms.milvus.vector_stores.connection import approve_configured_connection from litellm.repositories.table_repositories import ( ManagedVectorStoreIndexRepository, ManagedVectorStoresRepository, ) from litellm.types.vector_stores import ( - MILVUS_ADMIN_CONFIGURED_CONNECTION, VECTOR_STORE_OPENAI_PARAMS, LiteLLM_ManagedVectorStore, LiteLLM_ManagedVectorStoreIndex, @@ -459,14 +459,11 @@ class VectorStoreRegistry: f"custom_llm_provider is required for initializing vector store, got custom_llm_provider={custom_llm_provider}" ) - if custom_llm_provider == "milvus" and vector_store_litellm_params.get("milvus_transport") == "grpc": - vector_store_litellm_params[MILVUS_ADMIN_CONFIGURED_CONNECTION] = True - litellm_managed_vector_store = _MANAGED_VECTOR_STORE_ADAPTER.validate_python( { # mutable-ok: Pydantic validates the config mapping into a managed vector store "vector_store_id": vector_store_id, "custom_llm_provider": custom_llm_provider, - "litellm_params": vector_store_litellm_params, + "litellm_params": approve_configured_connection(custom_llm_provider, vector_store_litellm_params), "vector_store_name": vector_store_name, "vector_store_description": vector_store_litellm_params.get("vector_store_description"), "vector_store_metadata": vector_store_litellm_params.get("vector_store_metadata"), diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 89d75dd1a02..8ffbc3f5490 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -246,7 +246,7 @@ "limit": 109 }, "TRY300": { - "limit": 851 + "limit": 850 }, "UP028": { "limit": 2 diff --git a/test-quality-budget.json b/test-quality-budget.json index 9b26a1765ea..5fb3241dabf 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -21,6 +21,6 @@ "limit": 117 }, "TQ008": { - "limit": 10981 + "limit": 10978 } } 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 e946f6fc4b1..0bd8db38d93 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 @@ -45,10 +45,10 @@ class RecordingRouter: async def avector_store_search(self, **kwargs: object) -> VectorStoreSearchResponse: self.calls.append(kwargs) - if self.search_error is not None: - raise self.search_error vector_store_id = str(kwargs["vector_store_id"]) if vector_store_id in self.failing_vector_store_ids: + if self.search_error is not None: + raise self.search_error raise litellm.BadRequestError( message=f"no healthy deployments for {vector_store_id}", model="text-embedding-3-small", @@ -194,32 +194,35 @@ async def test_hook_rejects_an_untrusted_managed_milvus_grpc_connection( assert exc_info.value.detail == ( "This managed Milvus gRPC connection must be re-saved by a proxy admin before it can be used." ) - assert [call["vector_store_id"] for call in router.calls] == (["safe"] if vector_store_ids[0] == "safe" else []) + assert router.calls == [] assert "search_results" not in logging_obj.model_call_details assert warnings == [] @pytest.mark.asyncio @pytest.mark.parametrize("status_code", [401, 403]) -async def test_hook_propagates_search_authorization_errors( +async def test_hook_preserves_healthy_context_after_provider_authorization_errors( registry_with: RegisterStores, warnings: list[logging.LogRecord], status_code: int, ) -> None: registry_with("vs-denied", "vs-safe") error: Final = HTTPException(status_code=status_code, detail="Vector store access denied") - router: Final = RecordingRouter(search_error=error) + router: Final = RecordingRouter(failing_vector_store_ids=frozenset({"vs-denied"}), search_error=error) + logging_obj: Final = FakeLoggingObj({}) - with pytest.raises(HTTPException) as exc_info: - await _run_hook( - VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), - ["vs-denied", "vs-safe"], - FakeLoggingObj({}), - ) + _, messages, _ = await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), + ["vs-denied", "vs-safe"], + logging_obj, + ) - assert exc_info.value is error - assert [call["vector_store_id"] for call in router.calls] == ["vs-denied"] - assert warnings == [] + assert [call["vector_store_id"] for call in router.calls] == ["vs-denied", "vs-safe"] + assert messages[0]["content"] == "Context:\n\ncontext from vs-safe\n\n" + assert logging_obj.model_call_details["search_results"] == [_search_response("context from vs-safe")] + assert [record.getMessage() for record in warnings] == [ + f"Vector store search failed for vector_store_id=vs-denied, continuing without its context: {error}" + ] @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 9820052c37c..7f197b90200 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 @@ -28,18 +28,17 @@ from litellm.proxy.vector_store_endpoints.management_endpoints import ( new_vector_store, ) from litellm.proxy.vector_store_endpoints.utils import ( - MILVUS_ADMIN_CONFIGURED_CONNECTION, assert_proxy_admin_for_user_supplied_vector_store_connection, check_vector_store_permission, is_allowed_to_call_vector_store_endpoint, is_allowed_to_call_vector_store_files_endpoint, - prepare_milvus_connection_for_persistence, + prepare_vector_store_connection_for_persistence, ) from litellm.proxy.vector_store_files_endpoints.endpoints import ( _update_request_data_with_model_routing_hint, ) from litellm.types.utils import EmbeddingResponse, LlmProviders -from litellm.types.vector_stores import IndexCreateRequest, IndexListResponse +from litellm.types.vector_stores import MILVUS_ADMIN_CONFIGURED_CONNECTION, IndexCreateRequest, IndexListResponse from litellm.vector_stores.main import _direct_vector_store_embedding_executor from litellm.vector_stores.vector_store_registry import VectorStoreRegistry @@ -1013,7 +1012,7 @@ def test_config_vector_store_cannot_be_replaced_or_deleted_from_registry(): def test_admin_persistence_strips_forged_marker_and_adds_server_marker(): - params = prepare_milvus_connection_for_persistence( + params = prepare_vector_store_connection_for_persistence( custom_llm_provider="milvus/probe", litellm_params={ "milvus_transport": "grpc", @@ -1034,7 +1033,7 @@ def test_nested_provider_cannot_bypass_milvus_grpc_registration_authorization( provider: str, nested_provider: str ) -> None: with pytest.raises(HTTPException) as exc_info: - prepare_milvus_connection_for_persistence( + prepare_vector_store_connection_for_persistence( custom_llm_provider=provider, litellm_params={"custom_llm_provider": nested_provider, "milvus_transport": "grpc"}, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), @@ -1044,7 +1043,7 @@ def test_nested_provider_cannot_bypass_milvus_grpc_registration_authorization( def test_non_grpc_connection_update_keeps_replacement_semantics(): - params = prepare_milvus_connection_for_persistence( + params = prepare_vector_store_connection_for_persistence( custom_llm_provider="openai", litellm_params={"api_key": "new-key"}, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 74bdcf77652..7ef25d7080a 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22156 + "limit": 22147 }, "LIT002": { "limit": 26733 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16431 + "limit": 16420 }, "LIT011": { "limit": 5506