mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
test(integration): add strict vector store audit cells for routes, SDKs, workers and concurrency
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
ad1e28767e
commit
3ccaf28047
1 changed files with 354 additions and 2 deletions
|
|
@ -1,16 +1,20 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import jwt
|
||||
import openai
|
||||
import pytest
|
||||
import yaml
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
|
|
@ -21,6 +25,9 @@ from integration.authorization._guardrail_opt_out import upstream_observations
|
|||
from pydantic import JsonValue
|
||||
from redis import Redis
|
||||
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
|
||||
CONFIG_STORE_ID: Final = "vs_integration_config_store"
|
||||
PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml"
|
||||
REMOVE_OPENAI_API_BASE: Final = ("OPENAI_API_BASE",)
|
||||
|
|
@ -66,6 +73,20 @@ def _rag_query(
|
|||
return gateway.request("POST", path, _rag_query_body(model, marker, store_id), key=key)
|
||||
|
||||
|
||||
def _served_model(gateway: Gateway, scenario: Scenario) -> str:
|
||||
model: Final = scenario.model()
|
||||
statuses: Final = eventually(
|
||||
lambda: frozenset(
|
||||
_rag_query(gateway, model, f"lit6035 warm {uuid.uuid4().hex}", gateway.key).status_code for _ in range(8)
|
||||
),
|
||||
lambda seen: seen == frozenset({200}),
|
||||
seconds=30,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
assert statuses == frozenset({200}), statuses
|
||||
return model
|
||||
|
||||
|
||||
def _searches_for_marker(
|
||||
gateway: Gateway, marker: str, store_id: str = CONFIG_STORE_ID
|
||||
) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
|
|
@ -174,6 +195,7 @@ def strict_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[StrictG
|
|||
environment,
|
||||
config=config,
|
||||
remove_environment=REMOVE_OPENAI_API_BASE,
|
||||
workers=2,
|
||||
) as gateway:
|
||||
yield StrictGateway(gateway, upstream_gateway, signing_key, config, environment)
|
||||
|
||||
|
|
@ -255,7 +277,7 @@ def test_deny_by_default_rejects_ungranted_store_before_upstream_search(
|
|||
strict_gateway: StrictGateway, case: StrictCase
|
||||
) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
model: Final = _served_model(strict_gateway.gateway, scenario)
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
marker: Final = f"lit6035 deny by default {case} {uuid.uuid4().hex}"
|
||||
|
||||
|
|
@ -314,7 +336,7 @@ def _strict_granted_request(
|
|||
)
|
||||
def test_deny_by_default_searches_explicitly_granted_store(strict_gateway: StrictGateway, case: GrantedCase) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
model: Final = _served_model(strict_gateway.gateway, scenario)
|
||||
store_id: Final = (
|
||||
CONFIG_STORE_ID
|
||||
if case == "standalone_key_granted_registered_store"
|
||||
|
|
@ -584,3 +606,333 @@ def test_revoked_user_grant_stops_working_on_another_proxy(strict_gateway: Stric
|
|||
strict_gateway.gateway.post("/user/update", {"user_id": user, "object_permission": _permission_for_stores()})
|
||||
|
||||
assert eventually(peer_status, lambda status: status == 401, seconds=10, return_last_on_timeout=True) == 401
|
||||
|
||||
|
||||
def _store_searches(gateway: Gateway, marker: str) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
return tuple(
|
||||
observation
|
||||
for observation in upstream_observations(gateway)
|
||||
if str(observation["path"]).startswith("/vector_stores/") and marker in str(observation["body"])
|
||||
)
|
||||
|
||||
|
||||
def _marker_observations(gateway: Gateway, marker: str) -> tuple[Mapping[str, JsonValue], ...]:
|
||||
return tuple(observation for observation in upstream_observations(gateway) if marker in str(observation["body"]))
|
||||
|
||||
|
||||
def _messages_body(model: str, marker: str, store_id: str) -> JsonObject:
|
||||
body: Final[JsonObject] = {
|
||||
"model": model,
|
||||
"max_tokens": 16,
|
||||
"messages": _json_array({"role": "user", "content": marker}),
|
||||
"vector_store_ids": _json_array(store_id),
|
||||
}
|
||||
return body
|
||||
|
||||
|
||||
def _file_search_tools(store_id: JsonValue) -> JsonValue:
|
||||
return _json_array({"type": "file_search", "vector_store_ids": store_id})
|
||||
|
||||
|
||||
RequestShape: TypeAlias = Literal["messages_vector_store_ids", "chat_stream_file_search", "chat_vector_store_ids"]
|
||||
|
||||
|
||||
def _shape_request(shape: RequestShape, model: str, marker: str, store_id: str) -> tuple[str, JsonObject]:
|
||||
messages: Final = _json_array({"role": "user", "content": marker})
|
||||
if shape == "messages_vector_store_ids":
|
||||
return "/v1/messages", _messages_body(model, marker, store_id)
|
||||
if shape == "chat_stream_file_search":
|
||||
stream_body: Final[JsonObject] = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"stream": True,
|
||||
"tools": _file_search_tools(_json_array(store_id)),
|
||||
}
|
||||
return "/v1/chat/completions", stream_body
|
||||
chat_body: Final[JsonObject] = {"model": model, "messages": messages, "vector_store_ids": _json_array(store_id)}
|
||||
return "/v1/chat/completions", chat_body
|
||||
|
||||
|
||||
@pytest.mark.parametrize("shape", ("messages_vector_store_ids", "chat_stream_file_search", "chat_vector_store_ids"))
|
||||
def test_deny_by_default_covers_messages_streaming_and_top_level_store_ids(
|
||||
strict_gateway: StrictGateway, shape: RequestShape
|
||||
) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = _served_model(strict_gateway.gateway, scenario)
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
marker: Final = f"lit6035 shape {shape} {uuid.uuid4().hex}"
|
||||
key: Final = scenario.key(models=_json_array(model))
|
||||
path, body = _shape_request(shape, model, marker, store_id)
|
||||
|
||||
response: Final = strict_gateway.gateway.request("POST", path, body, key=key)
|
||||
|
||||
assert response.status_code == 401, response.text
|
||||
assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text
|
||||
assert _marker_observations(strict_gateway.upstream, marker) == ()
|
||||
|
||||
|
||||
SdkClient: TypeAlias = Literal["openai_sync", "openai_async", "anthropic_sync"]
|
||||
|
||||
|
||||
def _sdk_denial(gateway: Gateway, client: SdkClient, key: str, model: str, marker: str, store_id: str) -> int:
|
||||
base_url: Final = str(gateway.client.base_url)
|
||||
if client == "anthropic_sync":
|
||||
with pytest.raises(anthropic.AuthenticationError) as anthropic_denied:
|
||||
anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0).messages.create(
|
||||
model=model,
|
||||
max_tokens=16,
|
||||
messages=[{"role": "user", "content": marker}],
|
||||
extra_body={"vector_store_ids": _json_array(store_id)},
|
||||
)
|
||||
return anthropic_denied.value.status_code
|
||||
tools: Final = [{"type": "file_search", "vector_store_ids": [store_id]}]
|
||||
if client == "openai_sync":
|
||||
with pytest.raises(openai.AuthenticationError) as sync_denied:
|
||||
openai.OpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0).responses.create(
|
||||
model=model, input=marker, tools=tools
|
||||
)
|
||||
return sync_denied.value.status_code
|
||||
|
||||
async def create() -> None:
|
||||
async with openai.AsyncOpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0) as async_client:
|
||||
await async_client.responses.create(model=model, input=marker, tools=tools)
|
||||
|
||||
with pytest.raises(openai.AuthenticationError) as async_denied:
|
||||
asyncio.run(create())
|
||||
return async_denied.value.status_code
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client", ("openai_sync", "openai_async", "anthropic_sync"))
|
||||
def test_deny_by_default_rejects_sdk_clients_without_grant(strict_gateway: StrictGateway, client: SdkClient) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = _served_model(strict_gateway.gateway, scenario)
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
marker: Final = f"lit6035 sdk {client} {uuid.uuid4().hex}"
|
||||
key: Final = scenario.key(models=_json_array(model))
|
||||
|
||||
status: Final = _sdk_denial(strict_gateway.gateway, client, key, model, marker, store_id)
|
||||
|
||||
assert status == 401
|
||||
assert _marker_observations(strict_gateway.upstream, marker) == ()
|
||||
|
||||
|
||||
MalformedShape: TypeAlias = Literal[
|
||||
"ids_string", "ids_int", "ids_empty_string", "tool_ids_string", "tool_is_string", "oversized_id"
|
||||
]
|
||||
|
||||
|
||||
def _malformed_body(shape: MalformedShape, model: str, marker: str) -> JsonObject:
|
||||
messages: Final = _json_array({"role": "user", "content": marker})
|
||||
if shape == "tool_ids_string":
|
||||
tool_ids_body: Final[JsonObject] = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"tools": _file_search_tools(f"vs_{uuid.uuid4().hex}"),
|
||||
}
|
||||
return tool_ids_body
|
||||
if shape == "tool_is_string":
|
||||
string_tool_body: Final[JsonObject] = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"tools": _json_array("file_search"),
|
||||
}
|
||||
return string_tool_body
|
||||
ids: Final[Mapping[MalformedShape, JsonValue]] = MappingProxyType(
|
||||
{
|
||||
"ids_string": f"vs_{uuid.uuid4().hex}",
|
||||
"ids_int": _json_array(123),
|
||||
"ids_empty_string": _json_array(""),
|
||||
"oversized_id": _json_array("vs_" + "x" * 5000),
|
||||
}
|
||||
)
|
||||
body: Final[JsonObject] = {"model": model, "messages": messages, "vector_store_ids": ids[shape]}
|
||||
return body
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"shape", ("ids_string", "ids_int", "ids_empty_string", "tool_ids_string", "tool_is_string", "oversized_id")
|
||||
)
|
||||
def test_deny_by_default_malformed_store_ids_never_search_a_store(
|
||||
strict_gateway: StrictGateway, shape: MalformedShape
|
||||
) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = _served_model(strict_gateway.gateway, scenario)
|
||||
marker: Final = f"lit6035 malformed {shape} {uuid.uuid4().hex}"
|
||||
key: Final = scenario.key(models=_json_array(model))
|
||||
|
||||
response: Final = strict_gateway.gateway.request(
|
||||
"POST", "/v1/chat/completions", _malformed_body(shape, model, marker), key=key
|
||||
)
|
||||
searches: Final = _store_searches(strict_gateway.upstream, marker)
|
||||
healthy: Final = strict_gateway.gateway.request(
|
||||
"POST", "/v1/rag/query", _rag_query_body(model, f"{marker} master", f"vs_{uuid.uuid4().hex}")
|
||||
)
|
||||
|
||||
assert response.status_code < 500, response.text
|
||||
assert searches == ()
|
||||
assert healthy.status_code == 200, healthy.text
|
||||
|
||||
|
||||
def test_deny_by_default_searches_a_duplicated_granted_store_once(strict_gateway: StrictGateway) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = _served_model(strict_gateway.gateway, scenario)
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
marker: Final = f"lit6035 duplicate {uuid.uuid4().hex}"
|
||||
key: Final = scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id))
|
||||
body: Final[JsonObject] = {
|
||||
**_rag_query_body(model, marker, store_id),
|
||||
"vector_store_ids": _json_array(store_id, store_id),
|
||||
}
|
||||
|
||||
response: Final = strict_gateway.gateway.request("POST", "/v1/rag/query", body, key=key)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert len(_searches_for_marker(strict_gateway.upstream, marker, store_id)) == 1
|
||||
|
||||
|
||||
def test_deny_by_default_treats_null_key_grant_list_as_no_grant(strict_gateway: StrictGateway) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = _served_model(strict_gateway.gateway, scenario)
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
marker: Final = f"lit6035 null grants {uuid.uuid4().hex}"
|
||||
null_grants: Final[JsonObject] = {"vector_stores": None}
|
||||
key: Final = scenario.key(models=_json_array(model), object_permission=null_grants)
|
||||
|
||||
response: Final = _rag_query(strict_gateway.gateway, model, marker, key, store_id=store_id)
|
||||
|
||||
assert response.status_code == 401, response.text
|
||||
assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text
|
||||
assert _marker_observations(strict_gateway.upstream, marker) == ()
|
||||
|
||||
|
||||
def _statuses_across_workers(gateway: Gateway, model: str, key: str, store_id: str) -> frozenset[int]:
|
||||
return frozenset(
|
||||
_rag_query(gateway, model, f"lit6035 grant change {uuid.uuid4().hex}", key, store_id=store_id).status_code
|
||||
for _ in range(6)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("scope", ("key", "team"))
|
||||
def test_deny_by_default_grant_and_revoke_take_effect_on_every_worker(
|
||||
strict_gateway: StrictGateway, scope: Literal["key", "team"]
|
||||
) -> None:
|
||||
gateway: Final = strict_gateway.gateway
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = _served_model(gateway, scenario)
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
team: Final = scenario.team(models=_json_array(model), object_permission=_permission_for_stores(store_id))
|
||||
key: Final = (
|
||||
scenario.key(models=_json_array(model))
|
||||
if scope == "key"
|
||||
else scenario.key(
|
||||
team_id=team, models=_json_array(model), object_permission=_permission_for_stores(store_id)
|
||||
)
|
||||
)
|
||||
team_key: Final = scope == "team"
|
||||
update_path: Final = "/team/update" if team_key else "/key/update"
|
||||
identity: Final[JsonObject] = {"team_id": team} if team_key else {"key": key}
|
||||
if team_key:
|
||||
gateway.post(update_path, {**identity, "object_permission": _permission_for_stores()})
|
||||
|
||||
def statuses() -> frozenset[int]:
|
||||
return _statuses_across_workers(gateway, model, key, store_id)
|
||||
|
||||
before: Final = eventually(statuses, lambda seen: seen == frozenset({401}), return_last_on_timeout=True)
|
||||
gateway.post(update_path, {**identity, "object_permission": _permission_for_stores(store_id)})
|
||||
granted: Final = eventually(statuses, lambda seen: seen == frozenset({200}), return_last_on_timeout=True)
|
||||
gateway.post(update_path, {**identity, "object_permission": _permission_for_stores()})
|
||||
revoked: Final = eventually(statuses, lambda seen: seen == frozenset({401}), return_last_on_timeout=True)
|
||||
|
||||
assert (before, granted, revoked) == (frozenset({401}), frozenset({200}), frozenset({401}))
|
||||
|
||||
|
||||
def _cli_session_token(user_id: str, team_id: str) -> str:
|
||||
cli_user: Final = LiteLLM_UserTable(user_id=user_id, user_role="internal_user", teams=[team_id], models=[])
|
||||
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team_id, team_alias="cli-team")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("granted_by", ("team", "user"))
|
||||
def test_deny_by_default_session_token_uses_only_the_resolved_team_grant(
|
||||
strict_gateway: StrictGateway, granted_by: Literal["team", "user"], monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"))
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = _served_model(strict_gateway.gateway, scenario)
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
marker: Final = f"lit6035 session token {granted_by} {uuid.uuid4().hex}"
|
||||
team_grants: Final = _permission_for_stores(store_id if granted_by == "team" else "vs_some_other_store")
|
||||
user: Final = scenario.user(
|
||||
user_role="internal_user",
|
||||
object_permission=_permission_for_stores(store_id if granted_by == "user" else "vs_some_other_store"),
|
||||
)
|
||||
team: Final = scenario.team(
|
||||
models=_json_array(model),
|
||||
object_permission=team_grants,
|
||||
members_with_roles=_json_array({"role": "user", "user_id": user}),
|
||||
)
|
||||
|
||||
response: Final = _rag_query(
|
||||
strict_gateway.gateway, model, marker, _cli_session_token(user, team), store_id=store_id
|
||||
)
|
||||
|
||||
if granted_by == "team":
|
||||
assert response.status_code == 200, response.text
|
||||
assert len(_searches_for_marker(strict_gateway.upstream, marker, store_id)) == 1
|
||||
return
|
||||
assert response.status_code == 401, response.text
|
||||
assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text
|
||||
assert _marker_observations(strict_gateway.upstream, marker) == ()
|
||||
|
||||
|
||||
BurstRoute: TypeAlias = Literal["chat_retrieval_config", "search_route"]
|
||||
|
||||
|
||||
def _burst_request(
|
||||
gateway: Gateway, route: BurstRoute, model: str, key: str, store_id: str, marker: str
|
||||
) -> httpx.Response:
|
||||
if route == "search_route":
|
||||
return gateway.request("POST", f"/v1/vector_stores/{store_id}/search", {"query": marker}, key=key)
|
||||
return _rag_query(gateway, model, marker, key, store_id=store_id, path="/v1/chat/completions")
|
||||
|
||||
|
||||
def test_deny_by_default_concurrent_burst_only_searches_granted_stores(strict_gateway: StrictGateway) -> None:
|
||||
gateway: Final = strict_gateway.gateway
|
||||
routes: Final[tuple[BurstRoute, ...]] = ("chat_retrieval_config", "search_route")
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = _served_model(gateway, scenario)
|
||||
store_id: Final = CONFIG_STORE_ID
|
||||
granted_key: Final = scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id))
|
||||
ungranted_key: Final = scenario.key(models=_json_array(model))
|
||||
plan: Final = tuple(
|
||||
(routes[index % 2], index // 2 % 2 == 0, f"lit6035 burst {index} {uuid.uuid4().hex}") for index in range(32)
|
||||
)
|
||||
upstream_observations(strict_gateway.upstream)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=10) as pool:
|
||||
responses: Final = tuple(
|
||||
pool.map(
|
||||
lambda item: _burst_request(
|
||||
gateway, item[0], model, granted_key if item[1] else ungranted_key, store_id, item[2]
|
||||
),
|
||||
plan,
|
||||
)
|
||||
)
|
||||
observed: Final = tuple(
|
||||
(str(observation["path"]), str(observation["body"]))
|
||||
for observation in upstream_observations(strict_gateway.upstream)
|
||||
)
|
||||
|
||||
def hits(route: BurstRoute, marker: str) -> tuple[int, int]:
|
||||
expected_path: Final = (
|
||||
"/chat/completions" if route == "chat_retrieval_config" else f"/vector_stores/{store_id}/search"
|
||||
)
|
||||
on_route: Final = sum(path.endswith(expected_path) and marker in body for path, body in observed)
|
||||
return on_route, sum(marker in body for _, body in observed)
|
||||
|
||||
assert tuple(response.status_code for response in responses) == tuple(
|
||||
200 if granted else 401 for _, granted, _ in plan
|
||||
), tuple(response.text[:300] for response in responses if response.status_code not in (200, 401))
|
||||
assert tuple(hits(route, marker)[0] for route, _, marker in plan) == tuple(
|
||||
1 if granted else 0 for _, granted, _ in plan
|
||||
)
|
||||
assert tuple(hits(route, marker)[1] for route, granted, marker in plan if not granted) == (0,) * 16
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue