mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
test(proxy): cover vector_store_deny_by_default through a real proxy with key, team and JWT identities
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b94047f2ac
commit
d3c21f716e
2 changed files with 301 additions and 7 deletions
|
|
@ -1,5 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from pathlib import Path
|
||||
|
|
@ -7,16 +9,20 @@ from types import MappingProxyType
|
|||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
import jwt
|
||||
import pytest
|
||||
import yaml
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from integration.authorization._guardrail_opt_out import upstream_observations
|
||||
from pydantic import JsonValue
|
||||
|
||||
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",)
|
||||
JWT_KEY_ID: Final = "integration-vector-store-jwt-key"
|
||||
JsonObject: TypeAlias = dict[str, JsonValue]
|
||||
|
||||
|
||||
|
|
@ -99,6 +105,209 @@ def no_registry_gateways(tmp_path_factory: pytest.TempPathFactory) -> Iterator[t
|
|||
yield no_registry_gateway, upstream_gateway
|
||||
|
||||
|
||||
class StrictGateway:
|
||||
def __init__(self, gateway: Gateway, upstream: Gateway, signing_key: rsa.RSAPrivateKey) -> None:
|
||||
self.gateway: Final = gateway
|
||||
self.upstream: Final = upstream
|
||||
self._signing_key: Final = signing_key
|
||||
|
||||
def jwt(self, subject: str, groups: tuple[str, ...] = ()) -> str:
|
||||
claims: Final[JsonObject] = {
|
||||
"sub": subject,
|
||||
"groups": _json_array(*groups),
|
||||
"iat": int(time.time()),
|
||||
"exp": int(time.time()) + 300,
|
||||
}
|
||||
return jwt.encode(claims, self._signing_key, algorithm="RS256", headers={"kid": JWT_KEY_ID})
|
||||
|
||||
|
||||
def _strict_config(directory: Path, *, deny_by_default: bool = True) -> Path:
|
||||
config: Final = object_value(yaml.safe_load(PROXY_CONFIG.read_text()))
|
||||
general_settings: Final = object_value(config["general_settings"])
|
||||
strict: Final[JsonObject] = {
|
||||
**config,
|
||||
"general_settings": {
|
||||
**general_settings,
|
||||
"vector_store_deny_by_default": deny_by_default,
|
||||
"enable_jwt_auth": True,
|
||||
"litellm_jwtauth": {
|
||||
"user_id_jwt_field": "sub",
|
||||
"team_ids_jwt_field": "groups",
|
||||
"team_allowed_routes": _json_array("openai_routes", "info_routes", "/v1/rag/query"),
|
||||
},
|
||||
},
|
||||
}
|
||||
path: Final = directory / f"proxy_vector_store_deny_by_default_{deny_by_default}.yaml"
|
||||
path.write_text(yaml.safe_dump(strict))
|
||||
return path
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def strict_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[StrictGateway]:
|
||||
signing_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
public_jwk: Final = jwt.algorithms.RSAAlgorithm.to_jwk(signing_key.public_key())
|
||||
jwks_body: Final = json.dumps({"keys": [{**json.loads(public_jwk), "kid": JWT_KEY_ID}]}).encode()
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "GET", request
|
||||
return Reply(body=jwks_body)
|
||||
|
||||
with gateway_from_environment() as upstream_gateway, wire_server(respond) as jwks:
|
||||
directory: Final = tmp_path_factory.mktemp("rag_query_deny_by_default")
|
||||
with owned_proxy(
|
||||
upstream_gateway,
|
||||
directory,
|
||||
{**_openai_environment(upstream_gateway), "JWT_PUBLIC_KEY_URL": jwks.url},
|
||||
config=_strict_config(directory),
|
||||
remove_environment=REMOVE_OPENAI_API_BASE,
|
||||
) as gateway:
|
||||
yield StrictGateway(gateway, upstream_gateway, signing_key)
|
||||
|
||||
|
||||
StrictCase: TypeAlias = Literal[
|
||||
"standalone_key_no_permission",
|
||||
"team_key_empty_key_grants",
|
||||
"team_key_empty_team_grants",
|
||||
"multi_store_one_ungranted",
|
||||
"rag_alias_no_permission",
|
||||
"chat_retrieval_config_no_permission",
|
||||
"jwt_user_without_grant",
|
||||
]
|
||||
|
||||
|
||||
def _strict_denied_request(
|
||||
strict: StrictGateway, scenario: Scenario, case: StrictCase, model: str, marker: str, store_id: str
|
||||
) -> tuple[httpx.Response, str]:
|
||||
models: Final = _json_array(model)
|
||||
granted: Final = _permission_for_stores(store_id)
|
||||
empty: Final = _permission_for_stores()
|
||||
if case == "standalone_key_no_permission":
|
||||
key = scenario.key(models=models)
|
||||
return strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=key), (
|
||||
"key_vector_store_access_denied"
|
||||
)
|
||||
if case in ("team_key_empty_key_grants", "team_key_empty_team_grants"):
|
||||
key_grants_store: Final = case == "team_key_empty_team_grants"
|
||||
team = scenario.team(models=models, object_permission=empty if key_grants_store else granted)
|
||||
key = scenario.key(team_id=team, models=models, object_permission=granted if key_grants_store else empty)
|
||||
error_type: Final = "team_vector_store_access_denied" if key_grants_store else "key_vector_store_access_denied"
|
||||
return strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=key), (
|
||||
error_type
|
||||
)
|
||||
if case == "multi_store_one_ungranted":
|
||||
key = scenario.key(models=models, object_permission=_permission_for_stores(store_id))
|
||||
body: Final[JsonObject] = {
|
||||
**_rag_query_body(model, marker, store_id),
|
||||
"tools": _json_array({"type": "file_search", "vector_store_ids": _json_array(CONFIG_STORE_ID)}),
|
||||
}
|
||||
return strict.gateway.request("POST", "/v1/chat/completions", body, key=key), "key_vector_store_access_denied"
|
||||
if case == "rag_alias_no_permission":
|
||||
key = scenario.key(models=models)
|
||||
return strict.gateway.request("POST", "/rag/query", _rag_query_body(model, marker, store_id), key=key), (
|
||||
"key_vector_store_access_denied"
|
||||
)
|
||||
if case == "chat_retrieval_config_no_permission":
|
||||
key = scenario.key(models=models)
|
||||
return (
|
||||
strict.gateway.request("POST", "/v1/chat/completions", _rag_query_body(model, marker, store_id), key=key),
|
||||
"key_vector_store_access_denied",
|
||||
)
|
||||
user: Final = scenario.user(user_role="internal_user")
|
||||
return (
|
||||
strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=strict.jwt(user)),
|
||||
"user_vector_store_access_denied",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"case",
|
||||
(
|
||||
"standalone_key_no_permission",
|
||||
"team_key_empty_key_grants",
|
||||
"team_key_empty_team_grants",
|
||||
"multi_store_one_ungranted",
|
||||
"rag_alias_no_permission",
|
||||
"chat_retrieval_config_no_permission",
|
||||
"jwt_user_without_grant",
|
||||
),
|
||||
)
|
||||
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()
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
marker: Final = f"lit6035 deny by default {case} {uuid.uuid4().hex}"
|
||||
|
||||
response, error_type = _strict_denied_request(strict_gateway, scenario, case, model, marker, store_id)
|
||||
observations: Final = tuple(
|
||||
observation
|
||||
for observation in upstream_observations(strict_gateway.upstream)
|
||||
if marker in str(observation["body"])
|
||||
)
|
||||
|
||||
assert response.status_code == 401, f"{response.text}; scripted_upstream_observations={observations!r}"
|
||||
assert response.json()["error"]["type"] == error_type, response.text
|
||||
assert observations == ()
|
||||
|
||||
|
||||
GrantedCase: TypeAlias = Literal[
|
||||
"standalone_key_granted_registered_store",
|
||||
"team_key_both_grant_unregistered_store",
|
||||
"jwt_team_member_team_grant_only",
|
||||
"jwt_user_personal_grant",
|
||||
"master_key_without_grants",
|
||||
]
|
||||
|
||||
|
||||
def _strict_granted_request(
|
||||
strict: StrictGateway, scenario: Scenario, case: GrantedCase, model: str, marker: str, store_id: str
|
||||
) -> httpx.Response:
|
||||
models: Final = _json_array(model)
|
||||
granted: Final = _permission_for_stores(store_id)
|
||||
body: Final = _rag_query_body(model, marker, store_id)
|
||||
if case in ("standalone_key_granted_registered_store", "team_key_both_grant_unregistered_store"):
|
||||
team = scenario.team(models=models, object_permission=granted) if case.startswith("team") else None
|
||||
key = scenario.key(models=models, object_permission=granted, **({} if team is None else {"team_id": team}))
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body, key=key)
|
||||
if case == "jwt_team_member_team_grant_only":
|
||||
team = scenario.team(models=models, object_permission=granted)
|
||||
member: Final = scenario.member(team)
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(member, (team,)))
|
||||
if case == "jwt_user_personal_grant":
|
||||
user: Final = scenario.user(user_role="internal_user", object_permission=granted)
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(user))
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"case",
|
||||
(
|
||||
"standalone_key_granted_registered_store",
|
||||
"team_key_both_grant_unregistered_store",
|
||||
"jwt_team_member_team_grant_only",
|
||||
"jwt_user_personal_grant",
|
||||
"master_key_without_grants",
|
||||
),
|
||||
)
|
||||
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()
|
||||
store_id: Final = (
|
||||
CONFIG_STORE_ID
|
||||
if case == "standalone_key_granted_registered_store"
|
||||
else f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
)
|
||||
marker: Final = f"lit6035 granted {case} {uuid.uuid4().hex}"
|
||||
|
||||
response: Final = _strict_granted_request(strict_gateway, scenario, case, model, marker, store_id)
|
||||
assert response.status_code == 200, response.text
|
||||
|
||||
searches: Final = _searches_for_marker(strict_gateway.upstream, marker, store_id)
|
||||
assert len(searches) == 1, searches
|
||||
assert marker in str(object_value(searches[0]["body"])["query"]), searches
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("scope", "error_type"),
|
||||
(("key", "key_vector_store_access_denied"), ("team", "team_vector_store_access_denied")),
|
||||
|
|
@ -220,3 +429,31 @@ def test_rag_query_alias_denies_store_when_team_allowlist_excludes(gateway: Gate
|
|||
assert response.status_code == 401, response.text
|
||||
assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text
|
||||
assert _searches_for_marker(gateway, marker) == ()
|
||||
|
||||
|
||||
def test_explicit_false_flag_keeps_legacy_vector_store_outcomes(tmp_path: Path) -> None:
|
||||
with gateway_from_environment() as upstream_gateway:
|
||||
with owned_proxy(
|
||||
upstream_gateway,
|
||||
tmp_path,
|
||||
_openai_environment(upstream_gateway),
|
||||
config=_strict_config(tmp_path, deny_by_default=False),
|
||||
remove_environment=REMOVE_OPENAI_API_BASE,
|
||||
) as gateway:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
allowed_marker: Final = f"flag-false-allowed-{uuid.uuid4().hex}"
|
||||
denied_marker: Final = f"flag-false-denied-{uuid.uuid4().hex}"
|
||||
no_permission_key: Final = scenario.key(models=_json_array(model))
|
||||
excluding_key: Final = scenario.key(
|
||||
models=_json_array(model), object_permission=_permission_for_stores("vs_some_other_store")
|
||||
)
|
||||
|
||||
allowed: Final = _rag_query(gateway, model, allowed_marker, no_permission_key)
|
||||
denied: Final = _rag_query(gateway, model, denied_marker, excluding_key)
|
||||
|
||||
assert allowed.status_code == 200, allowed.text
|
||||
assert len(_searches_for_marker(upstream_gateway, allowed_marker)) == 1
|
||||
assert denied.status_code == 401, denied.text
|
||||
assert denied.json()["error"]["type"] == "key_vector_store_access_denied"
|
||||
assert _searches_for_marker(upstream_gateway, denied_marker) == ()
|
||||
|
|
|
|||
|
|
@ -1846,6 +1846,8 @@ async def _common_checks_for_rag_query(
|
|||
team_permission: SimpleNamespace | None = None,
|
||||
user_object: LiteLLM_UserTable | None = None,
|
||||
user_permission: SimpleNamespace | LiteLLM_ObjectPermissionTable | None = None,
|
||||
request_body: Mapping[str, object] | None = None,
|
||||
vector_store_registry: VectorStoreRegistry | None = None,
|
||||
) -> bool:
|
||||
permissions: Final = {
|
||||
"key-permission": key_permission,
|
||||
|
|
@ -1857,21 +1859,25 @@ async def _common_checks_for_rag_query(
|
|||
side_effect=lambda where: permissions.get(where["object_permission_id"])
|
||||
)
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
|
||||
request_body = {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "what is in this KB?"}],
|
||||
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
|
||||
}
|
||||
body: Final = (
|
||||
dict(request_body)
|
||||
if request_body is not None
|
||||
else {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "what is in this KB?"}],
|
||||
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
|
||||
}
|
||||
)
|
||||
with (
|
||||
patch( # test-quality-ok: production auth reads these module globals; no dependency injection seam exists
|
||||
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
|
||||
),
|
||||
patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists
|
||||
"litellm.vector_store_registry", None
|
||||
"litellm.vector_store_registry", vector_store_registry
|
||||
),
|
||||
):
|
||||
return await common_checks(
|
||||
request_body=request_body,
|
||||
request_body=body,
|
||||
team_object=team_object,
|
||||
user_object=user_object,
|
||||
end_user_object=None,
|
||||
|
|
@ -2171,6 +2177,57 @@ async def test_user_owned_standalone_key_cannot_use_owner_grants_under_deny_by_d
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("deny_by_default", "key_vector_stores", "denied_by"),
|
||||
[
|
||||
(True, ["KBSTOREA", "KBSTOREB"], None),
|
||||
(True, ["KBSTOREA"], _KEY_DENIED),
|
||||
(True, ["KBSTOREB"], _KEY_DENIED),
|
||||
(False, ["KBSTOREA"], _KEY_DENIED),
|
||||
],
|
||||
ids=["enabled-grants-both", "enabled-missing-second", "enabled-missing-first", "disabled-missing-second"],
|
||||
)
|
||||
async def test_every_requested_vector_store_needs_a_grant(
|
||||
deny_by_default: bool, key_vector_stores: list[str], denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": deny_by_default},
|
||||
_virtual_key(object_permission_id="key-permission"),
|
||||
SimpleNamespace(vector_stores=key_vector_stores),
|
||||
request_body={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "what is in these KBs?"}],
|
||||
"tools": [{"type": "file_search", "vector_store_ids": ["KBSTOREB"]}],
|
||||
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
|
||||
},
|
||||
vector_store_registry=VectorStoreRegistry(),
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"valid_token",
|
||||
[_virtual_key(), UserAPIKeyAuth(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER)],
|
||||
ids=["standalone-key", "keyless-user"],
|
||||
)
|
||||
async def test_request_without_vector_stores_is_unaffected_by_deny_by_default(valid_token: UserAPIKeyAuth):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
valid_token,
|
||||
None,
|
||||
user_object=_user_row(None),
|
||||
request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}]},
|
||||
vector_store_registry=VectorStoreRegistry(),
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def test_can_object_call_model_with_alias():
|
||||
"""Test that can_object_call_model works with model aliases"""
|
||||
from litellm import Router
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue