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:
mrinal 2026-10-02 23:45:15 +00:00
parent ad1e28767e
commit 3ccaf28047

View file

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