mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): reject invalid vector_store_deny_by_default at config load and return 400 for malformed vector_store_ids
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
50df20c174
commit
3e5fe50faa
4 changed files with 178 additions and 21 deletions
|
|
@ -15,6 +15,7 @@ import re
|
|||
import time
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
|
||||
from functools import partial
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias
|
||||
|
||||
|
|
@ -146,7 +147,6 @@ from litellm.router import Router
|
|||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
from litellm.types.proxy.model_access_group_budget import ModelAccessGroupBudget
|
||||
from litellm.utils import get_utc_datetime
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
|
||||
from .auth_checks_organization import (
|
||||
add_team_org_context_to_request_body,
|
||||
|
|
@ -1286,15 +1286,7 @@ async def common_checks(
|
|||
team_object=team_object,
|
||||
valid_token=valid_token,
|
||||
user_object=user_object,
|
||||
deny_by_default=ConfigGeneralSettings.model_validate(
|
||||
MappingProxyType(
|
||||
{
|
||||
"vector_store_deny_by_default": _typed_request_body(general_settings).get(
|
||||
"vector_store_deny_by_default", False
|
||||
)
|
||||
}
|
||||
)
|
||||
).vector_store_deny_by_default,
|
||||
deny_by_default=_vector_store_deny_by_default(_typed_request_body(general_settings)),
|
||||
)
|
||||
|
||||
# 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path)
|
||||
|
|
@ -6644,6 +6636,73 @@ def _is_strict_vector_store_identity(valid_token: UserAPIKeyAuth | None) -> bool
|
|||
)
|
||||
|
||||
|
||||
def _vector_store_deny_by_default(general_settings: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
Startup rejects a non-boolean value from the config file. A non-boolean value that reaches
|
||||
general_settings another way enables the policy, so only vector store requests are denied.
|
||||
"""
|
||||
try:
|
||||
return ConfigGeneralSettings.model_validate(
|
||||
MappingProxyType(
|
||||
{"vector_store_deny_by_default": general_settings.get("vector_store_deny_by_default", False)}
|
||||
)
|
||||
).vector_store_deny_by_default
|
||||
except ValidationError:
|
||||
return True
|
||||
|
||||
|
||||
_VECTOR_STORE_IDS_ADAPTER: Final[TypeAdapter[list[str]]] = TypeAdapter(list[str])
|
||||
_TOOLS_ADAPTER: Final[TypeAdapter[list[object]]] = TypeAdapter(list[object])
|
||||
_TOOL_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _validated_vector_store_ids(value: object) -> tuple[str, ...]:
|
||||
if value is None:
|
||||
return ()
|
||||
try:
|
||||
return tuple(_VECTOR_STORE_IDS_ADAPTER.validate_python(value, strict=True))
|
||||
except ValidationError:
|
||||
raise _malformed_vector_store_ids() from None
|
||||
|
||||
|
||||
def _malformed_vector_store_ids() -> ProxyException:
|
||||
return ProxyException(
|
||||
message="vector_store_ids must be a list of strings",
|
||||
type="invalid_request_error",
|
||||
param="vector_store_ids",
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
|
||||
def _tools(tools: object) -> tuple[object, ...]:
|
||||
try:
|
||||
return tuple(_TOOLS_ADAPTER.validate_python(tools, strict=True))
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def _tool_vector_store_ids(tool: object) -> tuple[str, ...]:
|
||||
try:
|
||||
tool_fields: Final = _TOOL_ADAPTER.validate_python(tool, strict=True)
|
||||
except ValidationError:
|
||||
return ()
|
||||
return _validated_vector_store_ids(tool_fields.get("vector_store_ids"))
|
||||
|
||||
|
||||
def _strict_requested_vector_store_ids(request_body: Mapping[str, object]) -> tuple[str, ...]:
|
||||
"""
|
||||
Same fields VectorStoreRegistry.get_vector_store_ids_to_run reads, but a vector_store_ids that is
|
||||
not a list of strings is a 400 instead of being skipped or iterated, and tools that are not
|
||||
objects name no store.
|
||||
"""
|
||||
return tuple(
|
||||
chain(
|
||||
_validated_vector_store_ids(request_body.get("vector_store_ids")),
|
||||
chain.from_iterable(_tool_vector_store_ids(tool) for tool in _tools(request_body.get("tools"))),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _require_vector_store_grant(
|
||||
object_type: Literal["key", "team", "user"],
|
||||
vector_store_ids_to_run: Sequence[str],
|
||||
|
|
@ -6694,7 +6753,8 @@ async def vector_store_access_check(
|
|||
rescue nor restrict a key or team request.
|
||||
|
||||
Raises ProxyException (401, `{key,team,user}_vector_store_access_denied`) on the first identity
|
||||
that does not grant a requested store.
|
||||
that does not grant a requested store, and with the flag on, ProxyException (400,
|
||||
`invalid_request_error`) when a `vector_store_ids` field is not a list of strings.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
|
||||
|
|
@ -6705,18 +6765,18 @@ async def vector_store_access_check(
|
|||
verbose_proxy_logger.debug("Prisma client not found, skipping vector store access check")
|
||||
return True
|
||||
|
||||
vector_store_registry: Final = (
|
||||
VectorStoreRegistry()
|
||||
if litellm.vector_store_registry is None and deny_by_default
|
||||
else litellm.vector_store_registry
|
||||
)
|
||||
registry_ids: Final = (
|
||||
vector_store_registry.get_vector_store_ids_to_run(
|
||||
non_default_params=request_body, tools=request_body.get("tools", None)
|
||||
_strict_requested_vector_store_ids(_typed_request_body(request_body))
|
||||
if deny_by_default
|
||||
else (
|
||||
litellm.vector_store_registry.get_vector_store_ids_to_run(
|
||||
non_default_params=request_body, tools=request_body.get("tools", None)
|
||||
)
|
||||
if litellm.vector_store_registry is not None
|
||||
else None
|
||||
)
|
||||
if vector_store_registry is not None
|
||||
else None
|
||||
) or ()
|
||||
or ()
|
||||
)
|
||||
rag_vector_store_id: Final = _get_rag_query_vector_store_id(_typed_request_body(request_body))
|
||||
rag_ids: Final = (rag_vector_store_id,) if rag_vector_store_id is not None else ()
|
||||
vector_store_ids_to_run: Final = tuple(dict.fromkeys((*registry_ids, *rag_ids)))
|
||||
|
|
|
|||
|
|
@ -6657,6 +6657,14 @@ class ProxyConfig:
|
|||
if general_settings is None:
|
||||
general_settings = {}
|
||||
|
||||
typed_general_settings: Final = _GENERAL_SETTINGS_VIEW.validate_python(general_settings)
|
||||
if "vector_store_deny_by_default" in typed_general_settings:
|
||||
ConfigGeneralSettings.model_validate(
|
||||
MappingProxyType(
|
||||
{"vector_store_deny_by_default": typed_general_settings["vector_store_deny_by_default"]}
|
||||
)
|
||||
)
|
||||
|
||||
if general_settings.get("mcp_advertised_versions") is not None:
|
||||
from litellm.types.mcp import MCPAdvertisedVersions
|
||||
|
||||
|
|
|
|||
|
|
@ -2275,6 +2275,78 @@ async def test_deny_by_default_reads_requested_vector_stores_without_a_registry(
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("request_body", "key_vector_stores", "denied_by"),
|
||||
[
|
||||
({"tools": [1]}, None, None),
|
||||
({"tools": 1}, None, None),
|
||||
({"tools": [{"type": "file_search", "vector_store_ids": None}]}, None, None),
|
||||
({"tools": [1, {"type": "file_search", "vector_store_ids": ["KBSTOREA"]}]}, ["KBSTOREA"], None),
|
||||
({"tools": [1, {"type": "file_search", "vector_store_ids": ["KBSTOREA"]}]}, ["KBSTOREB"], _KEY_DENIED),
|
||||
],
|
||||
ids=["int-tool", "non-list-tools", "null-tool-ids", "int-tool-beside-granted", "int-tool-beside-ungranted"],
|
||||
)
|
||||
async def test_deny_by_default_ignores_tools_that_name_no_vector_store(
|
||||
request_body: Mapping[str, object], key_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
_virtual_key(object_permission_id=None if key_vector_stores is None else "key-permission"),
|
||||
None if key_vector_stores is None else SimpleNamespace(vector_stores=key_vector_stores),
|
||||
request_body={"model": "gpt-4o-mini", "input": "what is in this KB?", **request_body},
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"request_body",
|
||||
[
|
||||
{"tools": [{"type": "file_search", "vector_store_ids": "KBSTOREA"}]},
|
||||
{"tools": [{"type": "file_search", "vector_store_ids": [1]}]},
|
||||
{"vector_store_ids": "KBSTOREA"},
|
||||
],
|
||||
ids=["string-tool-ids", "int-tool-id", "string-top-level-ids"],
|
||||
)
|
||||
async def test_deny_by_default_rejects_malformed_vector_store_ids_as_bad_request(request_body: Mapping[str, object]):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
_virtual_key(object_permission_id="key-permission"),
|
||||
SimpleNamespace(vector_stores=["KBSTOREA"]),
|
||||
request_body={"model": "gpt-4o-mini", "input": "what is in this KB?", **request_body},
|
||||
)
|
||||
assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == (
|
||||
"invalid_request_error",
|
||||
"vector_store_ids",
|
||||
"400",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("request_body", "denied_by"),
|
||||
[
|
||||
({"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}]}, None),
|
||||
(None, _KEY_DENIED),
|
||||
],
|
||||
ids=["no-vector-store", "ungranted-vector-store"],
|
||||
)
|
||||
@pytest.mark.parametrize("flag_value", [None, "enabled"], ids=["null", "string"])
|
||||
async def test_invalid_deny_by_default_value_only_denies_vector_store_requests(
|
||||
flag_value: object, request_body: Mapping[str, object] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": flag_value}, _virtual_key(), None, request_body=request_body
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_keyless_user_grant_is_read_through_the_object_permission_cache():
|
||||
cache: Final = UserApiKeyCache(default_in_memory_ttl=60)
|
||||
|
|
|
|||
|
|
@ -2293,6 +2293,23 @@ async def test_load_config_yaml_vector_store_deny_by_default_is_boolean(
|
|||
assert ConfigGeneralSettings.model_validate(dict(general_settings)).vector_store_deny_by_default is expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("yaml_value", ["", "enabled"], ids=["null", "string"])
|
||||
async def test_load_config_rejects_non_boolean_vector_store_deny_by_default(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, yaml_value: str
|
||||
):
|
||||
config_file: Final = tmp_path / "vector_store.yaml"
|
||||
config_file.write_text(
|
||||
f"model_list: []\nlitellm_settings: {{}}\ngeneral_settings:\n vector_store_deny_by_default: {yaml_value}\n"
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
|
||||
with pytest.raises(ValidationError, match="vector_store_deny_by_default"):
|
||||
await ProxyConfig().load_config(router=None, config_file_path=str(config_file))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_load_config_resolves_router_settings_plugins(tmp_path, monkeypatch):
|
||||
"""Regression: router_settings.plugins dotted-path strings must be resolved to
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue