mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(vector-stores): normalize cached parameters before approval checks
This commit is contained in:
parent
3f55217265
commit
a5fa90f611
8 changed files with 123 additions and 63 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -246,7 +246,7 @@
|
|||
"limit": 109
|
||||
},
|
||||
"TRY300": {
|
||||
"limit": 850
|
||||
"limit": 849
|
||||
},
|
||||
"UP028": {
|
||||
"limit": 2
|
||||
|
|
|
|||
|
|
@ -21,6 +21,6 @@
|
|||
"limit": 117
|
||||
},
|
||||
"TQ008": {
|
||||
"limit": 10978
|
||||
"limit": 10975
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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({})
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue