test(integration): move legacy proxy, router and Redis tests into tests/integration (#44128)

* test(integration): move legacy proxy, router and Redis tests into tests/integration

Port 39 legacy tests to the integration tier that owns them, running against
the scripted upstream, local Postgres and Redis, test-owned wire peers and
owned proxies. Delete 5 legacy tests whose contract is already owned by an
existing integration test, and remove the legacy functions, files and helpers
left unused.

* test(integration): cover recovery of a spent key after its budget is raised
This commit is contained in:
yuneng-jiang 2026-10-01 23:00:54 -07:00 • committed by GitHub
parent 826b21aab6
commit a93c396a5c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
43 changed files with 1830 additions and 3418 deletions

View file

@ -14,78 +14,6 @@ from litellm.caching.caching import DualCache
from litellm.exceptions import BlockedPiiEntityError
@pytest.mark.asyncio
async def test_presidio_with_entities_config():
"""Test for Presidio guardrail with entities config - requires actual Presidio API"""
# Setup the guardrail with specific entities config
litellm._turn_on_debug()
pii_entities_config = {
PiiEntityType.CREDIT_CARD: PiiAction.MASK,
PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK,
}
presidio_guardrail = _OPTIONAL_PresidioPIIMasking(
pii_entities_config=pii_entities_config,
presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"),
presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"),
)
# Test text with different PII types
test_text = "My credit card number is 4111-1111-1111-1111, my email is test@example.com, and my phone is 555-123-4567"
# Test the analyze request configuration
analyze_request = presidio_guardrail._get_presidio_analyze_request_payload(
text=test_text, presidio_config=None, request_data={}
)
# Verify entities were passed correctly
assert "entities" in analyze_request
assert set(analyze_request["entities"]) == set(pii_entities_config.keys())
# Test the check_pii method - this will call the actual Presidio API
redacted_text = await presidio_guardrail.check_pii(
text=test_text, output_parse_pii=True, presidio_config=None, request_data={}
)
# Verify PII has been masked/replaced/redacted in the result
assert "4111-1111-1111-1111" not in redacted_text
assert "test@example.com" not in redacted_text
# Since this entity is not in the config, it should not be masked
assert "555-123-4567" in redacted_text
# The specific replacements will vary based on Presidio's implementation
print(f"Redacted text: {redacted_text}")
@pytest.mark.asyncio
async def test_presidio_apply_guardrail():
"""Test for Presidio guardrail apply guardrail - requires actual Presidio API"""
litellm._turn_on_debug()
presidio_guardrail = _OPTIONAL_PresidioPIIMasking(
pii_entities_config={},
presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"),
presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"),
)
test_text = (
"My credit card number is 4111-1111-1111-1111 and my email is test@example.com"
)
response = await presidio_guardrail.apply_guardrail(
inputs={"texts": [test_text]},
request_data={},
input_type="request",
)
print("response from apply guardrail for presidio: ", response)
# Extract the modified text from the response
modified_text = response["texts"][0] if response.get("texts") else ""
# assert the default config masks the credit card and email
assert "4111-1111-1111-1111" not in modified_text
assert "test@example.com" not in modified_text
@pytest.mark.asyncio
async def test_presidio_with_blocked_entities():
"""Test for Presidio guardrail with blocked entities - requires actual Presidio API"""
@ -174,58 +102,6 @@ async def test_presidio_pre_call_hook_with_blocked_entities():
assert excinfo.value.guardrail_name == presidio_guardrail.guardrail_name
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", ["completion", "acompletion"])
async def test_presidio_pre_call_hook_with_different_call_types(call_type):
"""Test for Presidio guardrail pre-call hook with both completion and acompletion call types"""
# Setup the guardrail with specific entities config
pii_entities_config = {
PiiEntityType.CREDIT_CARD: PiiAction.MASK,
PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK,
}
presidio_guardrail = _OPTIONAL_PresidioPIIMasking(
pii_entities_config=pii_entities_config,
presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"),
presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"),
)
# Create a sample request with PII data
data = {
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{
"role": "user",
"content": "My credit card is 4111-1111-1111-1111 and my email is test@example.com. My phone number is 555-123-4567",
},
],
"model": "gpt-5-mini",
}
# Mock objects needed for the pre-call hook
user_api_key_dict = UserAPIKeyAuth(api_key="test_key")
cache = DualCache()
# Call the pre-call hook with the specified call type
modified_data = await presidio_guardrail.async_pre_call_hook(
user_api_key_dict=user_api_key_dict, cache=cache, data=data, call_type=call_type
)
# Verify the messages have been modified to mask PII
assert (
modified_data["messages"][0]["content"] == "You are a helpful assistant."
) # System prompt should be unchanged
user_message = modified_data["messages"][1]["content"]
assert "4111-1111-1111-1111" not in user_message
assert "test@example.com" not in user_message
# Since this entity is not in the config, it should not be masked
assert "555-123-4567" in user_message
print(f"Modified user message for call_type={call_type}: {user_message}")
@pytest.mark.parametrize(
"base_url",
[

View file

@ -0,0 +1,33 @@
import uuid
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
def test_a_key_bound_to_a_user_id_with_no_user_row_serves_and_attributes_spend_to_that_id(gateway: Gateway) -> None:
user_id: Final = f"integration-absent-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
key: Final = scenario.key(user_id=user_id, models=[model])
assert read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id=%s', (user_id,)) == []
response: Final = gateway.chat(model, key=key, text=f"unknown user {uuid.uuid4().hex}")
assert object_value(response["usage"])["total_tokens"] == 40
rows: Final = eventually(
lambda: read_rows(
'SELECT "user", spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (str(response["id"]),)
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows[0]["user"] == user_id
assert float(str(rows[0]["spend"])) == pytest.approx(0.06)
eventually(
lambda: read_rows(
'SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (sha256(key.encode()).hexdigest(),)
),
lambda values: len(values) == 1 and float(str(values[0]["spend"])) == pytest.approx(0.06),
seconds=70,
)

View file

@ -0,0 +1,107 @@
from dataclasses import dataclass
from hashlib import sha256
from typing import Final
import httpx
from integration._support.client import Gateway, Scenario, object_value, string_value
from integration._support.database import read_rows
from pydantic import JsonValue
PERMISSION_ERROR: Final = "team_member_permission_error"
@dataclass(frozen=True, slots=True)
class Member:
team_id: str
team_key: str
member_key: str
def _member(scenario: Scenario, permissions: list[JsonValue] | None) -> Member:
team_id: Final = scenario.team() if permissions is None else scenario.team(team_member_permissions=permissions)
team_key: Final = scenario.key(team_id=team_id, metadata={"owner": "team"})
member: Final = scenario.member(team_id)
return Member(team_id, team_key, scenario.key(user_id=member))
def _team_key_row(key: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT team_id, metadata FROM "LiteLLM_VerificationToken" WHERE token = %s',
(sha256(key.encode()).hexdigest(),),
)
def _team_key_count(team_id: str) -> int:
return len(read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE team_id = %s', (team_id,)))
def _refused(response: httpx.Response, status: int, error_type: str | None = None) -> None:
assert response.status_code == status, response.text
if error_type is not None:
assert object_value(response.json()["error"])["type"] == error_type, response.text
def test_default_member_permissions_only_allow_reading_team_keys(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
member: Final = _member(scenario, None)
generated: Final = gateway.request("POST", "/key/generate", {"team_id": member.team_id}, key=member.member_key)
_refused(generated, 401, PERMISSION_ERROR)
updated: Final = gateway.request(
"POST",
"/key/update",
{"key": member.team_key, "team_id": "ATTACKER_TEAM_ID", "metadata": {"owner": "member"}},
key=member.member_key,
)
_refused(updated, 401, PERMISSION_ERROR)
_refused(gateway.request("POST", "/key/delete", {"keys": [member.team_key]}, key=member.member_key), 403)
_refused(gateway.request("POST", "/key/regenerate", {"key": member.team_key}, key=member.member_key), 401)
info: Final = gateway.request("GET", "/key/info", key=member.member_key, params={"key": member.team_key})
assert info.status_code == 200, info.text
assert object_value(info.json()["info"])["team_id"] == member.team_id
assert _team_key_row(member.team_key) == [{"team_id": member.team_id, "metadata": {"owner": "team"}}]
assert _team_key_count(member.team_id) == 1
def test_update_and_delete_permissions_let_a_member_edit_but_not_create_delete_or_regenerate(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
member: Final = _member(scenario, ["/key/update", "/key/delete", "/key/info"])
updated: Final = gateway.request(
"POST",
"/key/update",
{"key": member.team_key, "team_id": member.team_id, "metadata": {"owner": "member"}},
key=member.member_key,
)
assert updated.status_code == 200, updated.text
assert _team_key_row(member.team_key) == [{"team_id": member.team_id, "metadata": {"owner": "member"}}]
_refused(gateway.request("POST", "/key/delete", {"keys": [member.team_key]}, key=member.member_key), 403)
generated: Final = gateway.request("POST", "/key/generate", {"team_id": member.team_id}, key=member.member_key)
_refused(generated, 401, PERMISSION_ERROR)
regenerated: Final = gateway.request(
"POST", "/key/regenerate", {"key": member.team_key, "team_id": member.team_id}, key=member.member_key
)
_refused(regenerated, 401)
assert _team_key_count(member.team_id) == 1
def test_generate_permission_lets_a_member_create_team_keys_but_not_change_existing_ones(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
member: Final = _member(scenario, ["/key/generate"])
generated: Final = gateway.request("POST", "/key/generate", {"team_id": member.team_id}, key=member.member_key)
assert generated.status_code == 200, generated.text
created: Final = string_value(generated.json()["key"])
scenario.cleanups.callback(scenario.delete_key, created)
assert _team_key_row(created) == [{"team_id": member.team_id, "metadata": {}}]
updated: Final = gateway.request(
"POST",
"/key/update",
{"key": member.team_key, "team_id": member.team_id, "metadata": {"owner": "member"}},
key=member.member_key,
)
_refused(updated, 401, PERMISSION_ERROR)
assert _team_key_row(member.team_key) == [{"team_id": member.team_id, "metadata": {"owner": "team"}}]
_refused(gateway.request("POST", "/key/delete", {"keys": [member.team_key]}, key=member.member_key), 403)
regenerated: Final = gateway.request(
"POST", "/key/regenerate", {"key": member.team_key, "team_id": member.team_id}, key=member.member_key
)
_refused(regenerated, 401, PERMISSION_ERROR)
assert _team_key_count(member.team_id) == 2

View file

@ -0,0 +1,105 @@
import uuid
from collections.abc import Iterator
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, object_value, string_value
from pydantic import JsonValue
@pytest.fixture
def upstream(gateway: Gateway) -> Iterator[httpx.Client]:
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as client:
client.get("/__observations").raise_for_status()
yield client
def _observed_models(upstream: httpx.Client) -> list[JsonValue]:
observed: Final = upstream.get("/__observations")
observed.raise_for_status()
return [request["body"]["model"] for request in observed.json()["requests"]]
def _chat(gateway: Gateway, model: str, key: str) -> httpx.Response:
return gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"team model {uuid.uuid4().hex}"}]},
key=key,
)
def _ids(listing: dict[str, JsonValue]) -> set[JsonValue]:
data: Final = listing["data"]
assert isinstance(data, list)
return {object_value(entry)["id"] for entry in data}
def test_a_model_created_for_a_team_is_listed_in_that_teams_models(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team(models=[])
other: Final = scenario.team(models=[])
model: Final = scenario.model(model_info={"team_id": team})
own_models: Final = object_value(gateway.get("/team/info", {"team_id": team})["team_info"])["models"]
other_models: Final = object_value(gateway.get("/team/info", {"team_id": other})["team_info"])["models"]
assert isinstance(own_models, list) and isinstance(other_models, list)
assert model in own_models
assert model not in other_models
def test_a_team_model_is_listed_and_served_only_for_keys_of_its_team(gateway: Gateway, upstream: httpx.Client) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team(models=[])
other: Final = scenario.team(models=[])
provider_model: Final = f"team-scoped-{uuid.uuid4().hex}"
model: Final = scenario.model(model=f"openai/{provider_model}", model_info={"team_id": team})
team_key: Final = scenario.key(team_id=team)
other_key: Final = scenario.key(team_id=other)
assert model in _ids(object_value(gateway.request("GET", "/models", key=team_key).json()))
assert model not in _ids(object_value(gateway.request("GET", "/models", key=other_key).json()))
served: Final = _chat(gateway, model, team_key)
assert served.status_code == 200, served.text
refused: Final = _chat(gateway, model, other_key)
assert refused.status_code == 400, refused.text
assert _observed_models(upstream) == [provider_model]
def _v2_team_public_names(gateway: Gateway, key: str, model: str) -> list[JsonValue]:
response: Final = gateway.request("GET", "/v2/model/info", key=key, params={"model_name": model})
assert response.status_code == 200, response.text
return [entry["model_info"].get("team_public_model_name") for entry in response.json()["data"]]
def test_v2_model_info_reports_a_team_model_to_team_and_non_team_keys(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team(models=[])
model: Final = scenario.model(model_info={"team_id": team})
assert _v2_team_public_names(gateway, scenario.key(team_id=team), model) == [model]
assert _v2_team_public_names(gateway, scenario.key(), model) == [model]
@pytest.mark.parametrize("set_on", ["team_new", "team_update"])
def test_team_model_alias_routes_a_team_key_to_its_target(
gateway: Gateway, upstream: httpx.Client, set_on: str
) -> None:
with gateway.scenario() as scenario:
provider_model: Final = f"team-alias-{uuid.uuid4().hex}"
model: Final = scenario.model(model=f"openai/{provider_model}")
alias: Final = f"alias-{uuid.uuid4().hex}"
team: Final = (
scenario.team(models=[model], model_aliases={alias: model})
if set_on == "team_new"
else scenario.team(models=[model])
)
if set_on == "team_update":
gateway.post("/team/update", {"team_id": team, "model_aliases": {alias: model}})
key: Final = scenario.key(team_id=team, models=[model])
response: Final = _chat(gateway, alias, key)
assert response.status_code == 200, response.text
assert string_value(response.json()["model"]) == alias
assert _observed_models(upstream) == [provider_model]
unaliased: Final = _chat(gateway, f"alias-{uuid.uuid4().hex}", key)
assert unaliased.status_code == 403, unaliased.text
assert unaliased.json()["error"]["type"] == "key_model_access_denied"
assert _observed_models(upstream) == []

View file

@ -0,0 +1,90 @@
import uuid
from collections.abc import Iterator
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, gateway_from_environment
from integration._support.process import owned_proxy
from pydantic import JsonValue
pytestmark: Final = pytest.mark.timeout(180)
def _deployment(model_name: str, model: str, upstream_url: str) -> dict[str, JsonValue]:
return {
"model_name": model_name,
"litellm_params": {"model": model, "api_base": f"{upstream_url}/v1", "api_key": "synthetic-wildcard-key"},
}
@pytest.fixture(scope="module")
def candidate(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("wildcard-access")
with gateway_from_environment() as base:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = [
_deployment("*", "openai/*", base.upstream_url),
_deployment("anthropic/*", "openai/*", base.upstream_url),
_deployment("groq/*", "openai/*", base.upstream_url),
_deployment("good-model", "openai/good-model-upstream", base.upstream_url),
]
path: Final = directory / "wildcard-access.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(base, directory, {}, config=path) as proxy:
yield proxy
@pytest.fixture
def upstream(candidate: Gateway) -> Iterator[httpx.Client]:
with httpx.Client(base_url=candidate.upstream_url, timeout=5, trust_env=False) as client:
client.get("/__observations").raise_for_status()
yield client
def _chat(gateway: Gateway, model: str, key: str) -> httpx.Response:
return gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"wildcard {uuid.uuid4().hex}"}]},
key=key,
)
def _observed_models(upstream: httpx.Client) -> list[JsonValue]:
return [request["body"]["model"] for request in upstream.get("/__observations").json()["requests"]]
def test_an_all_models_key_reaches_a_model_served_only_by_the_catch_all_deployment(
candidate: Gateway, upstream: httpx.Client
) -> None:
with candidate.scenario() as scenario:
key: Final = scenario.key(models=["*"])
unlisted: Final = f"unlisted-{uuid.uuid4().hex}"
response: Final = _chat(candidate, unlisted, key)
assert response.status_code == 200, response.text
assert _observed_models(upstream) == [unlisted]
def test_a_key_without_models_inherits_the_users_exact_and_wildcard_grants(
candidate: Gateway, upstream: httpx.Client
) -> None:
with candidate.scenario() as scenario:
user_id: Final = scenario.user(models=["good-model", "anthropic/*"])
key: Final = scenario.key(user_id=user_id, models=[])
wildcard_model: Final = f"claude-{uuid.uuid4().hex}"
assert _chat(candidate, f"anthropic/{wildcard_model}", key).status_code == 200
assert _chat(candidate, "good-model", key).status_code == 200
assert _observed_models(upstream) == [wildcard_model, "good-model-upstream"]
denied: Final = tuple(
_chat(candidate, outside, key)
for outside in (f"groq/{wildcard_model}", f"bedrock/anthropic.{wildcard_model}")
)
assert [(response.status_code, response.json()["error"]["type"]) for response in denied] == [
(403, "user_model_access_denied")
] * 2, [response.text for response in denied]
assert _observed_models(upstream) == []
assert _chat(candidate, f"groq/{wildcard_model}", candidate.key).status_code == 200
assert _observed_models(upstream) == [wildcard_model]

View file

@ -0,0 +1,34 @@
import uuid
from typing import Final
import httpx
from integration._support.client import Gateway, object_value
def test_health_check_of_a_model_added_through_the_api_calls_its_upstream_and_reports_it_healthy(
gateway: Gateway,
) -> None:
with (
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
provider_model: Final = f"health-{uuid.uuid4().hex}"
model: Final = scenario.model(model=f"openai/{provider_model}")
key: Final = scenario.key(models=[model])
listed: Final = gateway.request("GET", "/v2/model/info", key=key, params={"model": model})
assert listed.status_code == 200, listed.text
assert [entry["model_name"] for entry in listed.json()["data"]] == [model]
assert (
object_value(gateway.chat(model, key=key, text=f"health {uuid.uuid4().hex}")["usage"])["total_tokens"] == 40
)
upstream.get("/__observations").raise_for_status()
health: Final = gateway.request("GET", "/health", params={"model": model})
assert health.status_code == 200, health.text
report: Final = health.json()
assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report
assert [(endpoint["model"], endpoint["api_base"]) for endpoint in report["healthy_endpoints"]] == [
(f"openai/{provider_model}", f"{gateway.upstream_url}/v1")
]
assert [request["body"]["model"] for request in upstream.get("/__observations").json()["requests"]] == [
provider_model
]

View file

@ -0,0 +1,89 @@
import uuid
from concurrent.futures import ThreadPoolExecutor
from typing import Final
import httpx
from integration._support.client import Gateway, object_value, string_value
from integration._support.database import read_rows
from pydantic import JsonValue
CONCURRENT_CREATES: Final = 8
def _membership_rows(organization_id: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT user_id, user_role FROM "LiteLLM_OrganizationMembership" WHERE organization_id = %s',
(organization_id,),
)
def _listed(gateway: Gateway, organization_id: str) -> dict[str, JsonValue]:
response: Final = gateway.request("GET", "/organization/list")
assert response.status_code == 200, response.text
entries: Final = response.json()
assert isinstance(entries, list)
matches: Final = [entry for entry in entries if entry["organization_id"] == organization_id]
assert len(matches) == 1, f"{organization_id} listed {len(matches)} times"
return object_value(matches[0])
def test_concurrent_creates_with_one_alias_each_persist_a_distinct_organization(gateway: Gateway) -> None:
alias: Final = f"integration-{uuid.uuid4().hex}"
def create(_: int) -> httpx.Response:
return gateway.request("POST", "/organization/new", {"organization_alias": alias})
with ThreadPoolExecutor(max_workers=CONCURRENT_CREATES) as pool:
responses: Final = tuple(pool.map(create, range(CONCURRENT_CREATES)))
created: Final = tuple(response.json() for response in responses if response.status_code == 200)
with gateway.scenario() as scenario:
for body in created:
scenario.cleanups.callback(
scenario.delete_organization, string_value(body["organization_id"]), string_value(body["budget_id"])
)
assert [response.status_code for response in responses] == [200] * CONCURRENT_CREATES, [
response.text for response in responses
]
identities: Final = {string_value(body["organization_id"]) for body in created}
assert len(identities) == CONCURRENT_CREATES
rows: Final = read_rows(
'SELECT organization_id FROM "LiteLLM_OrganizationTable" WHERE organization_alias = %s', (alias,)
)
assert {string_value(row["organization_id"]) for row in rows} == identities
def test_list_returns_each_organization_with_its_budget_and_members(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
organization_id: Final = scenario.organization(max_budget=3.5, tpm_limit=120)
member: Final = scenario.org_member(organization_id, role="internal_user")
listed: Final = _listed(gateway, organization_id)
budget: Final = object_value(listed["litellm_budget_table"])
assert (budget["max_budget"], budget["tpm_limit"]) == (3.5, 120)
members: Final = listed["members"]
assert isinstance(members, list)
assert [(object_value(entry)["user_id"], object_value(entry)["user_role"]) for entry in members] == [
(member, "internal_user")
]
def test_member_role_update_and_removal_persist_and_read_back(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
organization_id: Final = scenario.organization()
member: Final = scenario.org_member(organization_id, role="internal_user")
assert _membership_rows(organization_id) == [{"user_id": member, "user_role": "internal_user"}]
updated: Final = gateway.request(
"PATCH",
"/organization/member_update",
{"organization_id": organization_id, "user_id": member, "role": "org_admin"},
)
assert updated.status_code == 200, updated.text
assert _membership_rows(organization_id) == [{"user_id": member, "user_role": "org_admin"}]
info_members: Final = gateway.get("/organization/info", {"organization_id": organization_id})["members"]
assert isinstance(info_members, list)
assert [object_value(entry)["user_role"] for entry in info_members] == ["org_admin"]
removed: Final = gateway.request(
"DELETE", "/organization/member_delete", {"organization_id": organization_id, "user_id": member}
)
assert removed.status_code == 200, removed.text
assert _membership_rows(organization_id) == []
assert _listed(gateway, organization_id)["members"] == []

View file

@ -0,0 +1,53 @@
import uuid
from collections.abc import Callable, Iterator, Mapping
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, gateway_from_environment
from integration._support.otlp_sink import Span, SpanSinks, recorded_spans
from integration._support.process import owned_proxy
from pydantic import JsonValue
pytestmark: Final = pytest.mark.timeout(180)
AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path]
@pytest.fixture(scope="module")
def gateway(audit_sinks: SpanSinks) -> Iterator[Gateway]:
with gateway_from_environment() as base:
yield base
def _traces(spans: tuple[Span, ...]) -> dict[str, frozenset[str]]:
trace_ids: Final = {span["trace_id"] for span in spans}
return {trace: frozenset(span["name"] for span in spans if span["trace_id"] == trace) for trace in trace_ids}
def test_default_otel_logger_puts_datastore_model_and_spend_writer_spans_in_the_request_trace(
gateway: Gateway, audit_sinks: SpanSinks, otel_audit_config: AuditConfigWriter, tmp_path: Path
) -> None:
config: Final = otel_audit_config(tmp_path, {})
overrides: Final = {"OTEL_EXPORTER": "http/json", "OTEL_ENDPOINT": audit_sinks.operator}
with owned_proxy(gateway, tmp_path, overrides, config=config) as candidate, candidate.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
start, _ = recorded_spans(audit_sinks.operator)
response: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"otel v1 {uuid.uuid4().hex}"}]},
key=key,
)
assert response.status_code == 200, response.text
expected: Final = frozenset({"postgres", "redis", "raw_gen_ai_request", "batch_write_to_db"})
traces: Final = eventually(
lambda: _traces(recorded_spans(audit_sinks.operator, start)[1]),
lambda grouped: any(expected <= names for names in grouped.values()),
seconds=60,
return_last_on_timeout=True,
)
assert any(expected <= names for names in traces.values()), {
trace: sorted(names) for trace, names in traces.items()
}

View file

@ -0,0 +1,141 @@
import json
import re
import uuid
from collections.abc import Iterator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from itertools import chain
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, string_value
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
CARD: Final = "4111-1111-1111-1111"
EMAIL: Final = "jane.doe@example.com"
PHONE: Final = "555-123-4567"
SYSTEM_PROMPT: Final = "You are a helpful assistant."
RECOGNIZERS: Final = {
"CREDIT_CARD": re.escape(CARD),
"EMAIL_ADDRESS": re.escape(EMAIL),
"PHONE_NUMBER": re.escape(PHONE),
}
def _detect(entity: str, text: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(
{"entity_type": entity, "start": match.start(), "end": match.end(), "score": 0.95}
for match in re.finditer(RECOGNIZERS[entity], text)
)
def _analyze(request: Request) -> Reply:
assert request.target == "/analyze", request.target
body: Final = json.loads(request.body)
requested: Final = body.get("entities") or list(RECOGNIZERS)
findings: Final = list(chain.from_iterable(_detect(entity, body["text"]) for entity in requested))
return Reply(body=json.dumps(findings).encode())
def _anonymize(request: Request) -> Reply:
assert request.target == "/anonymize", request.target
body: Final = json.loads(request.body)
spans: Final = sorted(body["analyzer_results"], key=lambda item: item["start"])
pieces: Final = [
body["text"][(spans[index - 1]["end"] if index else 0) : span["start"]] + f"<{span['entity_type']}>"
for index, span in enumerate(spans)
]
tail: Final = body["text"][spans[-1]["end"] :] if spans else body["text"]
return Reply(body=json.dumps({"text": "".join(pieces) + tail, "items": []}).encode())
@dataclass(frozen=True, slots=True)
class Presidio:
name: str
analyzer: Wire
anonymizer: Wire
@contextmanager
def _presidio(gateway: Gateway, mode: str, entities: Mapping[str, str] | None) -> Iterator[Presidio]:
name: Final = f"presidio-{uuid.uuid4().hex}"
with wire_server(_analyze) as analyzer, wire_server(_anonymize) as anonymizer:
created: Final = gateway.request(
"POST",
"/guardrails",
{
"guardrail": {
"guardrail_name": name,
"litellm_params": {
"guardrail": "presidio",
"mode": mode,
"default_on": False,
"presidio_analyzer_api_base": analyzer.url,
"presidio_anonymizer_api_base": anonymizer.url,
**({} if entities is None else {"pii_entities_config": dict(entities)}),
},
}
},
)
assert created.status_code == 200, created.text
try:
yield Presidio(name, analyzer, anonymizer)
finally:
deleted: Final = gateway.request("DELETE", f"/guardrails/{created.json()['guardrail_id']}")
assert deleted.status_code == 200, deleted.text
def _requested_entities(analyzer: Wire) -> list[JsonValue]:
return [json.loads(request.body).get("entities") for request in analyzer.drain()]
def test_pre_call_masks_only_the_configured_entities_before_the_provider_sees_the_prompt(gateway: Gateway) -> None:
with (
_presidio(gateway, "pre_call", {"CREDIT_CARD": "MASK", "EMAIL_ADDRESS": "MASK"}) as presidio,
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
model: Final = scenario.model()
upstream.get("/__observations").raise_for_status()
user_text: Final = f"{uuid.uuid4().hex} card {CARD}, email {EMAIL}, phone {PHONE}"
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"guardrails": [presidio.name],
"messages": [{"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user_text}],
},
)
assert response.status_code == 200, response.text
observed: Final = upstream.get("/__observations").json()["requests"]
assert len(observed) == 1
messages: Final = observed[0]["body"]["messages"]
assert messages[0] == {"role": "system", "content": SYSTEM_PROMPT}
forwarded: Final = string_value(messages[1]["content"])
assert CARD not in forwarded and EMAIL not in forwarded, forwarded
assert "<CREDIT_CARD>" in forwarded and "<EMAIL_ADDRESS>" in forwarded, forwarded
assert PHONE in forwarded, forwarded
requested: Final = _requested_entities(presidio.analyzer)
assert requested and all(sorted(entities) == ["CREDIT_CARD", "EMAIL_ADDRESS"] for entities in requested), (
requested
)
@pytest.mark.parametrize("entities", [None, {}])
def test_apply_guardrail_with_the_default_config_masks_every_detected_entity(
gateway: Gateway, entities: Mapping[str, str] | None
) -> None:
with _presidio(gateway, "pre_call", entities) as presidio:
response: Final = gateway.request(
"POST",
"/guardrails/apply_guardrail",
{"guardrail_name": presidio.name, "text": f"card {CARD} and email {EMAIL}"},
)
assert response.status_code == 200, response.text
masked: Final = string_value(response.json()["response_text"])
assert masked == "card <CREDIT_CARD> and email <EMAIL_ADDRESS>", masked
assert _requested_entities(presidio.analyzer) == [None]
assert len(presidio.anonymizer.drain()) == 1

View file

@ -0,0 +1,209 @@
import asyncio
import itertools
import json
import ssl
import threading
import uuid
from collections.abc import Iterator
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from typing import Final
import pytest
import websockets
from integration._support.client import Gateway, gateway_from_environment
from integration._support.process import owned_proxy
from integration._support.tls import server_context, write_self_signed_cert
from pydantic import JsonValue
from websockets.asyncio.server import ServerConnection, serve
pytestmark: Final = pytest.mark.timeout(180)
PROVIDER_MODEL: Final = "ws-peer-model"
USAGE: Final = {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7}
TERMINAL: Final = frozenset({"response.completed", "response.failed", "error"})
@dataclass(frozen=True, slots=True)
class Peer:
url: str
paths: SimpleQueue[str]
frames: SimpleQueue[dict[str, JsonValue]]
def _events(response_id: str, text: str) -> tuple[dict[str, JsonValue], ...]:
message: Final = {
"type": "message",
"id": f"msg_{response_id}",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": text, "annotations": []}],
}
response: Final = {"id": response_id, "object": "response", "created_at": 1700000000, "model": PROVIDER_MODEL}
return (
{"type": "response.created", "response": {**response, "status": "in_progress", "output": []}},
{
"type": "response.output_text.delta",
"item_id": f"msg_{response_id}",
"output_index": 0,
"content_index": 0,
"delta": text,
},
{
"type": "response.completed",
"response": {**response, "status": "completed", "output": [message], "usage": USAGE},
},
)
async def _answer(
connection: ServerConnection, paths: SimpleQueue[str], frames: SimpleQueue[dict[str, JsonValue]]
) -> None:
paths.put(connection.request.path if connection.request is not None else "")
turns: Final = itertools.count(1)
async for raw in connection:
frame: Final = json.loads(raw)
frames.put(frame)
if frame.get("type") != "response.create":
continue
for event in _events(f"resp_peer_{next(turns)}", "seven"):
await connection.send(json.dumps(event))
async def _serve(
tls: ssl.SSLContext,
paths: SimpleQueue[str],
frames: SimpleQueue[dict[str, JsonValue]],
ports: SimpleQueue[int],
stop: asyncio.Event,
) -> None:
async with serve(lambda connection: _answer(connection, paths, frames), "127.0.0.1", 0, ssl=tls) as server:
ports.put(next(iter(server.sockets)).getsockname()[1])
await stop.wait()
@contextmanager
def responses_peer(cert: tuple[Path, Path]) -> Iterator[Peer]:
loop: Final = asyncio.new_event_loop()
stop: Final = asyncio.Event()
paths: Final = SimpleQueue[str]()
frames: Final = SimpleQueue[dict[str, JsonValue]]()
ports: Final = SimpleQueue[int]()
thread: Final = threading.Thread(
target=loop.run_until_complete, args=(_serve(server_context(*cert), paths, frames, ports, stop),), daemon=True
)
thread.start()
try:
yield Peer(f"https://127.0.0.1:{ports.get(timeout=10)}/v1", paths, frames)
finally:
loop.call_soon_threadsafe(stop.set)
thread.join(timeout=10)
loop.close()
def _create(model: str, text: str, previous_response_id: str | None = None) -> str:
return json.dumps(
{
"type": "response.create",
"model": model,
"store": True,
"input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]}],
**({} if previous_response_id is None else {"previous_response_id": previous_response_id}),
}
)
async def _turn(connection: websockets.ClientConnection, frame: str) -> tuple[dict[str, JsonValue], ...]:
await connection.send(frame)
return await _until_terminal(connection, ())
async def _until_terminal(
connection: websockets.ClientConnection, received: tuple[dict[str, JsonValue], ...]
) -> tuple[dict[str, JsonValue], ...]:
event: Final = json.loads(await asyncio.wait_for(connection.recv(), timeout=20))
collected: Final = (*received, event)
if event.get("type") in TERMINAL or len(collected) >= 50:
return collected
return await _until_terminal(connection, collected)
async def _session(
proxy_url: str, key: str, model: str, texts: tuple[str, ...]
) -> tuple[tuple[dict[str, JsonValue], ...], ...]:
proxy: Final = proxy_url.rstrip("/").replace("http://", "ws://")
async with websockets.connect(
f"{proxy}/v1/responses?model={model}",
additional_headers={"Authorization": f"Bearer {key}"},
open_timeout=10,
) as connection:
first: Final = await _turn(connection, _create(model, texts[0]))
if len(texts) == 1:
return (first,)
previous: Final = str(first[-1]["response"]["id"])
second: Final = await _turn(connection, _create(model, texts[1], previous))
return (first, second)
def _completed(events: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]:
assert events[-1]["type"] == "response.completed", [event.get("type") for event in events]
return events[-1]["response"]
@pytest.fixture(scope="module")
def cert(tmp_path_factory: pytest.TempPathFactory) -> tuple[Path, Path]:
return write_self_signed_cert(tmp_path_factory.mktemp("responses-ws-cert"))
@pytest.fixture(scope="module")
def candidate(tmp_path_factory: pytest.TempPathFactory, cert: tuple[Path, Path]) -> Iterator[Gateway]:
with gateway_from_environment() as base:
with owned_proxy(base, tmp_path_factory.mktemp("responses-ws"), {"SSL_CERT_FILE": str(cert[0])}) as proxy:
yield proxy
def _drain(queue: SimpleQueue[dict[str, JsonValue]]) -> tuple[dict[str, JsonValue], ...]:
return tuple(queue.get_nowait() for _ in range(queue.qsize()))
def test_a_response_create_frame_streams_from_the_provider_socket_back_to_the_client(
candidate: Gateway, cert: tuple[Path, Path]
) -> None:
with responses_peer(cert) as peer, candidate.scenario() as scenario:
model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url)
key: Final = scenario.key(models=[model])
text: Final = f"say seven {uuid.uuid4().hex}"
(events,) = asyncio.run(_session(str(candidate.client.base_url), key, model, (text,)))
assert [event["type"] for event in events] == [
"response.created",
"response.output_text.delta",
"response.completed",
]
completed: Final = _completed(events)
assert completed["status"] == "completed"
assert completed["usage"] == USAGE
assert peer.paths.get_nowait() == f"/v1/responses?model={PROVIDER_MODEL}"
(forwarded,) = _drain(peer.frames)
assert forwarded["type"] == "response.create"
assert forwarded["model"] == PROVIDER_MODEL
assert forwarded["input"][0]["content"][0]["text"] == text
def test_previous_response_id_from_the_first_turn_reaches_the_provider_as_its_own_id(
candidate: Gateway, cert: tuple[Path, Path]
) -> None:
with responses_peer(cert) as peer, candidate.scenario() as scenario:
model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url)
key: Final = scenario.key(models=[model])
texts: Final = (f"remember seven {uuid.uuid4().hex}", f"which number {uuid.uuid4().hex}")
first, second = asyncio.run(_session(str(candidate.client.base_url), key, model, texts))
assert _completed(first)["status"] == "completed"
assert str(_completed(first)["id"]).startswith("resp_") and _completed(first)["id"] != "resp_peer_1"
assert _completed(second)["status"] == "completed"
assert peer.paths.qsize() == 1
forwarded: Final = _drain(peer.frames)
assert [frame["input"][0]["content"][0]["text"] for frame in forwarded] == list(texts)
assert "previous_response_id" not in forwarded[0]
assert forwarded[1]["previous_response_id"] == "resp_peer_1"

View file

@ -0,0 +1,72 @@
import uuid
from collections.abc import Iterator
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, gateway_from_environment
from integration._support.process import owned_proxy
pytestmark: Final = pytest.mark.timeout(180)
MODEL: Final = "regional-model"
UPSTREAM_BY_REGION: Final = {"eu": "regional-eu-upstream", "us": "regional-us-upstream"}
CALLS: Final = 5
@pytest.fixture(scope="module")
def candidate(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("region-routing")
with gateway_from_environment() as base:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = [
{
"model_name": MODEL,
"litellm_params": {
"model": f"openai/{upstream}",
"api_base": f"{base.upstream_url}/v1",
"api_key": "synthetic-region-key",
"region_name": region,
},
}
for region, upstream in UPSTREAM_BY_REGION.items()
]
config["router_settings"] = {**config["router_settings"], "enable_pre_call_checks": True}
path: Final = directory / "region-routing.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(base, directory, {}, config=path) as proxy:
yield proxy
@pytest.mark.parametrize("region", ["eu", "us"])
def test_an_end_users_allowed_region_pins_every_call_to_that_regions_deployment(
candidate: Gateway, region: str
) -> None:
with (
candidate.scenario() as scenario,
httpx.Client(base_url=candidate.upstream_url, timeout=5, trust_env=False) as upstream,
):
end_user: Final = f"integration-end-user-{uuid.uuid4().hex}"
candidate.post("/end_user/new", {"user_id": end_user, "allowed_model_region": region})
scenario.cleanups.callback(candidate.post, "/end_user/delete", {"user_ids": [end_user]})
key: Final = scenario.key(models=[MODEL])
upstream.get("/__observations").raise_for_status()
responses: Final = tuple(
candidate.request(
"POST",
"/v1/chat/completions",
{
"model": MODEL,
"user": end_user,
"messages": [{"role": "user", "content": f"region {uuid.uuid4().hex}"}],
},
key=key,
)
for _ in range(CALLS)
)
assert [response.status_code for response in responses] == [200] * CALLS, [r.text for r in responses]
assert [response.headers.get("x-litellm-model-region") for response in responses] == [region] * CALLS
observed: Final = upstream.get("/__observations").json()["requests"]
assert [request["body"]["model"] for request in observed] == [UPSTREAM_BY_REGION[region]] * CALLS

View file

@ -0,0 +1,27 @@
import uuid
from typing import Final
import httpx
from integration._support.client import Gateway, object_value
def test_zero_parallel_slots_refuse_before_the_provider_and_one_slot_serves(gateway: Gateway) -> None:
with (
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
model: Final = scenario.model()
blocked: Final = scenario.key(models=[model], max_parallel_requests=0)
allowed: Final = scenario.key(models=[model], max_parallel_requests=1)
upstream.get("/__observations").raise_for_status()
refused: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"no slots {uuid.uuid4().hex}"}]},
key=blocked,
)
assert refused.status_code == 429, refused.text
assert upstream.get("/__observations").json()["requests"] == []
served: Final = tuple(gateway.chat(model, key=allowed, text=f"one slot {uuid.uuid4().hex}") for _ in range(2))
assert [object_value(response["usage"])["total_tokens"] for response in served] == [40, 40]
assert len(upstream.get("/__observations").json()["requests"]) == 2

View file

@ -0,0 +1,67 @@
import uuid
from collections.abc import Iterator
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, gateway_from_environment
from integration._support.process import owned_proxy
pytestmark: Final = pytest.mark.timeout(180)
MODEL: Final = "tagged-model"
DEPLOYMENT_BY_TAG: Final = {"teamA": "team-a-deployment", "teamB": "team-b-deployment"}
CALLS: Final = 5
@pytest.fixture(scope="module")
def candidate(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("team-tag-routing")
with gateway_from_environment() as base:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = [
{
"model_name": MODEL,
"litellm_params": {
"model": f"openai/{deployment}",
"api_base": f"{base.upstream_url}/v1",
"api_key": "synthetic-tag-key",
"tags": [tag],
},
"model_info": {"id": deployment},
}
for tag, deployment in DEPLOYMENT_BY_TAG.items()
]
config["router_settings"] = {**config["router_settings"], "enable_tag_filtering": True}
path: Final = directory / "team-tag-routing.yaml"
path.write_text(yaml.safe_dump(config))
with owned_proxy(base, directory, {}, config=path) as proxy:
yield proxy
@pytest.mark.parametrize("tag", ["teamA", "teamB"])
def test_a_teams_tags_route_every_call_of_its_keys_to_the_matching_deployment(candidate: Gateway, tag: str) -> None:
with (
candidate.scenario() as scenario,
httpx.Client(base_url=candidate.upstream_url, timeout=5, trust_env=False) as upstream,
):
team_id: Final = scenario.team(tags=[tag])
key: Final = scenario.key(team_id=team_id)
upstream.get("/__observations").raise_for_status()
responses: Final = tuple(
candidate.request(
"POST",
"/v1/chat/completions",
{"model": MODEL, "messages": [{"role": "user", "content": f"tagged {uuid.uuid4().hex}"}]},
key=key,
)
for _ in range(CALLS)
)
assert [response.status_code for response in responses] == [200] * CALLS, [r.text for r in responses]
assert [response.headers.get("x-litellm-model-id") for response in responses] == [
DEPLOYMENT_BY_TAG[tag]
] * CALLS
observed: Final = upstream.get("/__observations").json()["requests"]
assert [request["body"]["model"] for request in observed] == [DEPLOYMENT_BY_TAG[tag]] * CALLS

View file

@ -0,0 +1,74 @@
import asyncio
import os
from collections.abc import Iterator
from datetime import datetime, timedelta, timezone
from typing import Final
import litellm
import pytest
from litellm.caching.dual_cache import DualCache
from litellm.caching.redis_cache import RedisCache
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.types.utils import BudgetConfig
from redis import Redis
WINDOWS: Final = {"openai": ("1d", 86400), "vertex_ai": ("1h", 3600)}
SPEND_KEYS: Final = {provider: f"provider_spend:{provider}:{window}" for provider, (window, _) in WINDOWS.items()}
START_KEYS: Final = tuple(f"provider_budget_start_time:{provider}" for provider in WINDOWS)
@pytest.fixture
def redis_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[Redis]:
monkeypatch.setattr(litellm, "callbacks", [])
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), decode_responses=True) as client:
client.delete(*SPEND_KEYS.values(), *START_KEYS)
yield client
client.delete(*SPEND_KEYS.values(), *START_KEYS)
def _limiter() -> RouterBudgetLimiting:
return RouterBudgetLimiting(
dual_cache=DualCache(redis_cache=RedisCache(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]))),
provider_budget_config={
provider: BudgetConfig(budget_duration=window, max_budget=100) for provider, (window, _) in WINDOWS.items()
},
)
async def _windows_opened(redis_client: Redis) -> bool:
for _ in range(100):
if all(int(redis_client.ttl(key)) > 0 for key in SPEND_KEYS.values()):
return True
await asyncio.sleep(0.1)
return False
@pytest.mark.asyncio
async def test_spend_written_to_redis_by_another_instance_is_pulled_into_memory(redis_client: Redis) -> None:
limiter: Final = _limiter()
assert await _windows_opened(redis_client)
elsewhere: Final = {SPEND_KEYS["openai"]: 50.0, SPEND_KEYS["vertex_ai"]: 75.0}
for key, value in elsewhere.items():
redis_client.set(key, str(value), keepttl=True)
await limiter._sync_in_memory_spend_with_redis()
in_memory: Final = {key: await limiter.dual_cache.in_memory_cache.async_get_cache(key) for key in elsewhere}
assert in_memory == elsewhere
assert await limiter._get_current_provider_spend("openai") == 50.0
@pytest.mark.asyncio
async def test_budget_reset_time_follows_the_redis_window_expiry(redis_client: Redis) -> None:
limiter: Final = _limiter()
assert await _windows_opened(redis_client)
assert await limiter._get_current_provider_budget_reset_at("anthropic") is None
reset_times: Final = {
provider: await limiter._get_current_provider_budget_reset_at(provider) for provider in WINDOWS
}
now: Final = datetime.now(timezone.utc)
drift: Final = {
provider: abs(
(datetime.fromisoformat(str(reset_times[provider])) - (now + timedelta(seconds=seconds))).total_seconds()
)
for provider, (_, seconds) in WINDOWS.items()
}
assert all(seconds < 5 for seconds in drift.values()), (reset_times, drift)

View file

@ -0,0 +1,81 @@
import json
import os
import uuid
from itertools import chain
from typing import Final
import litellm
import pytest
from integration._support.wire import Reply, Request, wire_server
from litellm import Router
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from prometheus_client import REGISTRY
CHAT_RESPONSE: Final = json.dumps(
{
"id": "chatcmpl_redis_service_metrics",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "metrics"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
}
).encode()
LABELS: Final = {"redis": "redis"}
def _reply(request: Request) -> Reply:
return Reply(body=CHAT_RESPONSE)
def _redis_metrics() -> tuple[float, float, float]:
failed_metrics: Final = tuple(
metric for metric in REGISTRY.collect() if metric.name == "litellm_redis_failed_requests"
)
samples: Final = chain.from_iterable(metric.samples for metric in failed_metrics)
failed: Final = sum(sample.value for sample in samples if sample.name.endswith("_total"))
return (
REGISTRY.get_sample_value("litellm_redis_total_requests_total", LABELS) or 0.0,
REGISTRY.get_sample_value("litellm_redis_latency_count", LABELS) or 0.0,
failed,
)
@pytest.mark.asyncio
async def test_router_redis_traffic_is_counted_in_the_prometheus_service_metrics(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(litellm, "service_callback", ["prometheus_system"])
with wire_server(_reply) as wire:
router: Final = Router(
model_list=[
{
"model_name": "redis-metrics",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": f"{wire.url}/v1",
"api_key": "synthetic-redis-metrics-key",
"tpm": tpm,
},
}
for tpm in (100, 1000)
],
routing_strategy="usage-based-routing-v2",
redis_host=os.environ["REDIS_HOST"],
redis_port=int(os.environ["REDIS_PORT"]),
)
before: Final = _redis_metrics()
responses: Final = [
await router.acompletion(
model="redis-metrics", messages=[{"role": "user", "content": f"metrics {uuid.uuid4().hex}"}]
)
for _ in range(2)
]
await GLOBAL_LOGGING_WORKER.flush()
after: Final = _redis_metrics()
assert [response.usage.total_tokens for response in responses] == [7, 7]
assert len(wire.drain()) == 2
total_delta, latency_delta, failed_delta = (now - then for now, then in zip(after, before, strict=True))
assert total_delta > 0, (before, after)
assert latency_delta > 0, (before, after)
assert failed_delta == 0, (before, after)

View file

@ -0,0 +1,107 @@
import os
import socket
import ssl
import threading
import uuid
from collections.abc import Iterator
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from typing import Final
import pytest
from integration._support.tls import server_context, write_self_signed_cert
from litellm import Router
from redis import Redis
PAYLOAD: Final = {"transport": "tls"}
@dataclass(frozen=True, slots=True)
class TlsRelay:
url: str
handshakes: SimpleQueue[str]
def _pipe(source: socket.socket, sink: socket.socket) -> None:
try:
while chunk := source.recv(65536):
sink.sendall(chunk)
except OSError:
pass
finally:
sink.close()
def _serve(listener: socket.socket, context: ssl.SSLContext, handshakes: SimpleQueue[str]) -> None:
while True:
try:
raw, _ = listener.accept()
except OSError:
return
try:
secured = context.wrap_socket(raw, server_side=True)
except (ssl.SSLError, OSError):
raw.close()
continue
handshakes.put(str(secured.version()))
backend = socket.create_connection((os.environ["REDIS_HOST"], int(os.environ["REDIS_PORT"])))
threading.Thread(target=_pipe, args=(secured, backend), daemon=True).start()
threading.Thread(target=_pipe, args=(backend, secured), daemon=True).start()
@contextmanager
def tls_relay(directory: Path) -> Iterator[TlsRelay]:
cert: Final = write_self_signed_cert(directory)
handshakes: Final = SimpleQueue[str]()
with socket.create_server(("127.0.0.1", 0)) as listener:
thread: Final = threading.Thread(target=_serve, args=(listener, server_context(*cert), handshakes), daemon=True)
thread.start()
port: Final = listener.getsockname()[1]
yield TlsRelay(f"rediss://127.0.0.1:{port}/0?ssl_ca_certs={cert[0]}", handshakes)
listener.close()
thread.join(timeout=5)
def _router(redis_url: str) -> Router:
return Router(
model_list=[
{
"model_name": "tls-cache",
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "synthetic-tls-key"},
}
],
redis_url=redis_url,
)
def _plain_redis() -> Redis:
return Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), decode_responses=True)
@pytest.mark.asyncio
async def test_async_router_cache_built_from_a_rediss_url_talks_tls_to_redis(tmp_path: Path) -> None:
with tls_relay(tmp_path) as relay, _plain_redis() as plain:
cache: Final = _router(relay.url).cache.redis_cache
assert cache is not None
assert await cache.ping() is True
key: Final = f"tls-async-{uuid.uuid4().hex}"
await cache.async_set_cache(key, PAYLOAD, ttl=60)
assert plain.exists(key) == 1
assert await cache.async_get_cache(key) == PAYLOAD
assert relay.handshakes.qsize() >= 1
assert relay.handshakes.get_nowait().startswith("TLS")
def test_sync_router_cache_built_from_a_rediss_url_talks_tls_to_redis(tmp_path: Path) -> None:
with tls_relay(tmp_path) as relay, _plain_redis() as plain:
cache: Final = _router(relay.url).cache.redis_cache
assert cache is not None
assert cache.sync_ping() is True
key: Final = f"tls-sync-{uuid.uuid4().hex}"
cache.set_cache(key, PAYLOAD, ttl=60)
assert plain.exists(key) == 1
assert cache.get_cache(key) == PAYLOAD
assert relay.handshakes.qsize() >= 1
assert relay.handshakes.get_nowait().startswith("TLS")

View file

@ -0,0 +1,89 @@
import json
import os
import uuid
from collections.abc import Iterator
from typing import Final
import pytest
from integration._support.wire import Reply, Request, Wire, wire_server
from litellm import Router
from litellm.caching.dual_cache import DualCache
from litellm.caching.redis_cache import RedisCache
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.proxy._types import AlertType
from litellm.types.integrations.slack_alerting import SlackAlertingCacheKeys
from redis import Redis
REPORT_SENT_KEY: Final = SlackAlertingCacheKeys.report_sent_key.value
FAILED_REQUESTS: Final = 3
API_BASE: Final = "http://daily-report-upstream.invalid/v1"
@pytest.fixture
def redis_client() -> Iterator[Redis]:
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), decode_responses=True) as client:
client.delete(REPORT_SENT_KEY)
yield client
client.delete(REPORT_SENT_KEY)
def _accept(request: Request) -> Reply:
return Reply(body=b"ok", content_type="text/plain")
def _pod(webhook: Wire) -> SlackAlerting:
redis_cache: Final = RedisCache(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]))
return SlackAlerting(
internal_usage_cache=DualCache(redis_cache=redis_cache),
alerting=["slack"],
alert_types=[AlertType.daily_reports],
alerting_args={"daily_report_frequency": 0},
default_webhook_url=webhook.url,
)
def _router(deployment_id: str) -> Router:
return Router(
model_list=[
{
"model_name": "daily-report",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": API_BASE,
"api_key": "synthetic-daily-report-key",
},
"model_info": {"id": deployment_id},
}
]
)
@pytest.mark.asyncio
async def test_the_report_timestamp_one_pod_stores_in_redis_drives_the_next_pods_daily_report(
redis_client: Redis,
) -> None:
deployment_id: Final = f"daily-report-{uuid.uuid4().hex}"
failed_key: Final = f"{deployment_id}:{SlackAlertingCacheKeys.failed_requests_key.value}"
redis_client.set(failed_key, json.dumps(FAILED_REQUESTS), ex=300)
router: Final = _router(deployment_id)
with wire_server(_accept) as webhook:
first_pod: Final = _pod(webhook)
assert await first_pod._run_scheduler_helper(llm_router=router) is False
stored: Final = redis_client.get(REPORT_SENT_KEY)
assert stored is not None
first_sent: Final = json.loads(stored)
assert isinstance(first_sent, float), stored
await first_pod.flush_queue()
assert webhook.drain() == ()
second_pod: Final = _pod(webhook)
assert await second_pod._run_scheduler_helper(llm_router=router) is True
await second_pod.flush_queue()
delivered: Final = webhook.drain()
assert len(delivered) == 1
text: Final = json.loads(delivered[0].body)["text"]
assert f"Failed Requests: `{FAILED_REQUESTS}`" in text, text
assert API_BASE in text, text
assert json.loads(redis_client.get(failed_key) or "null") == 0
assert float(json.loads(redis_client.get(REPORT_SENT_KEY) or "null")) >= first_sent
redis_client.delete(failed_key)

View file

@ -0,0 +1,97 @@
import json
import os
import uuid
from collections.abc import Iterator
from typing import Final
import pytest
from integration._support.client import eventually
from integration._support.wire import Reply, Request, Wire, wire_server
from litellm import Router
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from redis import Redis
COUNTER_TTL_SECONDS: Final = 60
CHAT_RESPONSE: Final = json.dumps(
{
"id": "chatcmpl_usage_counter_ttl",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ttl"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 7, "completion_tokens": 4, "total_tokens": 11},
}
).encode()
def _reply(request: Request) -> Reply:
assert request.target == "/v1/chat/completions", request.target
return Reply(body=CHAT_RESPONSE)
@pytest.fixture
def redis_client() -> Iterator[Redis]:
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), decode_responses=True) as client:
yield client
def _router(wire: Wire, deployment_id: str) -> Router:
return Router(
model_list=[
{
"model_name": "usage-ttl",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": f"{wire.url}/v1",
"api_key": "synthetic-usage-ttl-key",
"tpm": 1440,
},
"model_info": {"id": deployment_id},
}
],
routing_strategy="usage-based-routing-v2",
redis_host=os.environ["REDIS_HOST"],
redis_port=int(os.environ["REDIS_PORT"]),
)
def _counter_ttls(redis_client: Redis, deployment_id: str) -> dict[str, int]:
keys: Final = tuple(redis_client.scan_iter(match=f"{deployment_id}:*"))
return {key: int(redis_client.ttl(key)) for key in keys}
def _expiring(ttls: dict[str, int]) -> bool:
kinds: Final = {key.split(":")[-2] for key in ttls}
return "tpm" in kinds and all(0 < ttl <= COUNTER_TTL_SECONDS for ttl in ttls.values())
@pytest.mark.asyncio
async def test_async_usage_counters_land_in_redis_with_a_one_minute_expiry(redis_client: Redis) -> None:
deployment_id: Final = f"usage-ttl-{uuid.uuid4().hex}"
with wire_server(_reply) as wire:
router: Final = _router(wire, deployment_id)
response: Final = await router.acompletion(
model="usage-ttl", messages=[{"role": "user", "content": f"async {uuid.uuid4().hex}"}]
)
assert response.usage.total_tokens == 11
await GLOBAL_LOGGING_WORKER.flush()
ttls: Final = eventually(
lambda: _counter_ttls(redis_client, deployment_id), _expiring, seconds=15, return_last_on_timeout=True
)
assert _expiring(ttls), ttls
assert len(wire.drain()) == 1
def test_sync_usage_counters_land_in_redis_with_a_one_minute_expiry(redis_client: Redis) -> None:
deployment_id: Final = f"usage-ttl-{uuid.uuid4().hex}"
with wire_server(_reply) as wire:
router: Final = _router(wire, deployment_id)
response: Final = router.completion(
model="usage-ttl", messages=[{"role": "user", "content": f"sync {uuid.uuid4().hex}"}]
)
assert response.usage.total_tokens == 11
ttls: Final = eventually(
lambda: _counter_ttls(redis_client, deployment_id), _expiring, seconds=15, return_last_on_timeout=True
)
assert _expiring(ttls), ttls
assert len(wire.drain()) == 1

View file

@ -0,0 +1,73 @@
import uuid
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from pydantic import JsonValue
COST_PER_REQUEST: Final = 20 * 0.001 + 20 * 0.002
def _logged(key: str, requests: int) -> list[dict[str, JsonValue]]:
return eventually(
lambda: read_rows(
'SELECT model, to_char("startTime", \'YYYY-MM-DD\') AS day FROM "LiteLLM_SpendLogs" WHERE api_key=%s',
(sha256(key.encode()).hexdigest(),),
),
lambda rows: len(rows) == requests,
seconds=70,
)
def _team_entries(report: JsonValue, day: str, team_names: frozenset[str]) -> dict[str, dict[str, JsonValue]]:
assert isinstance(report, list), report
days: Final = [
object_value(row) for row in report if string_value(object_value(row)["group_by_day"]).startswith(day)
]
assert len(days) == 1, report
teams: Final = days[0]["teams"]
assert isinstance(teams, list)
return {
string_value(object_value(team)["team_name"]): object_value(team)
for team in teams
if object_value(team)["team_name"] in team_names
}
def test_default_report_groups_each_days_spend_by_team_with_per_key_breakdown(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
busy_alias: Final = f"integration-{uuid.uuid4().hex}"
quiet_alias: Final = f"integration-{uuid.uuid4().hex}"
busy: Final = scenario.team(team_alias=busy_alias, models=[model])
quiet: Final = scenario.team(team_alias=quiet_alias, models=[model])
busy_key: Final = scenario.key(team_id=busy, models=[model])
quiet_key: Final = scenario.key(team_id=quiet, models=[model])
traffic: Final = tuple(
gateway.chat(model, key=key, text=f"report {uuid.uuid4().hex}") for key in (busy_key, busy_key, quiet_key)
)
assert len({response["id"] for response in traffic}) == 3
busy_rows: Final = _logged(busy_key, 2)
_logged(quiet_key, 1)
day: Final = string_value(busy_rows[0]["day"])
stored_model: Final = busy_rows[0]["model"]
response: Final = gateway.request("GET", "/global/spend/report", params={"start_date": day, "end_date": day})
assert response.status_code == 200, response.text
entries: Final = _team_entries(response.json(), day, frozenset({busy_alias, quiet_alias}))
assert sorted(entries) == sorted((busy_alias, quiet_alias))
assert float(str(entries[busy_alias]["total_spend"])) == pytest.approx(2 * COST_PER_REQUEST)
assert float(str(entries[quiet_alias]["total_spend"])) == pytest.approx(COST_PER_REQUEST)
breakdown: Final = entries[busy_alias]["metadata"]
assert isinstance(breakdown, list)
assert [
(entry["model"], entry["api_key"], float(str(entry["spend"])), entry["total_tokens"])
for entry in map(object_value, breakdown)
] == [(stored_model, sha256(busy_key.encode()).hexdigest(), pytest.approx(2 * COST_PER_REQUEST), 80)]
filtered: Final = gateway.request(
"GET", "/global/spend/report", params={"start_date": day, "end_date": day, "team_id": quiet}
)
assert filtered.status_code == 200, filtered.text
only: Final = filtered.json()
assert len(only) == 1 and [object_value(team)["team_name"] for team in only[0]["teams"]] == [quiet_alias], only

View file

@ -0,0 +1,57 @@
import json
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, wire_server
PROMPT: Final = "a scripted sea otter"
PRICE_PER_IMAGE: Final = 0.25
def _image(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", "/images/generations")
return Reply(body=json.dumps({"created": 1700000000, "data": [{"b64_json": "aW1n"}]}).encode())
def test_identical_image_generations_each_charge_the_key(gateway: Gateway) -> None:
with wire_server(_image) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model="openai/dall-e-3",
api_base=wire.url,
api_key="synthetic-image-key",
output_cost_per_image=PRICE_PER_IMAGE,
)
key: Final = scenario.key(models=[model])
digest: Final = sha256(key.encode()).hexdigest()
body: Final = {"model": model, "prompt": PROMPT, "size": "1024x1024", "n": 1}
first: Final = gateway.request("POST", "/v1/images/generations", body, key=key)
assert first.status_code == 200, first.text
logged: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,)),
lambda rows: len(rows) == 1,
seconds=70,
)
charge: Final = float(str(logged[0]["spend"]))
assert charge == pytest.approx(PRICE_PER_IMAGE)
eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)),
lambda rows: float(str(rows[0]["spend"])) == pytest.approx(charge),
seconds=70,
)
repeat: Final = gateway.request("POST", "/v1/images/generations", body, key=key)
assert repeat.status_code == 200, repeat.text
rows: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,)),
lambda values: len(values) == 2,
seconds=70,
)
assert [float(str(row["spend"])) for row in rows] == pytest.approx([charge, charge])
eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)),
lambda values: float(str(values[0]["spend"])) == pytest.approx(2 * charge),
seconds=70,
)
assert len(wire.drain()) == 2

View file

@ -0,0 +1,84 @@
import uuid
from hashlib import sha256
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
def test_an_exhausted_key_is_refused_inference_but_can_still_read_its_own_info(gateway: Gateway) -> None:
with (
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
key: Final = scenario.key(models=[model], max_budget=0.06)
assert (
object_value(gateway.chat(model, key=key, text=f"spend {uuid.uuid4().hex}")["usage"])["total_tokens"] == 40
)
digest: Final = sha256(key.encode()).hexdigest()
eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)),
lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) >= 0.06,
seconds=70,
)
upstream.get("/__observations").raise_for_status()
denied: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"over budget {uuid.uuid4().hex}"}]},
key=key,
)
assert denied.status_code == 422, denied.text
error: Final = denied.json()["error"]
assert error["type"] == "budget_exceeded"
assert "Budget has been exceeded!" in error["message"]
assert upstream.get("/__observations").json()["requests"] == []
info: Final = gateway.request("GET", "/key/info", key=key, params={"key": key})
assert info.status_code == 200, info.text
own: Final = object_value(info.json()["info"])
assert float(str(own["spend"])) == pytest.approx(0.06)
assert own["max_budget"] == 0.06
def _bounded_chat(gateway: Gateway, model: str, key: str) -> httpx.Response:
return gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"max_tokens": 20,
"messages": [{"role": "user", "content": f"key recovery {uuid.uuid4().hex}"}],
},
key=key,
)
def test_raising_a_spent_keys_budget_restores_serving(gateway: Gateway) -> None:
with (
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
key: Final = scenario.key(models=[model], max_budget=0.06)
first: Final = _bounded_chat(gateway, model, key)
assert first.status_code == 200, first.text
eventually(
lambda: read_rows(
'SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (sha256(key.encode()).hexdigest(),)
),
lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) >= 0.06,
seconds=70,
)
eventually(lambda: _bounded_chat(gateway, model, key), lambda response: response.status_code != 200, seconds=30)
upstream.get("/__observations").raise_for_status()
denied: Final = _bounded_chat(gateway, model, key)
assert denied.status_code == 422, denied.text
assert object_value(denied.json()["error"])["type"] == "budget_exceeded"
assert upstream.get("/__observations").json()["requests"] == []
gateway.post("/key/update", {"key": key, "max_budget": 1.0})
served: Final = tuple(_bounded_chat(gateway, model, key) for _ in range(3))
assert [response.status_code for response in served] == [200, 200, 200], [response.text for response in served]
assert len(upstream.get("/__observations").json()["requests"]) == 3

View file

@ -0,0 +1,69 @@
import uuid
from dataclasses import dataclass
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, object_value
COST_PER_REQUEST: Final = 20 * 0.001 + 20 * 0.002
FIRST_BURST: Final = 6
SECOND_BURST: Final = 4
@dataclass(frozen=True, slots=True)
class Owners:
key: str
team_id: str
user_id: str
organization_id: str
def _reported(gateway: Gateway, owners: Owners) -> tuple[float, float, float, float]:
key_info: Final = object_value(gateway.get("/key/info", {"key": owners.key})["info"])
team_info: Final = object_value(gateway.get("/team/info", {"team_id": owners.team_id})["team_info"])
user_info: Final = object_value(gateway.get("/user/info", {"user_id": owners.user_id})["user_info"])
organization: Final = gateway.get("/organization/info", {"organization_id": owners.organization_id})
return (
float(str(key_info["spend"])),
float(str(team_info["spend"])),
float(str(user_info["spend"])),
float(str(organization["spend"])),
)
def _matches(observed: tuple[float, float, float, float], expected: float) -> bool:
return all(value == pytest.approx(expected, rel=1e-9) for value in observed)
def _burst(
gateway: Gateway, model: str, owners: Owners, requests: int, total_requests: int
) -> tuple[float, float, float, float]:
usage: Final = tuple(
object_value(gateway.chat(model, key=owners.key, text=f"burst {uuid.uuid4().hex}")["usage"])
for _ in range(requests)
)
assert [(entry["prompt_tokens"], entry["completion_tokens"]) for entry in usage] == [(20, 20)] * requests
return eventually(
lambda: _reported(gateway, owners),
lambda observed: _matches(observed, total_requests * COST_PER_REQUEST),
seconds=70,
return_last_on_timeout=True,
)
def test_every_burst_rolls_up_exactly_to_key_team_user_and_organization(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
organization_id: Final = scenario.organization()
team_id: Final = scenario.team(organization_id=organization_id, models=[model])
user_id: Final = scenario.user(user_role="internal_user")
owners: Final = Owners(
key=scenario.key(user_id=user_id, team_id=team_id, models=[model]),
team_id=team_id,
user_id=user_id,
organization_id=organization_id,
)
first: Final = _burst(gateway, model, owners, FIRST_BURST, FIRST_BURST)
assert first == pytest.approx((FIRST_BURST * COST_PER_REQUEST,) * 4, rel=1e-9), first
both: Final = _burst(gateway, model, owners, SECOND_BURST, FIRST_BURST + SECOND_BURST)
assert both == pytest.approx(((FIRST_BURST + SECOND_BURST) * COST_PER_REQUEST,) * 4, rel=1e-9), both

View file

@ -0,0 +1,72 @@
import uuid
from collections.abc import Iterator
from dataclasses import dataclass
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, Scenario, eventually, object_value
from integration._support.database import read_rows
TEAM_BUDGET: Final = 0.06
@dataclass(frozen=True, slots=True)
class ExhaustedTeam:
scenario: Scenario
upstream: httpx.Client
model: str
team_id: str
key: str
def _chat(gateway: Gateway, model: str, key: str) -> httpx.Response:
return gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"max_tokens": 20,
"messages": [{"role": "user", "content": f"team budget {uuid.uuid4().hex}"}],
},
key=key,
)
@pytest.fixture
def exhausted(gateway: Gateway) -> Iterator[ExhaustedTeam]:
with (
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
):
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team_id: Final = scenario.team(models=[model], max_budget=TEAM_BUDGET)
key: Final = scenario.key(team_id=team_id, models=[model], max_budget=1.0)
first: Final = _chat(gateway, model, key)
assert first.status_code == 200, first.text
eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id=%s', (team_id,)),
lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) >= TEAM_BUDGET,
seconds=70,
)
eventually(lambda: _chat(gateway, model, key), lambda response: response.status_code != 200, seconds=30)
upstream.get("/__observations").raise_for_status()
yield ExhaustedTeam(scenario, upstream, model, team_id, key)
def test_the_team_budget_blocks_a_key_whose_own_budget_has_room(gateway: Gateway, exhausted: ExhaustedTeam) -> None:
denied: Final = _chat(gateway, exhausted.model, exhausted.key)
assert denied.status_code == 422, denied.text
error: Final = object_value(denied.json()["error"])
assert error["type"] == "budget_exceeded"
assert f"Budget has been exceeded! Team={exhausted.team_id}" in str(error["message"])
assert exhausted.upstream.get("/__observations").json()["requests"] == []
def test_raising_an_exhausted_team_budget_restores_serving(gateway: Gateway, exhausted: ExhaustedTeam) -> None:
denied: Final = _chat(gateway, exhausted.model, exhausted.key)
assert denied.status_code == 422, denied.text
gateway.post("/team/update", {"team_id": exhausted.team_id, "max_budget": 1.0})
served: Final = tuple(_chat(gateway, exhausted.model, exhausted.key) for _ in range(3))
assert [response.status_code for response in served] == [200, 200, 200], [response.text for response in served]
assert len(exhausted.upstream.get("/__observations").json()["requests"]) == 3

View file

@ -83,63 +83,6 @@ async def test_completion_with_caching_bad_call():
assert sl.mock_testing_sync_success_hook == 0
@pytest.mark.asyncio
async def test_router_with_caching():
"""
- Run router with usage-based-routing-v2
- Assert success callback gets called
"""
try:
def get_openai_params():
params = {
"model": "gpt-4.1-nano",
"api_key": os.environ["OPENAI_API_KEY"],
}
return params
model_list = [
{
"model_name": "azure/gpt-4",
"litellm_params": get_openai_params(),
"tpm": 100,
},
{
"model_name": "azure/gpt-4",
"litellm_params": get_openai_params(),
"tpm": 1000,
},
]
router = litellm.Router(
model_list=model_list,
set_verbose=True,
debug_level="DEBUG",
routing_strategy="usage-based-routing-v2",
redis_host=os.environ["REDIS_HOST"],
redis_port=os.environ["REDIS_PORT"],
redis_password=os.environ["REDIS_PASSWORD"],
)
litellm.service_callback = ["prometheus_system"]
sl = ServiceLogging(mock_testing=True)
sl.prometheusServicesLogger.mock_testing = True
router.cache.redis_cache.service_logger_obj = sl
messages = [{"role": "user", "content": "Hey, how's it going?"}]
response1 = await router.acompletion(model="azure/gpt-4", messages=messages)
response1 = await router.acompletion(model="azure/gpt-4", messages=messages)
assert sl.mock_testing_async_success_hook > 0
assert sl.mock_testing_sync_failure_hook == 0
assert sl.mock_testing_async_failure_hook == 0
assert sl.prometheusServicesLogger.mock_testing_success_calls > 0
except Exception as e:
pytest.fail(f"An exception occured - {str(e)}")
@pytest.mark.asyncio
async def test_service_logger_db_monitoring():
"""

View file

@ -356,62 +356,6 @@ async def test_increment_spend_in_current_window():
assert queued_op["ttl"] == ttl
@pytest.mark.asyncio
async def test_sync_in_memory_spend_with_redis():
"""
Test _sync_in_memory_spend_with_redis helper method
Expected behavior:
- Push all provider spend increments to Redis
- Fetch all current provider spend from Redis to update in-memory cache
"""
cleanup_redis()
provider_budget_config = {
"openai": BudgetConfig(time_period="1d", budget_limit=100),
"anthropic": BudgetConfig(time_period="1d", budget_limit=200),
}
provider_budget = RouterBudgetLimiting(
dual_cache=DualCache(
redis_cache=RedisCache(
host=os.getenv("REDIS_HOST"),
port=int(os.getenv("REDIS_PORT")),
password=os.getenv("REDIS_PASSWORD"),
)
),
provider_budget_config=provider_budget_config,
)
# Allow background _init_provider_budget_in_cache tasks to complete
# before overwriting Redis values (avoids race where init overwrites with 0.0)
await asyncio.sleep(0.5)
# Set some values in Redis
spend_key_openai = "provider_spend:openai:1d"
spend_key_anthropic = "provider_spend:anthropic:1d"
await provider_budget.dual_cache.redis_cache.async_set_cache(
key=spend_key_openai, value=50.0
)
await provider_budget.dual_cache.redis_cache.async_set_cache(
key=spend_key_anthropic, value=75.0
)
# Test syncing with Redis
await provider_budget._sync_in_memory_spend_with_redis()
# Verify in-memory cache was updated
openai_spend = await provider_budget.dual_cache.in_memory_cache.async_get_cache(
spend_key_openai
)
anthropic_spend = await provider_budget.dual_cache.in_memory_cache.async_get_cache(
spend_key_anthropic
)
assert float(openai_spend) == 50.0
assert float(anthropic_spend) == 75.0
@pytest.mark.asyncio
async def test_get_current_provider_spend():
"""
@ -446,59 +390,6 @@ async def test_get_current_provider_spend():
assert spend == 50.5
@pytest.mark.flaky(retries=6, delay=2)
@pytest.mark.asyncio
async def test_get_current_provider_budget_reset_at():
"""
Test _get_current_provider_budget_reset_at helper method
Scenarios:
1. Provider with no budget config returns None
2. Provider with budget config but no TTL returns None
3. Provider with budget config and TTL returns correct ISO timestamp
"""
cleanup_redis()
provider_budget = RouterBudgetLimiting(
dual_cache=DualCache(
redis_cache=RedisCache(
host=os.getenv("REDIS_HOST"),
port=int(os.getenv("REDIS_PORT")),
password=os.getenv("REDIS_PASSWORD"),
)
),
provider_budget_config={
"openai": BudgetConfig(budget_duration="1d", max_budget=100),
"vertex_ai": BudgetConfig(budget_duration="1h", max_budget=100),
},
)
await asyncio.sleep(2)
# Test provider with no budget config
reset_at = await provider_budget._get_current_provider_budget_reset_at("anthropic")
assert reset_at is None
# Test provider with budget config but no TTL
reset_at = await provider_budget._get_current_provider_budget_reset_at("openai")
assert reset_at is not None
reset_time = datetime.fromisoformat(reset_at.replace("Z", "+00:00"))
expected_time = datetime.now(timezone.utc) + timedelta(seconds=(24 * 60 * 60))
time_difference = abs((reset_time - expected_time).total_seconds())
assert time_difference < 5
# Test provider with budget config and TTL
reset_at = await provider_budget._get_current_provider_budget_reset_at("vertex_ai")
assert reset_at is not None
# Verify the timestamp format and approximate time
reset_time = datetime.fromisoformat(reset_at.replace("Z", "+00:00"))
expected_time = datetime.now(timezone.utc) + timedelta(seconds=3600)
# Allow for small time differences (within 5 seconds)
time_difference = abs((reset_time - expected_time).total_seconds())
assert time_difference < 5
@pytest.mark.asyncio
async def test_deployment_budget_limits_e2e_test():
"""

View file

@ -18,61 +18,6 @@ from litellm.caching import RedisCache, RedisClusterCache
## 2. 2 models - openai, azure - 2 diff model groups, 1 caching group
@pytest.mark.asyncio
async def test_router_async_caching_with_ssl_url():
"""
Tests when a redis url is passed to the router, if caching is correctly setup
"""
try:
router = Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
},
"tpm": 100000,
"rpm": 10000,
},
],
redis_url=os.getenv("REDIS_SSL_URL"),
)
response = await router.cache.redis_cache.ping()
print(f"response: {response}")
assert response == True
except Exception as e:
pytest.fail(f"An exception occurred - {str(e)}")
def test_router_sync_caching_with_ssl_url():
"""
Tests when a redis url is passed to the router, if caching is correctly setup
"""
try:
router = Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
},
"tpm": 100000,
"rpm": 10000,
},
],
redis_url=os.getenv("REDIS_SSL_URL"),
)
response = router.cache.redis_cache.sync_ping()
print(f"response: {response}")
assert response == True
except Exception as e:
pytest.fail(f"An exception occurred - {str(e)}")
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=1)
async def test_acompletion_caching_on_router():

View file

@ -505,159 +505,6 @@ async def test_router_completion_streaming():
"""
@pytest.mark.asyncio
async def test_router_caching_ttl():
"""
Confirm caching ttl's work as expected.
Relevant issue: https://github.com/BerriAI/litellm/issues/5609
"""
messages = [
{"role": "user", "content": "Hello, can you generate a 500 words poem?"}
]
model = "azure-model"
model_list = [
{
"model_name": "azure-model",
"litellm_params": {
"model": "azure/gpt-turbo",
"api_key": "os.environ/AZURE_FRANCE_API_KEY",
"api_base": "https://openai-france-1234.openai.azure.com",
"tpm": 1440,
"mock_response": "Hello world",
},
"model_info": {"id": 1},
}
]
router = Router(
model_list=model_list,
routing_strategy="usage-based-routing-v2",
set_verbose=False,
redis_host=os.getenv("REDIS_HOST"),
redis_password=os.getenv("REDIS_PASSWORD"),
redis_port=os.getenv("REDIS_PORT"),
)
assert router.cache.redis_cache is not None
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
increment_cache_kwargs = {}
with patch.object(
router.cache,
"async_increment_cache_pipeline",
new=AsyncMock(),
) as mock_client:
await router.acompletion(model=model, messages=messages)
# Async success callbacks are dispatched to GLOBAL_LOGGING_WORKER's
# background queue; drain it before asserting the mock was invoked.
await GLOBAL_LOGGING_WORKER.flush()
# mock_client.assert_called_once()
print(f"mock_client.call_args.kwargs: {mock_client.call_args.kwargs}")
print(f"mock_client.call_args.args: {mock_client.call_args.args}")
# Get the increment_list from the first positional argument or the keyword argument
increment_list = mock_client.call_args.kwargs.get(
"increment_list",
mock_client.call_args.args[0] if mock_client.call_args.args else None,
)
assert increment_list is not None
assert len(increment_list) > 0
# Check that TTL is set to 60 for all operations
for operation in increment_list:
assert operation["ttl"] == 60
# Get the first operation for testing the redis increment
first_operation = increment_list[0]
increment_cache_kwargs = {
"key": first_operation["key"],
"value": first_operation["increment_value"],
"ttl": first_operation["ttl"],
}
## call redis async increment and check if ttl correctly set
await router.cache.redis_cache.async_increment(**increment_cache_kwargs)
_redis_client = router.cache.redis_cache.init_async_client()
async with _redis_client as redis_client:
current_ttl = await redis_client.ttl(increment_cache_kwargs["key"])
assert current_ttl >= 0
print(f"current_ttl: {current_ttl}")
def test_router_caching_ttl_sync():
"""
Confirm caching ttl's work as expected.
Relevant issue: https://github.com/BerriAI/litellm/issues/5609
"""
messages = [
{"role": "user", "content": "Hello, can you generate a 500 words poem?"}
]
model = "azure-model"
model_list = [
{
"model_name": "azure-model",
"litellm_params": {
"model": "azure/gpt-turbo",
"api_key": "os.environ/AZURE_FRANCE_API_KEY",
"api_base": "https://openai-france-1234.openai.azure.com",
"tpm": 1440,
"mock_response": "Hello world",
},
"model_info": {"id": 1},
}
]
router = Router(
model_list=model_list,
routing_strategy="usage-based-routing-v2",
set_verbose=False,
redis_host=os.getenv("REDIS_HOST"),
redis_password=os.getenv("REDIS_PASSWORD"),
redis_port=os.getenv("REDIS_PORT"),
)
assert router.cache.redis_cache is not None
increment_cache_kwargs = {}
with patch.object(
router.cache.redis_cache,
"increment_cache",
new=MagicMock(),
) as mock_client:
router.completion(model=model, messages=messages)
print(mock_client.call_args_list)
mock_client.assert_called()
print(f"mock_client.call_args.kwargs: {mock_client.call_args.kwargs}")
print(f"mock_client.call_args.args: {mock_client.call_args.args}")
increment_cache_kwargs = {
"key": mock_client.call_args.args[0],
"value": mock_client.call_args.args[1],
"ttl": mock_client.call_args.kwargs["ttl"],
}
assert mock_client.call_args.kwargs["ttl"] == 60
## call redis async increment and check if ttl correctly set
router.cache.redis_cache.increment_cache(**increment_cache_kwargs)
_redis_client = router.cache.redis_cache.redis_client
current_ttl = _redis_client.ttl(increment_cache_kwargs["key"])
assert current_ttl >= 0
print(f"current_ttl: {current_ttl}")
def test_return_potential_deployments():
"""
Assert deployment at limit is filtered out

View file

@ -128,8 +128,6 @@ def test_init():
print("passed testing slack alerting init")
@pytest.fixture
def slack_alerting():
return SlackAlerting(
@ -326,52 +324,6 @@ async def test_daily_reports_completion(slack_alerting):
mock_send_alert.assert_awaited()
@pytest.mark.asyncio
async def test_daily_reports_redis_cache_scheduler():
redis_cache = RedisCache()
slack_alerting = SlackAlerting(
internal_usage_cache=DualCache(redis_cache=redis_cache)
)
# we need this to be 0 so it actualy sends the report
slack_alerting.alerting_args.daily_report_frequency = 0
router = litellm.Router(
model_list=[
{
"model_name": "gpt-5.5",
"litellm_params": {
"model": "gpt-5-mini",
},
}
]
)
with (
patch.object(slack_alerting, "send_alert", new=AsyncMock()) as mock_send_alert,
patch.object(
redis_cache, "async_set_cache", new=AsyncMock()
) as mock_redis_set_cache,
):
# initial call - expect empty
await slack_alerting._run_scheduler_helper(llm_router=router)
try:
json.dumps(mock_redis_set_cache.call_args[0][1])
except Exception as e:
pytest.fail(
"Cache value can't be json dumped - {}".format(
mock_redis_set_cache.call_args[0][1]
)
)
mock_redis_set_cache.assert_awaited_once()
# second call - expect empty
await slack_alerting._run_scheduler_helper(llm_router=router)
@pytest.mark.asyncio
@pytest.mark.skip(reason="Local test. Test if slack alerts are sent.")
async def test_send_llm_exception_to_slack():

View file

@ -1,241 +0,0 @@
"""
E2E tests for OpenAI Responses API WebSocket mode through the LiteLLM proxy.
Connects to ws://0.0.0.0:4000/v1/responses, sends response.create events,
and validates the streamed response events.
Requires:
- Proxy running: python -m litellm.proxy.proxy_cli --config <config> --port 4000
- Model configured in proxy (e.g. gpt-5-mini)
See: https://developers.openai.com/api/docs/guides/websocket-mode/
"""
import asyncio
import json
import os
import httpx
import pytest
# ── Configuration ─────────────────────────────────────────────────────────────
PROXY_BASE_URL = os.environ.get("LITELLM_PROXY_BASE_URL", "ws://0.0.0.0:4000")
PROXY_MASTER_KEY = os.environ.get("LITELLM_PROXY_KEY", "sk-1234")
PROXY_MODEL = os.environ.get("LITELLM_PROXY_RESPONSES_MODEL", "gpt-5-mini")
# ──────────────────────────────────────────────────────────────────────────────
def _generate_key() -> str:
"""Generate a key for testing via proxy key/generate endpoint."""
url = "http://0.0.0.0:4000/key/generate"
headers = {
"Authorization": f"Bearer {PROXY_MASTER_KEY}",
"Content-Type": "application/json",
}
response = httpx.post(url, headers=headers, json={}, timeout=10)
if response.status_code != 200:
raise Exception(
f"Key generation failed with status: {response.status_code}. "
"Is the proxy running?"
)
return response.json()["key"]
def _assert_basic_response(events: list[dict], label: str = "") -> None:
"""Assert that events contain response.created, response.completed, and usage."""
prefix = f"[{label}] " if label else ""
types = [e.get("type") for e in events]
assert len(events) > 0, f"{prefix}no events received"
assert (
"response.created" in types
), f"{prefix}missing response.created, got: {types}"
assert (
"response.completed" in types
), f"{prefix}missing response.completed, got: {types}"
completed = next(e for e in events if e.get("type") == "response.completed")
resp = completed.get("response", {})
assert (
resp.get("status") == "completed"
), f"{prefix}status != completed: {resp.get('status')}"
usage = resp.get("usage", {})
assert usage.get("input_tokens", 0) > 0, f"{prefix}input_tokens=0"
assert usage.get("output_tokens", 0) > 0, f"{prefix}output_tokens=0"
streaming_types = {
"response.output_item.added",
"response.content_part.added",
"response.output_text.delta",
"response.output_item.done",
}
found = streaming_types & set(types)
assert found, f"{prefix}no streaming delta events found, got: {types}"
@pytest.mark.asyncio
async def test_responses_websocket_proxy_basic():
"""
Sends a simple response.create event to the proxy WebSocket endpoint
and validates response.created, response.completed, and streaming events.
"""
try:
import websockets
except ImportError:
pytest.skip("websockets not installed")
try:
key = _generate_key()
except Exception as e:
pytest.skip(
f"Proxy not available or key generation failed: {e}. "
"Start proxy: python -m litellm.proxy.proxy_cli --config <config> --port 4000"
)
url = f"{PROXY_BASE_URL}/v1/responses?model={PROXY_MODEL}"
headers = {"Authorization": f"Bearer {key}"}
events: list[dict] = []
try:
async with websockets.connect(
url, additional_headers=headers, open_timeout=5
) as ws:
payload = {
"type": "response.create",
"model": PROXY_MODEL,
"store": False,
"input": [
{
"type": "message",
"role": "user",
"content": [
{"type": "input_text", "text": "Say hello in one word."}
],
}
],
"tools": [],
}
await ws.send(json.dumps(payload))
for _ in range(50):
msg = await asyncio.wait_for(ws.recv(), timeout=15)
event = json.loads(msg)
events.append(event)
if event.get("type") in (
"response.completed",
"response.failed",
"error",
):
break
except Exception as e:
pytest.fail(
f"WebSocket connection failed: {e}. "
"Ensure proxy is running and model is configured."
)
_assert_basic_response(events, "proxy-basic")
@pytest.mark.asyncio
async def test_responses_websocket_proxy_multi_turn():
"""
Sends two sequential response.create events with previous_response_id
to validate multi-turn conversation over a single WebSocket.
"""
try:
import websockets
except ImportError:
pytest.skip("websockets not installed")
try:
key = _generate_key()
except Exception as e:
pytest.skip(
f"Proxy not available or key generation failed: {e}. "
"Start proxy: python -m litellm.proxy.proxy_cli --config <config> --port 4000"
)
url = f"{PROXY_BASE_URL}/v1/responses?model={PROXY_MODEL}"
headers = {"Authorization": f"Bearer {key}"}
all_events: list[dict] = []
completed: list[dict] = []
first_id = None
try:
async with websockets.connect(
url, additional_headers=headers, open_timeout=5
) as ws:
# Turn 1
await ws.send(
json.dumps(
{
"type": "response.create",
"model": PROXY_MODEL,
"store": True,
"input": [
{
"type": "message",
"role": "user",
"content": [
{
"type": "input_text",
"text": "Remember the number 7. Just say OK.",
}
],
}
],
}
)
)
for _ in range(50):
msg = await asyncio.wait_for(ws.recv(), timeout=15)
event = json.loads(msg)
all_events.append(event)
if event.get("type") == "response.completed":
completed.append(event)
first_id = event.get("response", {}).get("id")
break
if event.get("type") in ("response.failed", "error"):
break
assert first_id, "Turn 1 never completed"
# Turn 2
await ws.send(
json.dumps(
{
"type": "response.create",
"model": PROXY_MODEL,
"store": True,
"previous_response_id": first_id,
"input": [
{
"type": "message",
"role": "user",
"content": [
{
"type": "input_text",
"text": "What number did I tell you to remember?",
}
],
}
],
}
)
)
for _ in range(50):
msg = await asyncio.wait_for(ws.recv(), timeout=15)
event = json.loads(msg)
all_events.append(event)
if event.get("type") == "response.completed":
completed.append(event)
break
if event.get("type") in ("response.failed", "error"):
break
except Exception as e:
pytest.fail(
f"WebSocket multi-turn failed: {e}. "
"Ensure proxy is running and model is configured."
)
assert (
len(completed) >= 2
), f"Expected 2 response.completed events, got {len(completed)}"
assert completed[1].get("response", {}).get("status") == "completed"

View file

@ -83,18 +83,6 @@ async def chat_completion(session, key: str, model: str):
return response
async def update_key_budget(session, key: str, max_budget: float):
"""Helper function to update a key's max budget"""
url = "http://0.0.0.0:4000/key/update"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"key": key,
"max_budget": max_budget,
}
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
@pytest.mark.asyncio
async def test_chat_completion_low_budget():
"""
@ -174,51 +162,6 @@ async def test_chat_completion_high_budget():
), "Should make at least one successful call before budget exceeded"
@pytest.mark.asyncio
async def test_chat_completion_budget_update():
"""
Test that requests continue working after updating a key's budget:
1. Create key with low budget
2. Make calls until budget exceeded
3. Update key with higher budget
4. Verify calls work again
"""
async with aiohttp.ClientSession() as session:
# Create key with very low budget
key_gen = await generate_key(session=session, max_budget=0.0000000005)
key = key_gen["key"]
# Make calls until budget exceeded
calls_made = await make_calls_until_budget_exceeded(
session=session,
key=key,
call_function=chat_completion,
model="fake-openai-endpoint",
)
assert (
calls_made > 0
), "Should make at least one successful call before budget exceeded"
# Update key with higher budget
await update_key_budget(session, key, max_budget=0.001)
# Verify calls work again
for _ in range(3):
try:
response = await chat_completion(
session=session, key=key, model="fake-openai-endpoint"
)
print("response: ", response)
assert (
response is not None
), "Should get valid response after budget update"
except Exception as e:
pytest.fail(
f"Request should succeed after budget update but got error: {e}"
)
@pytest.mark.parametrize(
"field",
[
@ -610,112 +553,4 @@ async def test_team_budget_enforcement_cli_sso_token():
), "Should make at least one successful call before team budget exceeded"
@pytest.mark.asyncio
async def test_team_and_key_budget_enforcement():
"""
Test budget enforcement when both team and key have budgets:
1. Create team with low budget
2. Create key with higher budget
3. Verify team budget is enforced first
"""
async with aiohttp.ClientSession() as session:
# Create team with very low budget
team_response = await create_team(session=session, max_budget=0.0000000005)
team_id = team_response["team_id"]
# Create key with higher budget
key_gen = await generate_team_key(
session=session,
team_id=team_id,
max_budget=0.001, # Higher than team budget
)
key = key_gen["key"]
# Make calls until budget exceeded
calls_made = await make_calls_until_budget_exceeded(
session=session,
key=key,
call_function=chat_completion,
model="fake-openai-endpoint",
)
assert (
calls_made > 0
), "Should make at least one successful call before team budget exceeded"
# Verify it was the team budget that was exceeded
try:
await chat_completion(
session=session, key=key, model="fake-openai-endpoint"
)
except Exception as e:
error_dict = e.body
assert (
"Budget has been exceeded! Team=" in error_dict["message"]
), "Error should mention team budget being exceeded"
assert team_id in error_dict["message"], "Error should mention team id"
async def update_team_budget(session, team_id: str, max_budget: float):
"""Helper function to update a team's max budget"""
url = "http://0.0.0.0:4000/team/update"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"team_id": team_id,
"max_budget": max_budget,
}
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
@pytest.mark.asyncio
async def test_team_budget_update():
"""
Test that requests continue working after updating a team's budget:
1. Create team with low budget
2. Create key for that team
3. Make calls until team budget exceeded
4. Update team with higher budget
5. Verify calls work again
"""
async with aiohttp.ClientSession() as session:
# Create team with very low budget
team_response = await create_team(session=session, max_budget=0.0000000005)
team_id = team_response["team_id"]
# Create key for team (no specific budget)
key_gen = await generate_team_key(session=session, team_id=team_id)
key = key_gen["key"]
# Make calls until budget exceeded
calls_made = await make_calls_until_budget_exceeded(
session=session,
key=key,
call_function=chat_completion,
model="fake-openai-endpoint",
)
assert (
calls_made > 0
), "Should make at least one successful call before team budget exceeded"
# Update team with higher budget
await update_team_budget(session, team_id, max_budget=0.001)
# Verify calls work again
for _ in range(3):
try:
response = await chat_completion(
session=session, key=key, model="fake-openai-endpoint"
)
print("response: ", response)
assert (
response is not None
), "Should get valid response after budget update"
except Exception as e:
pytest.fail(
f"Request should succeed after team budget update but got error: {e}"
)
# Verify it was the team budget that was exceeded

View file

@ -1,135 +0,0 @@
# What this tests ?
## Tests /chat/completions by generating a key and then making a chat completions request
import pytest
import asyncio
import aiohttp, openai
from openai import OpenAI, AsyncOpenAI
from typing import Optional, List, Union
from litellm._uuid import uuid
async def generate_key(
session,
models=[
"gpt-5.5",
"text-embedding-3-small",
"gpt-image-1",
"fake-openai-endpoint",
"mistral-embed",
],
):
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"models": models,
"duration": None,
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
async def chat_completion(session, key, model: Union[str, List] = "gpt-5.5"):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": model,
"messages": [
{"role": "user", "content": f"Hello! {str(uuid.uuid4())}"},
],
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
async def get_otel_spans(session, key):
url = "http://0.0.0.0:4000/otel-spans"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
async with session.get(url, headers=headers) as response:
status = response.status
response_text = await response.text()
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
@pytest.mark.asyncio
async def test_chat_completion_check_otel_spans():
"""
- Create key
Make chat completion call
- Create user
make chat completion call
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session)
key = key_gen["key"]
await chat_completion(session=session, key=key, model="fake-openai-endpoint")
await asyncio.sleep(3)
# /otel-spans requires proxy admin; use the master key.
otel_spans = await get_otel_spans(session=session, key="sk-1234")
print("otel_spans: ", otel_spans)
all_otel_spans = otel_spans["otel_spans"]
spans_grouped_by_parent = otel_spans["spans_grouped_by_parent"]
print("\n spans grouped by parent: ", spans_grouped_by_parent)
# The GET /otel-spans request itself produces auth spans that beat
# the chat-completion spans on start_time, so `most_recent_parent`
# points at the wrong trace. Pick the chat-completion trace by
# content: it's the one carrying the full set of expected markers.
chat_completion_markers = {
"postgres",
"redis",
"raw_gen_ai_request",
"batch_write_to_db",
}
parent_trace_spans = next(
spans
for spans in spans_grouped_by_parent.values()
if chat_completion_markers.issubset(spans)
)
print("Parent trace spans: ", parent_trace_spans)
# either 5 or 6 traces depending on how many redis calls were made
assert len(parent_trace_spans) >= 5
# 'postgres', 'redis', 'raw_gen_ai_request', 'litellm_request', 'Received Proxy Server Request' in the span
assert "postgres" in parent_trace_spans
assert "redis" in parent_trace_spans
assert "raw_gen_ai_request" in parent_trace_spans
assert "batch_write_to_db" in parent_trace_spans

View file

@ -1,490 +0,0 @@
"""
1. Default permissions for members in a team - allowed to call /key/info and /key/health
- Create a team, create a member in a team (role = "user")
Invalid Permissions:
- User tries creating a key with team_id = team_id -> expect to fail. Invalid Permissions
- User tries editing a key with team_id = team_id -> expect to fail. Invalid Permissions
- User tries deleting a key with team_id = team_id -> expect to fail. Invalid Permissions
- User tries regenerating a key with team_id = team_id -> expect to fail. Invalid Permissions
Valid Permissions:
- User tries calling /key/info with team_id, expect to get valid response
2. Permissions - members allowd to edit, delete keys but not allowed to create keys
- Create a team with member_permissions = ["/key/update", "/key/delete", "/key/info"]
- Create a member in the team with role = "user"
Valid Permissions:
- User tries editing a key with team_id = team_id -> expect to pass. Valid Permissions
- Note: Delete/regenerate require key ownership or team admin status, not just team member permissions
- User tries deleting a key with team_id = team_id -> expect to fail (403) unless user owns the key or is team admin
- User tries regenerating a key with team_id = team_id -> expect to fail (403) unless user owns the key or is team admin
Invalid Permissions:
- User tries creating a key with team_id = team_id -> expect to fail. Invalid Permissions
- User tries calling /key/info with team_id, expect to get valid response
3. Permissions - members allowed to create keys but not allowed to edit, delete keys
- Create a team with member_permissions = ["/key/generate"]
- Create a member in the team with role = "user"
Valid Permissions:
- User tries creating a key with team_id = team_id -> expect to pass. Valid Permissions
Invalid Permissions:
- User tries editing a key with team_id = team_id -> expect to fail. Invalid Permissions
- User tries deleting a key with team_id = team_id -> expect to fail. Invalid Permissions
- User tries regenerating a key with team_id = team_id -> expect to fail. Invalid Permissions
"""
import pytest
import asyncio
import aiohttp, openai
from litellm._uuid import uuid
import json
from litellm.proxy._types import ProxyErrorTypes
from typing import Optional
LITELLM_MASTER_KEY = "sk-1234"
async def create_team(session, key, member_permissions=None):
url = "http://0.0.0.0:4000/team/new"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {"team_member_permissions": member_permissions}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
raise Exception(response_text)
return await response.json()
async def create_user(session, key, user_id, team_id=None):
url = "http://0.0.0.0:4000/user/new"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {"user_id": user_id}
if team_id:
data["team_id"] = team_id
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
raise Exception(response_text)
return await response.json()
async def add_team_member(session, key, team_id, user_id, role="user"):
url = "http://0.0.0.0:4000/team/member_add"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {"team_id": team_id, "member": {"role": role, "user_id": user_id}}
print("Adding team member with data: ", data)
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
raise Exception(response_text)
return await response.json()
async def generate_key(session, key, team_id=None, user_id=None):
url = "http://0.0.0.0:4000/key/generate"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {}
if team_id:
data["team_id"] = team_id
if user_id:
data["user_id"] = user_id
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
return {"status": status, "error": response_text}
return await response.json()
async def key_info(session, key, key_id):
url = f"http://0.0.0.0:4000/key/info?key={key_id}"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
async with session.get(url, headers=headers) as response:
status = response.status
response_text = await response.text()
if status != 200:
return {"status": status, "error": response_text}
return await response.json()
async def update_key(
session: aiohttp.ClientSession,
key: str,
key_id: str,
team_id: Optional[str] = None,
):
"""
Update a key
Args:
key: key to use for authentication
key_id: key to update
"""
url = "http://0.0.0.0:4000/key/update"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {"key": key_id, "metadata": {"updated": True}}
if team_id:
data["team_id"] = team_id
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
return {"status": status, "error": response_text}
return await response.json()
async def delete_key(session, key, key_id):
url = "http://0.0.0.0:4000/key/delete"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {"keys": [key_id]}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
return {"status": status, "error": response_text}
return await response.json()
async def regenerate_key(session, key, key_id, team_id=None):
url = "http://0.0.0.0:4000/key/regenerate"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {"key": key_id}
if team_id:
data["team_id"] = team_id
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
return {"status": status, "error": response_text}
return await response.json()
@pytest.mark.asyncio()
async def test_default_member_permissions():
"""
Test default permissions for members in a team - allowed to call /key/info and /key/health
"""
async with aiohttp.ClientSession() as session:
master_key = LITELLM_MASTER_KEY
# Create a team
team_data = await create_team(session=session, key=master_key)
team_id = team_data["team_id"]
# create a team key
team_key_data = await generate_key(
session=session, key=master_key, team_id=team_id
)
team_key = team_key_data["key"]
# create a user
user_data = await create_user(
session=session,
key=master_key,
user_id=f"user_{uuid.uuid4().hex[:8]}",
team_id=team_id,
)
user_id = user_data["user_id"]
# Create a user key
print("New user data: ", user_data)
# Create a user key
user_key_data = await generate_key(
session=session, key=master_key, user_id=user_id
)
print("new user key: ", user_key_data)
user_key = user_key_data["key"]
# Test invalid permissions
# User tries creating a key with team_id
print(
"Regular team member trying to create a key with team_id. Expecting error."
)
create_result = await generate_key(
session=session, key=user_key, team_id=team_id
)
print("result: ", create_result)
assert (
"status" in create_result and create_result["status"] == 401
), "User should not be able to create keys for team"
error_data = json.loads(create_result["error"])
print("error response =", json.dumps(error_data, indent=4))
assert (
error_data["error"]["type"]
== ProxyErrorTypes.team_member_permission_error.value
), "Error should be a team member permission error"
# User tries editing a key with team_id
print("Regular team member trying to edit a key with team_id. Expecting error.")
update_result = await update_key(
session=session, key=user_key, key_id=team_key, team_id="ATTACKER_TEAM_ID"
)
assert (
"status" in update_result and update_result["status"] == 401
), "User should not be able to update keys for team"
error_data = json.loads(update_result["error"])
print("error response =", json.dumps(error_data, indent=4))
assert (
error_data["error"]["type"]
== ProxyErrorTypes.team_member_permission_error.value
), "Error should be a team member permission error"
# User tries deleting a key with team_id
print(
"Regular team member trying to delete a key with team_id. Expecting error."
)
delete_result = await delete_key(
session=session,
key=user_key,
key_id=team_key,
)
assert (
"status" in delete_result and delete_result["status"] == 403
), "User should not be able to delete keys for team"
error_data = json.loads(delete_result["error"])
print("error response =", json.dumps(error_data, indent=4))
# Delete endpoint now returns 403 with authorization error, not team_member_permission_error
assert "error" in error_data, "Error should contain error field"
# User tries regenerating a key with team_id
print(
"Regular team member trying to regenerate a key with team_id. Expecting error."
)
regenerate_result = await regenerate_key(
session=session,
key=user_key,
key_id=team_key,
)
assert (
"status" in regenerate_result and regenerate_result["status"] == 401
), "User should not be able to regenerate keys for team"
error_data = json.loads(regenerate_result["error"])
print("error response =", json.dumps(error_data, indent=4))
# Regenerate endpoint now returns 403 with authorization error, not team_member_permission_error
assert "error" in error_data, "Error should contain error field"
# Test valid permissions
# User tries calling /key/info with team_id
print(
"Regular team member trying to get key info with team_id. Expecting success."
)
info_result = await key_info(
session=session,
key=user_key,
key_id=team_key,
)
print("info result =", info_result)
assert "status" not in info_result, "Admin should be able to get key info"
@pytest.mark.asyncio()
async def test_edit_delete_permissions():
"""
Test permissions - members allowed to edit, delete keys but not allowed to create keys
"""
async with aiohttp.ClientSession() as session:
master_key = LITELLM_MASTER_KEY
# Create a team with specific member permissions
team_data = await create_team(
session=session,
key=master_key,
member_permissions=["/key/update", "/key/delete", "/key/info"],
)
team_id = team_data["team_id"]
# create a user in team=team_id
user_data = await create_user(
session=session,
key=master_key,
user_id=f"user_{uuid.uuid4().hex[:8]}",
team_id=team_id,
)
user_id = user_data["user_id"]
# Generate an admin key for the team
admin_key_data = await generate_key(session, master_key, team_id)
key_id = admin_key_data["key"]
# Create a user key
user_key_data = await generate_key(
session=session, key=master_key, user_id=user_id
)
user_key = user_key_data["key"]
# Test valid permissions
# User tries editing a key with team_id
update_result = await update_key(
session=session, key=user_key, key_id=key_id, team_id=team_id
)
assert (
"status" not in update_result
), "User should be able to update keys for team"
# User tries deleting a key with team_id
# Note: Even with /key/delete permission, users can only delete keys they own or if they're team admin
# The delete endpoint checks ownership/team admin status, not just team member permissions
delete_result = await delete_key(session=session, key=user_key, key_id=key_id)
assert (
"status" in delete_result and delete_result["status"] == 403
), "User should not be able to delete keys they don't own (even with /key/delete permission, ownership is required)"
# Test invalid permissions
# User tries creating a key with team_id
create_result = await generate_key(
session=session, key=user_key, team_id=team_id
)
assert (
"status" in create_result and create_result["status"] != 200
), "User should not be able to create keys for team"
# User tries regenerating a key with team_id
# Note: Even with /key/regenerate permission, users can only regenerate keys they own or if they're team admin
regenerate_result = await regenerate_key(
session=session, key=user_key, key_id=key_id, team_id=team_id
)
assert (
"status" in regenerate_result and regenerate_result["status"] == 401
), "User should not be able to regenerate keys they don't own (even with /key/regenerate permission, ownership is required)"
@pytest.mark.asyncio()
async def test_create_permissions():
"""
Test permissions - members allowed to create keys but not allowed to edit, delete keys
"""
async with aiohttp.ClientSession() as session:
master_key = LITELLM_MASTER_KEY
# Create a team with specific member permissions
team_data = await create_team(
session=session, key=master_key, member_permissions=["/key/generate"]
)
team_id = team_data["team_id"]
# Create a user in the team
user_id = f"user_{uuid.uuid4().hex[:8]}"
await add_team_member(
session=session,
key=master_key,
team_id=team_id,
user_id=user_id,
role="user",
)
# Generate an admin key for the team
admin_key_data = await generate_key(
session=session, key=master_key, team_id=team_id
)
admin_key = admin_key_data["key"]
key_id = admin_key_data["key"]
# Create a user key
user_key_data = await generate_key(
session=session, key=master_key, user_id=user_id
)
user_key = user_key_data["key"]
# Test valid permissions
# User tries creating a key with team_id
create_result = await generate_key(
session=session, key=user_key, team_id=team_id
)
print("success, user created key for team=", create_result)
assert "key" in create_result, "User should be able to create keys for team"
assert (
create_result["team_id"] == team_id
), "User should be able to create keys for team"
assert (
"status" not in create_result
), "User should be able to create keys for team"
# Test invalid permissions
# User tries editing a key with team_id
update_result = await update_key(
session=session, key=user_key, key_id=key_id, team_id=team_id
)
assert (
"status" in update_result and update_result["status"] != 200
), "User should not be able to update keys for team"
# User tries deleting a key with team_id
delete_result = await delete_key(session=session, key=user_key, key_id=key_id)
assert (
"status" in delete_result and delete_result["status"] == 403
), "User should not be able to delete keys for team"
# User tries regenerating a key with team_id
# User doesn't have /key/regenerate permission, so should get 401 (team member permission error)
regenerate_result = await regenerate_key(
session=session, key=user_key, key_id=key_id, team_id=team_id
)
assert (
"status" in regenerate_result and regenerate_result["status"] == 401
), "User should not be able to regenerate keys for team (no /key/regenerate permission)"
error_data = json.loads(regenerate_result["error"])
assert (
error_data["error"]["type"]
== ProxyErrorTypes.team_member_permission_error.value
), "Error should be a team member permission error"

View file

@ -36,45 +36,6 @@ async def chat_completion(
return await response.json(), response.headers
async def create_team_with_tags(session, key, tags: List[str]):
url = "http://0.0.0.0:4000/team/new"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"tags": tags,
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
raise Exception(response_text)
return await response.json()
async def create_key_with_team(session, key, team_id: str):
url = f"http://0.0.0.0:4000/key/generate"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"team_id": team_id,
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
if status != 200:
raise Exception(response_text)
return await response.json()
async def model_info_get_call(session, key, model_id: str):
# make get call pass "litellm_model_id" in query params
url = f"http://0.0.0.0:4000/model/info?litellm_model_id={model_id}"
@ -92,45 +53,6 @@ async def model_info_get_call(session, key, model_id: str):
return await response.json()
@pytest.mark.asyncio()
async def test_team_tag_routing():
async with aiohttp.ClientSession() as session:
key = LITELLM_MASTER_KEY
team_a_data = await create_team_with_tags(session, key, ["teamA"])
print("team_a_data=", team_a_data)
team_a_id = team_a_data["team_id"]
team_b_data = await create_team_with_tags(session, key, ["teamB"])
print("team_b_data=", team_b_data)
team_b_id = team_b_data["team_id"]
key_with_team_a = await create_key_with_team(session, key, team_a_id)
print("key_with_team_a=", key_with_team_a)
_key_with_team_a = key_with_team_a["key"]
for _ in range(5):
response_a, headers = await chat_completion(
session=session, key=_key_with_team_a
)
headers = dict(headers)
print(response_a)
print(headers)
assert (
headers["x-litellm-model-id"] == "team-a-model"
), "Model ID should be teamA"
key_with_team_b = await create_key_with_team(session, key, team_b_id)
_key_with_team_b = key_with_team_b["key"]
for _ in range(5):
response_b, headers = await chat_completion(session, _key_with_team_b)
headers = dict(headers)
print(response_b)
print(headers)
assert (
headers["x-litellm-model-id"] == "team-b-model"
), "Model ID should be teamB"
@pytest.mark.asyncio()
async def test_chat_completion_with_no_tags():
async with aiohttp.ClientSession() as session:

View file

@ -1,395 +0,0 @@
import pytest
import asyncio
import aiohttp
import time
import litellm
from litellm._uuid import uuid
"""
Tests to run
Basic Tests:
1. Basic Spend Accuracy Test:
- Make N requests, compute expected total spend locally from each response's usage
- Poll until batch writer has flushed spend to the DB
- Expect spend for Key, Team, User, Org (/info endpoints) to equal the computed total
2. Long term spend accuracy test (with 2 bursts of requests)
- Burst 1: compute expected from responses, verify
- Burst 2: compute expected from responses, verify total = burst1 + burst2
Additional Test Scenarios:
3. Concurrent Request Accuracy Test:
- Make 20 concurrent requests
- Check for race conditions in spend tracking
4. Error Case Test:
- Make 10 successful requests
- Make 5 failed requests
- Verify spend is only counted for successful requests
5. Mixed Request Type Test:
- Make different types of requests with varying costs
- Verify accurate total spend calculation
"""
# Upstream model the proxy is configured with (spend_tracking_config.yaml).
# The proxy computes spend using this model's pricing; the local ground-truth
# calculation uses the same pricing table via litellm.cost_per_token.
UPSTREAM_MODEL = "gpt-5-mini"
# Batch writer flush cadence in CI is ~2-7s (PROXY_BATCH_WRITE_AT=2 + up to 5s jitter).
# Poll every 2s for 60s — plenty of headroom for multiple ticks to land.
POLL_INTERVAL_SECONDS = 2
POLL_TIMEOUT_SECONDS = 60
TOLERANCE = 1e-10
def _make_test_session() -> aiohttp.ClientSession:
"""
Session tuned for CI reliability:
- force_close: avoid aiohttp reusing a TCP connection that the proxy/kernel
silently closed during the long idle window between setup POSTs and the
later poll loop (observed failure mode: ConnectionTimeoutError on the
first /key/info call after 20 chat completions).
- explicit connect timeout: surface a blocked proxy event loop quickly
instead of hanging on aiohttp's 5-minute default total timeout.
"""
return aiohttp.ClientSession(
connector=aiohttp.TCPConnector(force_close=True),
timeout=aiohttp.ClientTimeout(total=30, connect=10),
)
async def create_organization(session, organization_alias: str):
"""Helper function to create a new organization"""
url = "http://0.0.0.0:4000/organization/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"organization_alias": organization_alias}
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
async def create_team(session, org_id: str):
"""Helper function to create a new team under an organization"""
url = "http://0.0.0.0:4000/team/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"organization_id": org_id, "team_alias": f"test-team-{uuid.uuid4()}"}
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
async def create_user(session, org_id: str):
"""Helper function to create a new user"""
url = "http://0.0.0.0:4000/user/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"user_name": f"test-user-{uuid.uuid4()}"}
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
async def generate_key(session, user_id: str, team_id: str):
"""Helper function to generate a key for a specific user and team"""
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"user_id": user_id, "team_id": team_id}
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
async def chat_completion(session, key: str):
"""Make a chat completion request"""
from openai import AsyncOpenAI
from litellm._uuid import uuid
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000/v1")
response = await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Test message {uuid.uuid4()}"}],
)
return response
async def get_spend_info(session, entity_type: str, entity_id: str):
"""Helper function to get spend information for an entity"""
url = f"http://0.0.0.0:4000/{entity_type}/info"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
if entity_type == "key":
data = {"key": entity_id}
else:
data = {f"{entity_type}_id": entity_id}
async with session.get(url, headers=headers, params=data) as response:
return await response.json()
async def get_proxy_readiness(session):
"""Fetch authenticated readiness details. Used both as a fail-fast gate and as a diagnostic on poll timeout."""
url = "http://0.0.0.0:4000/health/readiness/details"
headers = {"Authorization": "Bearer sk-1234"}
async with session.get(url, headers=headers) as response:
return response.status, await response.json()
async def assert_proxy_healthy(session):
"""Fail fast if the proxy's DB or cache is not reachable — no point running the test."""
status, body = await get_proxy_readiness(session)
if status != 200 or body.get("db") != "connected":
pytest.fail(
f"Proxy /health/readiness/details unhealthy (status={status}). "
f"Cannot run spend accuracy test. Response: {body}"
)
print(f"Proxy readiness OK: {body}")
def compute_expected_spend(responses) -> float:
"""
Compute the expected total spend locally from each response's usage tokens,
using the same pricing table the proxy uses. This is the independent ground
truth we compare the proxy's reported spend against.
"""
total = 0.0
for r in responses:
usage = r.usage
prompt_cost, completion_cost = litellm.cost_per_token(
model=UPSTREAM_MODEL,
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
)
total += prompt_cost + completion_cost
return total
async def poll_key_spend_until(session, key: str, expected: float) -> float:
"""
Poll key spend until it matches `expected` within TOLERANCE, or timeout.
Returns the last observed spend either way; caller decides how to report.
"""
start = time.time()
last_spend = 0.0
while time.time() - start < POLL_TIMEOUT_SECONDS:
try:
key_info = await get_spend_info(session, "key", key)
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
print(
f"Transient transport error during spend poll: "
f"{type(exc).__name__}: {exc}. Retrying... "
f"({time.time() - start:.1f}s elapsed)"
)
await asyncio.sleep(POLL_INTERVAL_SECONDS)
continue
last_spend = key_info["info"]["spend"]
if abs(last_spend - expected) < TOLERANCE:
print(
f"Key spend reached expected {expected} after {time.time() - start:.1f}s"
)
return last_spend
print(
f"Key spend {last_spend}, expected {expected}, waiting... "
f"({time.time() - start:.1f}s elapsed)"
)
await asyncio.sleep(POLL_INTERVAL_SECONDS)
return last_spend
async def fail_with_diagnostics(session, stage: str, expected: float, observed: float):
"""Emit a failure with readiness state so CI output points at the real cause."""
_, readiness = await get_proxy_readiness(session)
pytest.fail(
f"{stage}: key spend did not match expected after {POLL_TIMEOUT_SECONDS}s poll. "
f"expected={expected}, observed={observed}, diff={expected - observed}. "
f"Proxy readiness: {readiness}"
)
@pytest.mark.asyncio
async def test_basic_spend_accuracy():
"""
Test basic spend accuracy across different entities:
1. Create org, team, user, and key
2. Make N requests, keeping each response
3. Compute expected spend locally from response usage (independent ground truth)
4. Poll until proxy-reported spend matches expected
5. Verify spend is consistent across key, team, user, and org entities
"""
NUM_LLM_REQUESTS = 20
async with _make_test_session() as session:
await assert_proxy_healthy(session)
org_response = await create_organization(
session=session, organization_alias=f"test-org-{uuid.uuid4()}"
)
print("org_response: ", org_response)
org_id = org_response["organization_id"]
team_response = await create_team(session, org_id)
print("team_response: ", team_response)
team_id = team_response["team_id"]
user_response = await create_user(session, org_id)
print("user_response: ", user_response)
user_id = user_response["user_id"]
key_response = await generate_key(session, user_id, team_id)
print("key_response: ", key_response)
key = key_response["key"]
responses = []
for i in range(NUM_LLM_REQUESTS):
response = await chat_completion(session, key)
responses.append(response)
print(f"Request {i + 1}/{NUM_LLM_REQUESTS} completed")
expected_spend = compute_expected_spend(responses)
assert expected_spend > 0, (
f"Locally computed expected spend is {expected_spend}. Either cost calc "
f"is broken or upstream returned zero tokens. "
f"Usage: {[r.usage.model_dump() for r in responses]}"
)
print(f"Expected total spend (local ground truth): {expected_spend}")
final_spend = await poll_key_spend_until(session, key, expected_spend)
if abs(final_spend - expected_spend) >= TOLERANCE:
await fail_with_diagnostics(
session,
stage="test_basic_spend_accuracy",
expected=expected_spend,
observed=final_spend,
)
# Allow a final scheduler tick for team/user/org aggregations to settle
await asyncio.sleep(5)
key_info = await get_spend_info(session, "key", key)
print("key_info: ", key_info)
team_info = await get_spend_info(session, "team", team_id)
print("team_info: ", team_info)
user_info = await get_spend_info(session, "user", user_id)
print("user_info: ", user_info)
org_info = await get_spend_info(session, "organization", org_id)
print("org_info: ", org_info)
assert (
abs(key_info["info"]["spend"] - expected_spend) < TOLERANCE
), f"Key spend {key_info['info']['spend']} does not match expected {expected_spend}"
assert (
abs(user_info["user_info"]["spend"] - expected_spend) < TOLERANCE
), f"User spend {user_info['user_info']['spend']} does not match expected {expected_spend}"
assert (
abs(team_info["team_info"]["spend"] - expected_spend) < TOLERANCE
), f"Team spend {team_info['team_info']['spend']} does not match expected {expected_spend}"
assert (
abs(org_info["spend"] - expected_spend) < TOLERANCE
), f"Organization spend {org_info['spend']} does not match expected {expected_spend}"
@pytest.mark.asyncio
async def test_long_term_spend_accuracy_with_bursts():
"""
Test long-term spend accuracy with multiple bursts of requests:
1. Create org, team, user, and key
2. Burst 1: make requests, compute expected locally, verify proxy matches
3. Burst 2: make more requests, verify proxy total == burst1 + burst2
4. Verify total spend is consistent across all entities
"""
BURST_1_REQUESTS = 22
BURST_2_REQUESTS = 12
async with _make_test_session() as session:
await assert_proxy_healthy(session)
org_response = await create_organization(
session=session, organization_alias=f"test-org-{uuid.uuid4()}"
)
print("org_response: ", org_response)
org_id = org_response["organization_id"]
team_response = await create_team(session, org_id)
print("team_response: ", team_response)
team_id = team_response["team_id"]
user_response = await create_user(session, org_id)
print("user_response: ", user_response)
user_id = user_response["user_id"]
key_response = await generate_key(session, user_id, team_id)
print("key_response: ", key_response)
key = key_response["key"]
print(f"Starting first burst of {BURST_1_REQUESTS} requests...")
burst_1_responses = []
for i in range(BURST_1_REQUESTS):
response = await chat_completion(session, key)
burst_1_responses.append(response)
print(f"Burst 1 - Request {i + 1}/{BURST_1_REQUESTS} completed")
burst_1_expected = compute_expected_spend(burst_1_responses)
assert burst_1_expected > 0, (
f"Burst 1 expected spend is {burst_1_expected}. "
f"Usage: {[r.usage.model_dump() for r in burst_1_responses]}"
)
print(f"Burst 1 expected spend: {burst_1_expected}")
final_burst_1 = await poll_key_spend_until(session, key, burst_1_expected)
if abs(final_burst_1 - burst_1_expected) >= TOLERANCE:
await fail_with_diagnostics(
session,
stage="test_long_term_spend_accuracy burst 1",
expected=burst_1_expected,
observed=final_burst_1,
)
print(f"Starting second burst of {BURST_2_REQUESTS} requests...")
burst_2_responses = []
for i in range(BURST_2_REQUESTS):
response = await chat_completion(session, key)
burst_2_responses.append(response)
print(f"Burst 2 - Request {i + 1}/{BURST_2_REQUESTS} completed")
total_expected = burst_1_expected + compute_expected_spend(burst_2_responses)
print(f"Total expected spend (burst 1 + burst 2): {total_expected}")
final_total = await poll_key_spend_until(session, key, total_expected)
if abs(final_total - total_expected) >= TOLERANCE:
await fail_with_diagnostics(
session,
stage="test_long_term_spend_accuracy total",
expected=total_expected,
observed=final_total,
)
await asyncio.sleep(5)
key_info = await get_spend_info(session, "key", key)
team_info = await get_spend_info(session, "team", team_id)
user_info = await get_spend_info(session, "user", user_id)
org_info = await get_spend_info(session, "organization", org_id)
print(f"Final key spend: {key_info['info']['spend']}")
print(f"Final team spend: {team_info['team_info']['spend']}")
print(f"Final user spend: {user_info['user_info']['spend']}")
print(f"Final org spend: {org_info['spend']}")
assert (
abs(key_info["info"]["spend"] - total_expected) < TOLERANCE
), f"Key spend {key_info['info']['spend']} does not match expected {total_expected}"
assert (
abs(user_info["user_info"]["spend"] - total_expected) < TOLERANCE
), f"User spend {user_info['user_info']['spend']} does not match expected {total_expected}"
assert (
abs(team_info["team_info"]["spend"] - total_expected) < TOLERANCE
), f"Team spend {team_info['team_info']['spend']} does not match expected {total_expected}"
assert (
abs(org_info["spend"] - total_expected) < TOLERANCE
), f"Organization spend {org_info['spend']} does not match expected {total_expected}"

