fix(vector-stores): normalize cached parameters before approval checks

This commit is contained in:
Yujong Lee 2026-09-05 12:14:49 -07:00
parent 3f55217265
commit a5fa90f611
8 changed files with 123 additions and 63 deletions

View file

@ -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

View file

@ -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:

View file

@ -246,7 +246,7 @@
"limit": 109
},
"TRY300": {
"limit": 850
"limit": 849
},
"UP028": {
"limit": 2

View file

@ -21,6 +21,6 @@
"limit": 117
},
"TQ008": {
"limit": 10978
"limit": 10975
}
}

View file

@ -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({})

View file

@ -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

View file

@ -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

View file

@ -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