diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index fe7d6810755..cd4f897cd1a 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,9 +1,9 @@ { "reportAny": { - "limit": 13426 + "limit": 13424 }, "reportArgumentType": { - "limit": 2158 + "limit": 2148 }, "reportAssignmentType": { "limit": 319 @@ -24,7 +24,7 @@ "limit": 19 }, "reportExplicitAny": { - "limit": 3365 + "limit": 3364 }, "reportFunctionMemberAccess": { "limit": 7 @@ -57,7 +57,7 @@ "limit": 5570 }, "reportMissingTypeArgument": { - "limit": 15258 + "limit": 15252 }, "reportMissingTypeStubs": { "limit": 40 @@ -99,19 +99,19 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 42997 + "limit": 42776 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 38223 + "limit": 38207 }, "reportUnknownParameterType": { - "limit": 19580 + "limit": 19578 }, "reportUnknownVariableType": { - "limit": 29820 + "limit": 29817 }, "reportUnnecessaryCast": { "limit": 110 @@ -123,7 +123,7 @@ "limit": 4 }, "reportUnnecessaryIsInstance": { - "limit": 812 + "limit": 811 }, "reportUntypedBaseClass": { "limit": 0 diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index 20c72f97f4e..542d06ed772 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -2,6 +2,7 @@ import json from collections.abc import Mapping from datetime import datetime, timezone +from types import MappingProxyType from typing import ( TYPE_CHECKING, Any, # noqa: TID251 # untyped non_default_params dict is the only source of the unknown key type @@ -48,6 +49,14 @@ def deserialize_litellm_params( return {} # mutable-ok: managed vector store rows expose JSON objects as dicts +def _normalized_vector_store(vector_store: LiteLLM_ManagedVectorStore) -> LiteLLM_ManagedVectorStore: + return _MANAGED_VECTOR_STORE_ADAPTER.validate_python( + MappingProxyType( + {**vector_store, "litellm_params": deserialize_litellm_params(vector_store.get("litellm_params"))} + ) + ) + + class VectorStoreIndexRegistry: def __init__(self, vector_store_indexes: list[LiteLLM_ManagedVectorStoreIndex] = []): self.vector_store_indexes: list[LiteLLM_ManagedVectorStoreIndex] = vector_store_indexes @@ -361,7 +370,7 @@ class VectorStoreRegistry: for vector_store in self.vector_stores: if vector_store.get("vector_store_id") == vector_store_id: # Create a copy to avoid modifying the registry - vector_store_copy = vector_store.copy() + vector_store_copy = _normalized_vector_store(vector_store) # Merge tool params if they exist if vector_store_id in params_by_id: @@ -410,7 +419,7 @@ class VectorStoreRegistry: if vector_store is not None: # Create a copy to avoid modifying the registry - vector_store_copy = vector_store.copy() + vector_store_copy = _normalized_vector_store(vector_store) # Merge tool params if they exist if vector_store_id in params_by_id: diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 8ffbc3f5490..aa546ba7f12 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -246,7 +246,7 @@ "limit": 109 }, "TRY300": { - "limit": 850 + "limit": 849 }, "UP028": { "limit": 2 diff --git a/test-quality-budget.json b/test-quality-budget.json index 5fb3241dabf..2a1213d42b6 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -21,6 +21,6 @@ "limit": 117 }, "TQ008": { - "limit": 10978 + "limit": 10975 } } 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 0bd8db38d93..8b03aa8b0b6 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 @@ -3,8 +3,8 @@ from collections.abc import Iterator from dataclasses import dataclass, field from typing import Final, Protocol +import httpx import pytest -from fastapi import HTTPException import litellm from litellm._logging import verbose_logger @@ -153,52 +153,6 @@ async def test_hook_searches_through_the_injected_router_with_the_request_metada assert messages[0]["content"] == "Context:\n\ncontext from vs-router\n\n" -@pytest.mark.asyncio -@pytest.mark.parametrize("vector_store_ids", [("legacy",), ("legacy", "safe"), ("safe", "legacy")]) -async def test_hook_rejects_an_untrusted_managed_milvus_grpc_connection( - monkeypatch: pytest.MonkeyPatch, - vector_store_ids: tuple[str, ...], - warnings: list[logging.LogRecord], -) -> None: - monkeypatch.setattr( - litellm, - "vector_store_registry", - VectorStoreRegistry( - vector_stores=[ - LiteLLM_ManagedVectorStore( - vector_store_id="legacy", - custom_llm_provider="milvus", - litellm_params={ - "milvus_transport": "grpc", - "api_base": "http://internal-milvus:19530", - }, - ), - LiteLLM_ManagedVectorStore( - vector_store_id="safe", - custom_llm_provider="bedrock", - ), - ], - ), - ) - router: Final = RecordingRouter() - logging_obj: Final = FakeLoggingObj({}) - - with pytest.raises(HTTPException) as exc_info: - await _run_hook( - VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)), - list(vector_store_ids), - logging_obj, - ) - - assert exc_info.value.status_code == 403 - assert exc_info.value.detail == ( - "This managed Milvus gRPC connection must be re-saved by a proxy admin before it can be used." - ) - 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_preserves_healthy_context_after_provider_authorization_errors( @@ -207,7 +161,12 @@ async def test_hook_preserves_healthy_context_after_provider_authorization_error status_code: int, ) -> None: registry_with("vs-denied", "vs-safe") - error: Final = HTTPException(status_code=status_code, detail="Vector store access denied") + error: Final = (litellm.AuthenticationError if status_code == 401 else litellm.PermissionDeniedError)( + message="Vector store access denied", + model="embedding-model", + llm_provider="milvus", + response=httpx.Response(status_code, request=httpx.Request("POST", "http://milvus/search")), + ) router: Final = RecordingRouter(failing_vector_store_ids=frozenset({"vs-denied"}), search_error=error) logging_obj: Final = FakeLoggingObj({}) 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 7f197b90200..39bd83beb98 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 @@ -10,6 +10,7 @@ from fastapi import HTTPException, Request import litellm from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( LiteLLM_ManagedVectorStore, + VectorStorePreCallHook, ) from litellm.llms.base_llm.vector_store.transformation import ( LiteLLMVectorStoreEmbeddingExecutor, @@ -42,6 +43,12 @@ from litellm.types.vector_stores import MILVUS_ADMIN_CONFIGURED_CONNECTION, Inde from litellm.vector_stores.main import _direct_vector_store_embedding_executor from litellm.vector_stores.vector_store_registry import VectorStoreRegistry +from tests.test_litellm.integrations.vector_store_integrations.test_vector_store_pre_call_hook import ( + FakeLoggingObj, + FakeProxyRuntime, + RecordingRouter, +) + def _serialize_litellm_params(litellm_params): """Serialize ``litellm_params`` to a string for substring assertions. @@ -3583,3 +3590,56 @@ def test_vector_store_search_missing_query_returns_400(prefix: str) -> None: assert response.status_code == 400, response.text assert "query" in response.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("serialized", [False, True]) +@pytest.mark.parametrize("file_search", [False, True]) +@pytest.mark.parametrize("vector_store_ids", [("legacy",), ("legacy", "safe"), ("safe", "legacy")]) +async def test_hook_rejects_an_untrusted_managed_milvus_grpc_connection( + monkeypatch: pytest.MonkeyPatch, + vector_store_ids: tuple[str, ...], + caplog: pytest.LogCaptureFixture, + serialized: bool, + file_search: bool, +) -> None: + params: Final = {"milvus_transport": "grpc", "api_base": "http://internal-milvus:19530"} + monkeypatch.setattr( + litellm, + "vector_store_registry", + VectorStoreRegistry( + vector_stores=[ + LiteLLM_ManagedVectorStore( + vector_store_id="legacy", + custom_llm_provider="milvus", + litellm_params=json.dumps(params) if serialized else params, + ), + LiteLLM_ManagedVectorStore( + vector_store_id="safe", + custom_llm_provider="bedrock", + ), + ], + ), + ) + router: Final = RecordingRouter() + logging_obj: Final = FakeLoggingObj({}) + + with pytest.raises(HTTPException) as exc_info: + await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=router)).async_get_chat_completion_prompt( + model="chat-model", + messages=[{"role": "user", "content": "what is litellm?"}], + non_default_params={} if file_search else {"vector_store_ids": list(vector_store_ids)}, + tools=[{"type": "file_search", "vector_store_ids": list(vector_store_ids)}] if file_search else None, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + litellm_logging_obj=logging_obj, + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == ( + "This managed Milvus gRPC connection must be re-saved by a proxy admin before it can be used." + ) + assert router.calls == [] + assert "search_results" not in logging_obj.model_call_details + assert "continuing without its context" not in caplog.text 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 00c8f7da5aa..ece38287a35 100644 --- a/tests/test_litellm/vector_stores/test_vector_store_registry.py +++ b/tests/test_litellm/vector_stores/test_vector_store_registry.py @@ -216,3 +216,35 @@ def test_search_uses_registry_credentials(): assert getattr(called_params, "aws_region_name") == "us-east-1" finally: litellm.vector_store_registry = original_registry + + +@pytest.mark.asyncio +@pytest.mark.parametrize("use_async", [False, True]) +@pytest.mark.parametrize("file_search", [False, True]) +async def test_serialized_cached_params_are_normalized_before_tool_merging( + use_async: bool, file_search: bool +) -> None: + params: Final = { + "milvus_transport": "grpc", + "api_base": "http://milvus:19530", + "_litellm_admin_configured_milvus_grpc": True, + } + serialized: Final = json.dumps(params) + store: Final = LiteLLM_ManagedVectorStore( + vector_store_id="documents", custom_llm_provider="milvus", litellm_params=serialized + ) + registry: Final = VectorStoreRegistry(vector_stores=[store]) + request_params: Final = {} if file_search else {"vector_store_ids": ["documents"]} + tools: Final = ( + [{"type": "file_search", "vector_store_ids": ["documents"], "max_num_results": 2}] if file_search else None + ) + + result: Final = ( + await registry.pop_vector_stores_to_run_with_db_fallback(request_params, tools=tools) + if use_async + else registry.pop_vector_stores_to_run(request_params, tools=tools) + ) + + assert len(result) == 1 + assert result[0]["litellm_params"] == ({**params, "max_num_results": 2} if file_search else params) + assert store["litellm_params"] == serialized diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 7ef25d7080a..be86146139b 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22147 + "limit": 22138 }, "LIT002": { "limit": 26733 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16420 + "limit": 16409 }, "LIT011": { "limit": 5506