View file

@ -1,311 +0,0 @@
import pytest
import asyncio
import aiohttp
import json
from openai import AsyncOpenAI
from litellm._uuid import uuid
from httpx import AsyncClient
import os
TEST_MASTER_KEY = "sk-1234"
PROXY_BASE_URL = "http://0.0.0.0:4000"
@pytest.mark.asyncio
async def test_team_model_alias():
"""
Test model alias functionality with teams:
1. Add a new model with model_name="gpt-4-team1" and litellm_params.model="gpt-4o"
2. Create a new team
3. Update team with model_alias mapping
4. Generate key for team
5. Make request with aliased model name
"""
client = AsyncClient(base_url=PROXY_BASE_URL)
headers = {"Authorization": f"Bearer {TEST_MASTER_KEY}"}
# Add new model
model_response = await client.post(
"/model/new",
json={
"model_name": "gpt-4o-team1",
"litellm_params": {
"model": "gpt-4o",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
headers=headers,
)
assert model_response.status_code == 200
# Create new team
team_response = await client.post(
"/team/new",
json={
"models": ["gpt-4o-team1"],
},
headers=headers,
)
assert team_response.status_code == 200
team_data = team_response.json()
team_id = team_data["team_id"]
# Update team with model alias
update_response = await client.post(
"/team/update",
json={"team_id": team_id, "model_aliases": {"gpt-4o": "gpt-4o-team1"}},
headers=headers,
)
assert update_response.status_code == 200
# Generate key for team
key_response = await client.post(
"/key/generate", json={"team_id": team_id}, headers=headers
)
assert key_response.status_code == 200
key = key_response.json()["key"]
# Make request with model alias
openai_client = AsyncOpenAI(api_key=key, base_url=f"{PROXY_BASE_URL}/v1")
response = await openai_client.chat.completions.create(
model="gpt-4o",
messages=[{"role": "user", "content": f"Test message {uuid.uuid4()}"}],
)
assert response is not None, "Should get valid response when using model alias"
# Cleanup - delete the model
model_id = model_response.json()["model_info"]["id"]
delete_response = await client.post(
"/model/delete",
json={"id": model_id},
headers={"Authorization": f"Bearer {TEST_MASTER_KEY}"},
)
assert delete_response.status_code == 200
@pytest.mark.asyncio
async def test_team_model_association():
"""
Test that models created with a team_id are properly associated with the team:
1. Create a new team
2. Add a model with team_id in model_info
3. Verify the model appears in team info
"""
client = AsyncClient(base_url=PROXY_BASE_URL)
headers = {"Authorization": f"Bearer {TEST_MASTER_KEY}"}
# Create new team
team_response = await client.post(
"/team/new",
json={
"models": [], # Start with empty model list
},
headers=headers,
)
assert team_response.status_code == 200
team_data = team_response.json()
team_id = team_data["team_id"]
# Add new model with team_id
model_response = await client.post(
"/model/new",
json={
"model_name": "gpt-4-team-test",
"litellm_params": {
"model": "gpt-4",
"custom_llm_provider": "openai",
"api_key": "fake_key",
},
"model_info": {"team_id": team_id},
},
headers=headers,
)
assert model_response.status_code == 200
# Get team info and verify model association
team_info_response = await client.get(
f"/team/info",
headers=headers,
params={"team_id": team_id},
)
assert team_info_response.status_code == 200
team_info = team_info_response.json()["team_info"]
print("team_info", json.dumps(team_info, indent=4))
# Verify the model is in team_models
assert (
"gpt-4-team-test" in team_info["models"]
), "Model should be associated with team"
# Cleanup - delete the model
model_id = model_response.json()["model_info"]["id"]
delete_response = await client.post(
"/model/delete",
json={"id": model_id},
headers=headers,
)
assert delete_response.status_code == 200
@pytest.mark.asyncio
async def test_team_model_visibility_in_models_endpoint():
"""
Test that team-specific models are only visible to the correct team in /models endpoint:
1. Create two teams
2. Add a model associated with team1
3. Generate keys for both teams
4. Verify team1's key can see the model in /models
5. Verify team2's key cannot see the model in /models
"""
client = AsyncClient(base_url=PROXY_BASE_URL)
headers = {"Authorization": f"Bearer {TEST_MASTER_KEY}"}
# Create team1
team1_response = await client.post(
"/team/new",
json={"models": []},
headers=headers,
)
assert team1_response.status_code == 200
team1_id = team1_response.json()["team_id"]
# Create team2
team2_response = await client.post(
"/team/new",
json={"models": []},
headers=headers,
)
assert team2_response.status_code == 200
team2_id = team2_response.json()["team_id"]
# Add model associated with team1
model_response = await client.post(
"/model/new",
json={
"model_name": "gpt-4-team-test",
"litellm_params": {
"model": "gpt-4",
"custom_llm_provider": "openai",
"api_key": "fake_key",
},
"model_info": {"team_id": team1_id},
},
headers=headers,
)
assert model_response.status_code == 200
# Generate keys for both teams
team1_key = (
await client.post("/key/generate", json={"team_id": team1_id}, headers=headers)
).json()["key"]
team2_key = (
await client.post("/key/generate", json={"team_id": team2_id}, headers=headers)
).json()["key"]
# Check models visibility for team1's key
team1_models = await client.get(
"/models", headers={"Authorization": f"Bearer {team1_key}"}
)
assert team1_models.status_code == 200
print("team1_models", json.dumps(team1_models.json(), indent=4))
assert any(
model["id"] == "gpt-4-team-test" for model in team1_models.json()["data"]
), "Team1 should see their model"
# Check models visibility for team2's key
team2_models = await client.get(
"/models", headers={"Authorization": f"Bearer {team2_key}"}
)
assert team2_models.status_code == 200
print("team2_models", json.dumps(team2_models.json(), indent=4))
assert not any(
model["id"] == "gpt-4-team-test" for model in team2_models.json()["data"]
), "Team2 should not see team1's model"
# Cleanup
model_id = model_response.json()["model_info"]["id"]
await client.post("/model/delete", json={"id": model_id}, headers=headers)
@pytest.mark.asyncio
async def test_team_model_visibility_in_model_info_endpoint():
"""
Test that team-specific models are visible to all users in /v2/model/info endpoint:
Note: /v2/model/info is used by the Admin UI to display model info
1. Create a team
2. Add a model associated with the team
3. Generate a team key
4. Verify both team key and non-team key can see the model in /v2/model/info
"""
client = AsyncClient(base_url=PROXY_BASE_URL)
headers = {"Authorization": f"Bearer {TEST_MASTER_KEY}"}
# Create team
team_response = await client.post(
"/team/new",
json={"models": []},
headers=headers,
)
assert team_response.status_code == 200
team_id = team_response.json()["team_id"]
# Add model associated with team
model_response = await client.post(
"/model/new",
json={
"model_name": "gpt-4-team-test",
"litellm_params": {
"model": "gpt-4",
"custom_llm_provider": "openai",
"api_key": "fake_key",
},
"model_info": {"team_id": team_id},
},
headers=headers,
)
assert model_response.status_code == 200
# Generate team key
team_key = (
await client.post("/key/generate", json={"team_id": team_id}, headers=headers)
).json()["key"]
# Generate non-team key
non_team_key = (
await client.post("/key/generate", json={}, headers=headers)
).json()["key"]
# Check model info visibility with team key
team_model_info = await client.get(
"/v2/model/info",
headers={"Authorization": f"Bearer {team_key}"},
params={"model_name": "gpt-4-team-test"},
)
assert team_model_info.status_code == 200
team_model_info = team_model_info.json()
print("Team 1 model info", json.dumps(team_model_info, indent=4))
assert any(
model["model_info"].get("team_public_model_name") == "gpt-4-team-test"
for model in team_model_info["data"]
), "Team1 should see their model"
# Check model info visibility with non-team key
non_team_model_info = await client.get(
"/v2/model/info",
headers={"Authorization": f"Bearer {non_team_key}"},
params={"model_name": "gpt-4-team-test"},
)
assert non_team_model_info.status_code == 200
non_team_model_info = non_team_model_info.json()
print("Non-team model info", json.dumps(non_team_model_info, indent=4))
assert any(
model["model_info"].get("team_public_model_name") == "gpt-4-team-test"
for model in non_team_model_info["data"]
), "Non-team should see the model"
# Cleanup
model_id = model_response.json()["model_info"]["id"]
await client.post("/model/delete", json={"id": model_id}, headers=headers)

View file

@ -118,45 +118,6 @@ async def test_end_user_new():
await asyncio.gather(*tasks)
@pytest.mark.asyncio
async def test_aaaend_user_specific_region():
"""
- Specify region user can make calls in
- Make a generic call
- assert returned api base is for model in region
Repeat 3 times
"""
key: str = ""
## CREATE USER ##
async with aiohttp.ClientSession() as session:
end_user_obj = await new_end_user(
session=session,
i=0,
user_id=str(uuid.uuid4()),
model_region="eu",
)
## MAKE CALL ##
key_gen = await generate_key(
session=session, i=0, models=["gpt-5-mini-end-user-test"]
)
key = key_gen["key"]
for _ in range(3):
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000", max_retries=0)
print("SENDING USER PARAM - {}".format(end_user_obj["user_id"]))
result = await client.chat.completions.with_raw_response.create(
model="gpt-5-mini-end-user-test",
messages=[{"role": "user", "content": "Hey!"}],
user=end_user_obj["user_id"],
)
assert result.headers.get("x-litellm-model-region") == "eu"
@pytest.mark.asyncio
async def test_enduser_tpm_limits_non_master_key():
"""

View file

@ -147,55 +147,6 @@ async def test_key_gen_bad_key():
pass
async def update_key(session, get_key, metadata: Optional[dict] = None):
"""
Make sure only models user has access to are returned
"""
url = "http://0.0.0.0:4000/key/update"
headers = {
"Authorization": "Bearer sk-1234",
"Content-Type": "application/json",
}
data = {"key": get_key}
if metadata is not None:
data["metadata"] = metadata
else:
data.update({"models": ["gpt-4"], "duration": "120s"})
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
async def update_proxy_budget(session):
"""
Make sure only models user has access to are returned
"""
url = "http://0.0.0.0:4000/user/update"
headers = {
"Authorization": f"Bearer sk-1234",
"Content-Type": "application/json",
}
data = {"user_id": "litellm-proxy-budget", "spend": 0}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
async def chat_completion(session, key, model="gpt-4"):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
@ -232,39 +183,6 @@ async def chat_completion(session, key, model="gpt-4"):
pass
async def image_generation(session, key, model="gpt-image-1"):
url = "http://0.0.0.0:4000/v1/images/generations"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": model,
"prompt": "A cute baby sea otter",
}
for i in range(3):
try:
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print("/images/generations response", response_text)
print()
if status != 200:
raise Exception(
f"Request did not return a 200 status code: {status}. Response: {response_text}"
)
return await response.json()
except Exception as e:
if "Request did not return a 200 status code" in str(e):
raise e
else:
pass
async def chat_completion_streaming(session, key, model="gpt-4"):
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000")
messages = [
@ -292,29 +210,6 @@ async def chat_completion_streaming(session, key, model="gpt-4"):
return prompt_tokens, completion_tokens
@pytest.mark.parametrize("metadata", [{"test": "new"}, {}])
@pytest.mark.asyncio
async def test_key_update(metadata):
"""
Create key
Update key with new model
Test key w/ model
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, i=0, metadata={"test": "test"})
key = key_gen["key"]
assert key_gen["metadata"]["test"] == "test"
updated_key = await update_key(
session=session,
get_key=key,
metadata=metadata,
)
print(f"updated_key['metadata']: {updated_key['metadata']}")
assert updated_key["metadata"] == metadata
await update_proxy_budget(session=session) # resets proxy spend
await chat_completion(session=session, key=key)
async def delete_key(session, get_key, auth_key="sk-1234"):
"""
Delete key
@ -583,61 +478,6 @@ async def test_aaaaakey_info_spend_values_streaming():
), f"Expected={rounded_response_cost}, Got={rounded_key_info_spend}"
@pytest.mark.flaky(retries=3, delay=1)
@pytest.mark.asyncio
async def test_key_info_spend_values_image_generation():
"""
Test to ensure spend is correctly calculated
- create key
- make image gen call
- assert cost is expected value
"""
async def retry_request(func, *args, _max_attempts=5, **kwargs):
for attempt in range(_max_attempts):
try:
return await func(*args, **kwargs)
except aiohttp.client_exceptions.ClientOSError as e:
if attempt + 1 == _max_attempts:
raise # re-raise the last ClientOSError if all attempts failed
print(f"Attempt {attempt+1} failed, retrying...")
async with aiohttp.ClientSession(
timeout=aiohttp.ClientTimeout(total=600)
) as session:
## Test Spend Update ##
# completion
key_gen = await generate_key(session=session, i=0)
key = key_gen["key"]
response = await image_generation(session=session, key=key)
await asyncio.sleep(5)
key_info = await retry_request(
get_key_info, session=session, get_key=key, call_key=key
)
spend = key_info["info"]["spend"]
assert spend > 0
# The record/replay proxy serves this identical second call from its
# cassette (free), but the proxy must still bill it. Spend logging is
# async/batched, so poll for the increase rather than reading once after a
# fixed sleep; a spend that never grows means the repeat was not billed
# (e.g. the proxy response cache is on), which this still catches.
await image_generation(session=session, key=key)
spend_after = spend
for _ in range(12):
await asyncio.sleep(5)
key_info = await retry_request(
get_key_info, session=session, get_key=key, call_key=key
)
spend_after = key_info["info"]["spend"]
if spend_after > spend:
break
assert spend_after > spend, (
"spend did not increase on an identical repeat image call; the repeat "
"was not billed (the proxy response cache may be on)"
)
@pytest.mark.skip(reason="Frequent check on ci/cd leads to read timeout issue.")
@pytest.mark.asyncio
async def test_key_with_budgets():
@ -684,33 +524,6 @@ async def test_key_with_budgets():
assert reset_at_init_value != reset_at_new_value
@pytest.mark.asyncio
async def test_key_crossing_budget():
"""
- Create key with budget with budget=0.00000001
- make a /chat/completions call
- wait 5s
- make a /chat/completions call - should fail with key crossed it's budget
- Check if value updated
"""
from litellm.proxy.utils import hash_token
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, i=0, budget=0.0000001)
key = key_gen["key"]
hashed_token = hash_token(token=key)
print(f"hashed_token: {hashed_token}")
response = await chat_completion(session=session, key=key)
print("response 1: ", response)
await asyncio.sleep(10)
with pytest.raises(Exception, match="Budget has been exceeded!") as exc_info:
response = await chat_completion(session=session, key=key)
e = exc_info.value
assert "Budget has been exceeded!" in str(e)
@pytest.mark.skip(reason="AWS Suspended Account")
@pytest.mark.asyncio
async def test_key_info_spend_values_sagemaker():
@ -736,32 +549,6 @@ async def test_key_info_spend_values_sagemaker():
# assert rounded_response_cost == rounded_key_info_spend
@pytest.mark.asyncio
async def test_key_rate_limit():
"""
Tests backoff/retry logic on parallel request error.
- Create key with max parallel requests 0
- run 2 requests -> both fail
- Create key with max parallel request 1
- run 2 requests
- both should succeed
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, i=0, max_parallel_requests=0)
new_key = key_gen["key"]
try:
await chat_completion(session=session, key=new_key)
pytest.fail(f"Expected this call to fail")
except Exception as e:
pass
key_gen = await generate_key(session=session, i=0, max_parallel_requests=1)
new_key = key_gen["key"]
try:
await chat_completion(session=session, key=new_key)
except Exception as e:
pytest.fail(f"Expected this call to work - {str(e)}")
@pytest.mark.asyncio
async def test_key_delete_ui():
"""
@ -845,43 +632,3 @@ async def test_key_model_list(model_access, model_access_level, model_endpoint):
assert len(model_list["data"]) == 1
@pytest.mark.asyncio
async def test_key_user_not_in_db():
"""
- Create a key with unique user-id (not in db)
- Check if key can make `/chat/completion` call
"""
my_unique_user = str(uuid.uuid4())
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(
session=session,
i=0,
user_id=my_unique_user,
)
key = key_gen["key"]
try:
await chat_completion(session=session, key=key)
except Exception as e:
pytest.fail(f"Expected this call to work - {str(e)}")
@pytest.mark.asyncio
async def test_key_over_budget():
"""
Test if key over budget is handled as expected.
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, i=0, budget=0.0000001)
key = key_gen["key"]
try:
await chat_completion(session=session, key=key)
except Exception as e:
pytest.fail(f"Expected this call to work - {str(e)}")
## CALL `/models` - expect to work
model_list = await get_key_info(session=session, get_key=key, call_key=key)
## CALL `/chat/completions` - expect to fail
with pytest.raises(Exception, match="Budget has been exceeded!") as exc_info:
await chat_completion(session=session, key=key)
e = exc_info.value
assert "Budget has been exceeded!" in str(e)

View file

@ -106,37 +106,6 @@ async def add_models(
return response_json
async def update_model(
session, model_id="123", model_name="azure-gpt-3.5", key="sk-1234"
):
url = "http://0.0.0.0:4000/model/update"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model_name": model_name,
"litellm_params": {
"model": "openai/gpt-4.1-nano",
"api_key": "os.environ/OPENAI_API_KEY",
},
"model_info": {"id": model_id},
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Add models {response_text}")
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
response_json = await response.json()
return response_json
async def get_model_info(session, key, litellm_model_id=None):
"""
Make sure only models user has access to are returned
@ -301,169 +270,6 @@ async def test_add_and_delete_models():
pass
async def add_model_for_health_checking(session, model_id="123"):
url = "http://0.0.0.0:4000/model/new"
headers = {
"Authorization": f"Bearer sk-1234",
"Content-Type": "application/json",
}
data = {
"model_name": f"azure-model-health-check-{model_id}",
"litellm_params": {
"model": "gpt-4.1-nano",
"api_key": os.getenv("OPENAI_API_KEY"),
},
"model_info": {"id": model_id},
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Add models {response_text}")
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
async def get_model_info_v2(session, key):
url = "http://0.0.0.0:4000/v2/model/info"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
async with session.get(url, headers=headers) as response:
status = response.status
response_text = await response.text()
print("response from v2/model/info")
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
async def get_specific_model_info_v2(session, key, model_name):
url = "http://0.0.0.0:4000/v2/model/info?debug=True&model=" + model_name
print("running /model/info check for model=", model_name)
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
async with session.get(url, headers=headers) as response:
status = response.status
response_text = await response.text()
print("response from v2/model/info")
print(response_text)
print()
_json_response = await response.json()
print("JSON response from /v2/model/info?model=", model_name, _json_response)
_model_info = _json_response["data"]
assert len(_model_info) == 1, f"Expected 1 model, got {len(_model_info)}"
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return _model_info[0]
async def get_model_health(session, key, model_name):
url = "http://0.0.0.0:4000/health?model=" + model_name
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
async with session.get(url, headers=headers) as response:
status = response.status
response_text = await response.json()
print("response from /health?model=", model_name)
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return response_text
@pytest.mark.asyncio
async def test_add_model_run_health():
"""
Add model
Call /model/info and v2/model/info
-> Admin UI calls v2/model/info
Call /chat/completions
Call /health
-> Ensure the health check for the endpoint is working as expected
"""
from litellm._uuid import uuid
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session)
key = key_gen["key"]
master_key = "sk-1234"
model_id = str(uuid.uuid4())
model_name = f"azure-model-health-check-{model_id}"
print("adding model", model_name)
await add_model_for_health_checking(session=session, model_id=model_id)
_old_model_info = await get_specific_model_info_v2(
session=session, key=key, model_name=model_name
)
print("model info before test", _old_model_info)
await asyncio.sleep(30)
print("calling /model/info")
await get_model_info(session=session, key=key)
print("calling v2/model/info")
await get_model_info_v2(session=session, key=key)
print("calling /chat/completions -> expect to work")
await chat_completion(session=session, key=key, model=model_name)
print("calling /health?model=", model_name)
_health_info = await get_model_health(
session=session, key=master_key, model_name=model_name
)
_healthy_endpooint = _health_info["healthy_endpoints"][0]
assert _health_info["healthy_count"] == 1
assert (
_healthy_endpooint["model"] == "gpt-4.1-nano"
) # this is the model that got added
# assert httpx client is is unchanges
await asyncio.sleep(10)
_model_info_after_test = await get_specific_model_info_v2(
session=session, key=key, model_name=model_name
)
print("model info after test", _model_info_after_test)
old_openai_client = _old_model_info["openai_client"]
new_openai_client = _model_info_after_test["openai_client"]
print("old openai client", old_openai_client)
print("new openai client", new_openai_client)
"""
PROD TEST - This is extremly important
The OpenAI client used should be the same after 30 seconds
It is a serious bug if the openai client does not match here
"""
assert (
old_openai_client == new_openai_client
), "OpenAI client does not match for the same model after 30 seconds"
# cleanup
await delete_model(session=session, model_id=model_id)
@pytest.mark.asyncio
async def test_get_personal_models_for_user():
"""
@ -506,52 +312,3 @@ async def test_model_group_info_e2e():
)
@pytest.mark.asyncio
async def test_team_model_e2e():
"""
Test team model e2e
- create team
- create user
- add user to team as admin
- add model to team
- update model
- delete model
"""
from tests.test_users import new_user
from tests.test_team import new_team
from litellm._uuid import uuid
async with aiohttp.ClientSession() as session:
# Creat a user
user_data = await new_user(session=session, i=0)
user_id = user_data["user_id"]
user_api_key = user_data["key"]
# Create a team
member_list = [
{"role": "admin", "user_id": user_id},
]
team_data = await new_team(session=session, member_list=member_list, i=0)
team_id = team_data["team_id"]
model_id = str(uuid.uuid4())
model_name = "my-test-model"
# Add model to team
model_data = await add_models(
session=session,
model_id=model_id,
model_name=model_name,
key=user_api_key,
team_id=team_id,
)
model_id = model_data["model_id"]
# Update model
model_data = await update_model(
session=session, model_id=model_id, model_name=model_name, key=user_api_key
)
model_id = model_data["model_id"]
# Delete model
await delete_model(session=session, model_id=model_id, key=user_api_key)

View file

@ -523,23 +523,6 @@ async def test_image_generation():
await image_generation(session=session, key=key_2)
@pytest.mark.flaky(retries=5, delay=1)
@pytest.mark.asyncio
async def test_openai_wildcard_chat_completion():
"""
- Create key for model = "*" -> this has access to all models
- proxy_server_config.yaml has model = *
- Make chat completion call
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, models=["*"])
key = key_gen["key"]
# call chat/completions with a model that the key was not created for + the model is not on the config.yaml
await chat_completion(session=session, key=key, model="gpt-3.5-turbo-0125")
@pytest.mark.asyncio
async def test_proxy_all_models():
"""

View file

@ -1,319 +0,0 @@
# What this tests ?
## Tests /organization endpoints.
import pytest
import asyncio
import aiohttp
import time, uuid
from openai import AsyncOpenAI
async def new_user(
session,
i,
user_id=None,
budget=None,
budget_duration=None,
models=["azure-models"],
team_id=None,
user_email=None,
):
url = "http://0.0.0.0:4000/user/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"models": models,
"aliases": {"mistral-7b": "gpt-3.5-turbo"},
"duration": None,
"max_budget": budget,
"budget_duration": budget_duration,
"user_email": user_email,
}
if user_id is not None:
data["user_id"] = user_id
if team_id is not None:
data["team_id"] = team_id
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response {i} (Status code: {status}):")
print(response_text)
print()
if status != 200:
raise Exception(
f"Request {i} did not return a 200 status code: {status}, response: {response_text}"
)
return await response.json()
async def new_organization(session, i, organization_alias, max_budget=None):
url = "http://0.0.0.0:4000/organization/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"organization_alias": organization_alias,
"models": ["azure-models"],
"max_budget": max_budget,
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response {i} (Status code: {status}):")
print(response_text)
print()
if status != 200:
raise Exception(f"Request {i} did not return a 200 status code: {status}")
return await response.json()
async def add_member_to_org(
session, i, organization_id, user_id, user_role="internal_user"
):
url = "http://0.0.0.0:4000/organization/member_add"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"organization_id": organization_id,
"member": {
"user_id": user_id,
"role": user_role,
},
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response {i} (Status code: {status}):")
print(response_text)
print()
if status != 200:
raise Exception(f"Request {i} did not return a 200 status code: {status}")
return await response.json()
async def update_member_role(
session, i, organization_id, user_id, user_role="internal_user"
):
url = "http://0.0.0.0:4000/organization/member_update"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"organization_id": organization_id,
"user_id": user_id,
"role": user_role,
}
async with session.patch(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response {i} (Status code: {status}):")
print(response_text)
print()
if status != 200:
raise Exception(f"Request {i} did not return a 200 status code: {status}")
return await response.json()
async def delete_member_from_org(session, i, organization_id, user_id):
url = "http://0.0.0.0:4000/organization/member_delete"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {
"organization_id": organization_id,
"user_id": user_id,
}
async with session.delete(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response {i} (Status code: {status}):")
print(response_text)
print()
if status != 200:
raise Exception(f"Request {i} did not return a 200 status code: {status}")
return await response.json()
async def delete_organization(session, i, organization_id):
url = "http://0.0.0.0:4000/organization/delete"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {"organization_ids": [organization_id]}
async with session.delete(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response {i} (Status code: {status}):")
print(response_text)
print()
if status != 200:
raise Exception(f"Request {i} did not return a 200 status code: {status}")
return await response.json()
async def list_organization(session, i):
url = "http://0.0.0.0:4000/organization/list"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
async with session.get(url, headers=headers) as response:
status = response.status
response_json = await response.json()
print(f"Response {i} (Status code: {status}):")
print()
if status != 200:
raise Exception(f"Request {i} did not return a 200 status code: {status}")
# Assert that budget info is returned for each organization
for org in response_json:
assert (
"litellm_budget_table" in org
), "Missing budget info in organization response"
# Optionally also check that it's not null
assert org["litellm_budget_table"] is not None, "Budget info is None"
return response_json
@pytest.mark.flaky(retries=5, delay=1)
@pytest.mark.asyncio
async def test_organization_new():
"""
Make 20 parallel calls to /organization/new. Assert all worked.
"""
organization_alias = f"Organization: {uuid.uuid4()}"
async with aiohttp.ClientSession() as session:
tasks = [
new_organization(
session=session, i=0, organization_alias=organization_alias
)
for i in range(1, 20)
]
await asyncio.gather(*tasks)
@pytest.mark.asyncio
async def test_organization_list():
"""
create 2 new Organizations
check if the Organization list is not empty
"""
organization_alias = f"Organization: {uuid.uuid4()}"
async with aiohttp.ClientSession() as session:
tasks = [
new_organization(
session=session, i=0, organization_alias=organization_alias
)
for i in range(1, 2)
]
await asyncio.gather(*tasks)
response_json = await list_organization(session, i=0)
print(len(response_json))
if len(response_json) == 0:
raise Exception("Return empty list of organization")
@pytest.mark.asyncio
async def test_organization_delete():
"""
create a new organization
delete the organization
check if the Organization list is set
"""
organization_alias = f"Organization: {uuid.uuid4()}"
async with aiohttp.ClientSession() as session:
tasks = [
new_organization(
session=session, i=0, organization_alias=organization_alias
)
]
await asyncio.gather(*tasks)
response_json = await list_organization(session, i=0)
print(len(response_json))
organization_id = response_json[0]["organization_id"]
await delete_organization(session, i=0, organization_id=organization_id)
response_json = await list_organization(session, i=0)
print(len(response_json))
@pytest.mark.asyncio
async def test_organization_member_flow():
"""
create a new organization
add a new member to the organization
check if the member is added to the organization
update the member's role in the organization
delete the member from the organization
check if the member is deleted from the organization
"""
organization_alias = f"Organization: {uuid.uuid4()}"
async with aiohttp.ClientSession() as session:
response_json = await new_organization(
session=session, i=0, organization_alias=organization_alias
)
organization_id = response_json["organization_id"]
response_json = await list_organization(session, i=0)
print(len(response_json))
new_user_response_json = await new_user(
session=session, i=0, user_email=f"test_user_{uuid.uuid4()}@example.com"
)
user_id = new_user_response_json["user_id"]
await add_member_to_org(
session, i=0, organization_id=organization_id, user_id=user_id
)
response_json = await list_organization(session, i=0)
print(len(response_json))
for orgs in response_json:
tmp_organization_id = orgs["organization_id"]
if (
tmp_organization_id is not None
and tmp_organization_id == organization_id
):
user_id = orgs["members"][0]["user_id"]
response_json = await list_organization(session, i=0)
print(len(response_json))
await update_member_role(
session,
i=0,
organization_id=organization_id,
user_id=user_id,
user_role="org_admin",
)
response_json = await list_organization(session, i=0)
print(len(response_json))
await delete_member_from_org(
session, i=0, organization_id=organization_id, user_id=user_id
)
response_json = await list_organization(session, i=0)
print(len(response_json))

View file

@ -221,23 +221,6 @@ async def get_predict_spend_logs(session):
return await response.json()
async def get_spend_report(session, start_date, end_date):
url = "http://0.0.0.0:4000/global/spend/report"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
async with session.get(
url, headers=headers, params={"start_date": start_date, "end_date": end_date}
) as response:
status = response.status
response_text = await response.text()
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
@pytest.mark.skip(reason="datetime in ci/cd gets set weirdly")
@pytest.mark.asyncio
async def test_get_predicted_spend_logs():
@ -308,37 +291,3 @@ async def test_spend_logs_high_traffic():
raise Exception("it worked!")
@pytest.mark.asyncio
async def test_spend_report_endpoint():
async with aiohttp.ClientSession(
timeout=aiohttp.ClientTimeout(total=600)
) as session:
import datetime
todays_date = datetime.date.today() + datetime.timedelta(days=1)
todays_date = todays_date.strftime("%Y-%m-%d")
print("todays_date", todays_date)
thirty_days_ago = (
datetime.date.today() - datetime.timedelta(days=30)
).strftime("%Y-%m-%d")
spend_report = await get_spend_report(
session=session, start_date=thirty_days_ago, end_date=todays_date
)
print("spend report", spend_report)
for row in spend_report:
date = row["group_by_day"]
teams = row["teams"]
for team in teams:
team_name = team["team_name"]
total_spend = team["total_spend"]
metadata = team["metadata"]
assert team_name is not None
print(f"Date: {date}")
print(f"Team: {team_name}")
print(f"Total Spend: {total_spend}")
print("Metadata: ", metadata)
print()

View file

@ -690,40 +690,6 @@ async def test_member_delete(dimension):
assert user_in_team is True
@pytest.mark.asyncio
async def test_team_alias():
"""
- Create team w/ model alias
- Create key for team
- Check if key works
"""
async with aiohttp.ClientSession() as session:
## Create admin
admin_user = f"{uuid.uuid4()}"
await new_user(session=session, i=0, user_id=admin_user)
## Create normal user
normal_user = f"{uuid.uuid4()}"
await new_user(session=session, i=0, user_id=normal_user)
## Create team with 1 admin and 1 user
member_list = [
{"role": "admin", "user_id": admin_user},
{"role": "user", "user_id": normal_user},
]
team_data = await new_team(
session=session,
i=0,
member_list=member_list,
model_aliases={"cheap-model": "gpt-3.5-turbo"},
)
## Create key
key_gen = await generate_key(
session=session, i=0, team_id=team_data["team_id"], models=["gpt-3.5-turbo"]
)
key = key_gen["key"]
## Test key
response = await chat_completion(session=session, key=key, model="cheap-model")
@pytest.mark.asyncio
async def test_users_in_team_budget():
"""

View file

@ -40,51 +40,6 @@ async def new_user(
return await response.json()
async def generate_key(
session,
i,
budget=None,
budget_duration=None,
models=["azure-models", "gpt-4", "dall-e-3"],
max_parallel_requests: Optional[int] = None,
user_id: Optional[str] = None,
team_id: Optional[str] = None,
metadata: Optional[dict] = None,
calling_key="sk-1234",
):
url = "http://0.0.0.0:4000/key/generate"
headers = {
"Authorization": f"Bearer {calling_key}",
"Content-Type": "application/json",
}
data = {
"models": models,
"aliases": {"mistral-7b": "gpt-3.5-turbo"},
"duration": None,
"max_budget": budget,
"budget_duration": budget_duration,
"max_parallel_requests": max_parallel_requests,
"user_id": user_id,
"team_id": team_id,
"metadata": metadata,
}
print(f"data: {data}")
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response {i} (Status code: {status}):")
print(response_text)
print()
if status != 200:
raise Exception(f"Request {i} did not return a 200 status code: {status}")
return await response.json()
@pytest.mark.asyncio
async def test_user_new():
"""
@ -260,62 +215,6 @@ async def test_global_proxy_budget_update():
assert new_new_spend > new_spend
@pytest.mark.asyncio
async def test_user_model_access():
"""
- Create user with model access
- Create key with user
- Call model that user has access to -> should work
- Call wildcard model that user has access to -> should work
- Call model that user does not have access to -> should fail
- Call wildcard model that user does not have access to -> should fail
"""
import openai
async with aiohttp.ClientSession() as session:
get_user = f"krrish_{time.time()}@berri.ai"
await new_user(
session=session,
i=0,
user_id=get_user,
models=["good-model", "anthropic/*"],
)
result = await generate_key(
session=session,
i=0,
user_id=get_user,
models=[], # assign no models. Allow inheritance from user
)
key = result["key"]
await chat_completion(
session=session,
key=key,
model="anthropic/claude-haiku-4-5-20251001",
)
await chat_completion(
session=session,
key=key,
model="good-model",
)
with pytest.raises(openai.PermissionDeniedError):
await chat_completion(
session=session,
key=key,
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
)
with pytest.raises(openai.PermissionDeniedError):
await chat_completion(
session=session,
key=key,
model="groq/claude-3-5-haiku-20241022",
)
import json
from litellm._uuid import uuid
import pytest