test: move 76 live top-level tests to offline unit and integration coverage (#45352)

* test: move 76 live top-level tests to offline unit and integration coverage

* test: keep the mock-param opt-in guard valid once the top-level senders are gone

* test: address review on offline replacements

* docs(test): drop the deleted harness test file from the harness check command

* test: make the key rebind, team member delete and routes integration tests exercise the legacy paths

* test: make fallback, rpm and spend integration contracts deterministic and clean up their rows, assert the rpm limit in usage-based routing

* test: expect the no-deployments error at the rpm limit and cover the strategy check without pre-call checks

* test: hand member cleanups to the scenario instead of growing a budget list, flatten callback kinds

* test: scope the admin health check to the test's own deployment

---------

Co-authored-by: yuneng <yuneng@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 10:33:35 -07:00 • committed by GitHub
parent a81f507bca
commit d5173e3d70
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
47 changed files with 1746 additions and 4409 deletions

View file

@ -0,0 +1,147 @@
import json
import re
import uuid
from collections import Counter
from itertools import chain
from pathlib import Path
from typing import Final
from pydantic import JsonValue
from tests.integration._support.client import Gateway, object_value
from tests.integration._support.process import owned_proxy
SAMPLES: Final = 4
REQUESTS_PER_INTERVAL: Final = 5
LEAK_MIN_NET_GROWTH: Final = 5
LEAK_MIN_GROWING_INTERVALS: Final = 2
ADDRESS: Final = re.compile(r" at 0x[0-9a-fA-F]+")
OBJECT: Final = re.compile(r"<([\w.]+) object")
BOUND_METHOD: Final = re.compile(r"bound method ([\w.]+)")
def _callback_type(text: str) -> str:
stripped: Final = ADDRESS.sub("", text)
if (instance := OBJECT.search(stripped)) is not None:
return instance.group(1).split(".")[-1]
if (method := BOUND_METHOD.search(stripped)) is not None:
return method.group(1)
return stripped.strip()
def _sample(candidate: Gateway) -> tuple[Counter[str], int]:
response: Final = candidate.request("GET", "/active/callbacks")
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
callbacks: Final = body["all_litellm_callbacks"]
alerting: Final = body["num_alerting"]
assert isinstance(callbacks, list) and isinstance(alerting, int), body
return Counter(_callback_type(str(callback)) for callback in callbacks), alerting
Samples = tuple[Counter[str], ...]
def _kinds(samples: Samples) -> frozenset[str]:
return frozenset(chain.from_iterable(samples))
def _series(samples: Samples, kind: str) -> tuple[int, ...]:
return tuple(sample.get(kind, 0) for sample in samples)
def _deltas(series: tuple[int, ...]) -> tuple[int, ...]:
return tuple(after - before for before, after in zip(series, series[1:]))
def _grows(series: tuple[int, ...]) -> bool:
deltas: Final = _deltas(series)
return (
all(delta >= 0 for delta in deltas)
and series[-1] - series[0] >= LEAK_MIN_NET_GROWTH
and sum(1 for delta in deltas if delta > 0) >= LEAK_MIN_GROWING_INTERVALS
)
def _grows_only_in_last_interval(series: tuple[int, ...]) -> bool:
deltas: Final = _deltas(series)
return all(delta >= 0 for delta in deltas) and [
index for index, delta in enumerate(deltas) if delta > 0
] == [len(deltas) - 1]
def _leaking(samples: Samples) -> dict[str, tuple[int, ...]]:
return {kind: _series(samples, kind) for kind in sorted(_kinds(samples)) if _grows(_series(samples, kind))}
def _config(directory: Path, upstream_url: str, model: str, extra: dict[str, JsonValue]) -> Path:
config: Final = directory / f"callback_leak_{uuid.uuid4().hex}.yaml"
general_settings: Final = extra.get("general_settings")
config.write_text(
json.dumps(
{
"model_list": [
{
"model_name": model,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": f"{upstream_url}/v1",
"api_key": "integration-provider-key",
},
}
],
"router_settings": extra.get("router_settings", {}),
"general_settings": {
"master_key": "os.environ/LITELLM_MASTER_KEY",
"database_url": "os.environ/DATABASE_URL",
**(general_settings if isinstance(general_settings, dict) else {}),
},
}
)
)
return config
def _interval(candidate: Gateway, model: str, index: int) -> tuple[Counter[str], int]:
for request in range(REQUESTS_PER_INTERVAL if index else 0):
assert candidate.chat(model, text=f"leak probe {index} {request}")["model"] == model
return _sample(candidate)
def _sample_under_traffic(candidate: Gateway, model: str) -> tuple[Samples, tuple[int, ...]]:
taken: Final = tuple(_interval(candidate, model, index) for index in range(SAMPLES))
callbacks: Final = tuple(sample for sample, _ in taken)
late_growth: Final = any(_grows_only_in_last_interval(_series(callbacks, kind)) for kind in _kinds(callbacks))
confirmed: Final = taken + ((_interval(candidate, model, SAMPLES),) if late_growth else ())
return tuple(sample for sample, _ in confirmed), tuple(alerting for _, alerting in confirmed)
def test_callback_registry_does_not_grow_with_traffic(gateway: Gateway, tmp_path: Path) -> None:
model: Final = f"integration-callback-leak-{uuid.uuid4().hex}"
config: Final = _config(tmp_path, gateway.upstream_url, model, {})
with owned_proxy(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as candidate:
samples, _ = _sample_under_traffic(candidate, model)
assert sum(samples[0].values()) > 0
assert _leaking(samples) == {}, samples
def test_callback_registry_does_not_grow_under_latency_routing_with_alerting(gateway: Gateway, tmp_path: Path) -> None:
model: Final = f"integration-callback-leak-{uuid.uuid4().hex}"
config: Final = _config(
tmp_path,
gateway.upstream_url,
model,
{
"router_settings": {"routing_strategy": "latency-based-routing"},
"general_settings": {
"alert_to_webhook_url": {"llm_exceptions": "http://127.0.0.1:9/integration-alerts"},
"alert_types": ["llm_exceptions", "db_exceptions"],
},
},
)
with owned_proxy(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as candidate:
samples, alerts = _sample_under_traffic(candidate, model)
assert sum(samples[0].values()) > 0
assert any(kind.startswith("LowestLatencyLoggingHandler") for kind in samples[0]), samples[0]
assert _leaking(samples) == {}, samples
assert len(set(alerts)) == 1, alerts

View file

@ -0,0 +1,152 @@
import uuid
from concurrent.futures import ThreadPoolExecutor
from hashlib import sha256
from typing import Final, Literal
import pytest
from pydantic import JsonValue
from tests.integration._support.client import Gateway, Scenario, delete_key_if_present, object_value, string_value
from tests.integration._support.database import read_rows
KEY_ALIAS: Final = "mistral-7b"
ModelAccess = Literal["all-team-models", "single-model"]
AccessLevel = Literal["key", "team"]
ModelEndpoint = Literal["/v1/models", "/model/info"]
def _owned_key(
gateway: Gateway, scenario: Scenario, fields: dict[str, JsonValue], *, key: str | None = None
) -> dict[str, JsonValue]:
response: Final = gateway.request("POST", "/key/generate", fields, key=key)
assert response.status_code == 200, response.text
created: Final = object_value(response.json())
scenario.cleanups.callback(delete_key_if_present, gateway, string_value(created["key"]))
return created
def _data(gateway: Gateway, path: str, key: str, params: dict[str, str] | None = None) -> list[JsonValue]:
response: Final = gateway.request("GET", path, key=key, params=params)
assert response.status_code == 200, response.text
data: Final = object_value(response.json())["data"]
assert isinstance(data, list)
return data
def _verification_row(key: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT user_id FROM "LiteLLM_VerificationToken" WHERE token = %s',
(sha256(key.encode()).hexdigest(),),
)
def test_concurrent_key_generation_persists_every_key(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
with ThreadPoolExecutor(max_workers=10) as pool:
keys: Final = list(pool.map(lambda _: scenario.key(models=[model]), range(10)))
assert len(set(keys)) == 10
for key in keys:
assert len(_verification_row(key)) == 1
def test_generated_key_exposes_hashed_token_and_timestamps(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
created: Final = _owned_key(gateway, scenario, {})
key: Final = string_value(created["key"])
assert created["token"] is not None
assert created["token"] != key
assert created["token"] == sha256(key.encode()).hexdigest()
assert created["token_id"] is not None
assert created["created_at"] is not None
assert created["updated_at"] is not None
def test_key_info_serves_admin_and_self_and_hides_unknown_keys(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
key: Final = scenario.key()
digest: Final = sha256(key.encode()).hexdigest()
admin: Final = gateway.request("GET", "/key/info", params={"key": key})
assert admin.status_code == 200, admin.text
assert object_value(admin.json())["key"] == key
explicit: Final = gateway.request("GET", "/key/info", params={"key": key}, key=key)
assert explicit.status_code == 200, explicit.text
assert object_value(explicit.json())["key"] == key
implicit: Final = gateway.request("GET", "/key/info", key=key)
assert implicit.status_code == 200, implicit.text
assert object_value(implicit.json())["key"] == digest
unknown: Final = gateway.request("GET", "/key/info", params={"key": f"sk-{uuid.uuid4()}"}, key=key)
assert unknown.status_code == 404, unknown.text
def test_model_info_is_filtered_to_the_models_a_key_can_use(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
admin_models: Final = _data(gateway, "/model/info", gateway.key)
user_models: Final = _data(gateway, "/model/info", key)
assert len(admin_models) > len(user_models)
assert [object_value(entry)["model_name"] for entry in user_models] == [model]
def test_proxy_admin_user_key_deletes_a_key_owned_by_another_user(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
admin_user: Final = scenario.user(user_role="proxy_admin")
owner: Final = scenario.user(user_role="internal_user")
admin_key: Final = string_value(_owned_key(gateway, scenario, {"user_id": admin_user})["key"])
victim: Final = string_value(_owned_key(gateway, scenario, {"user_id": owner})["key"])
deleted: Final = gateway.request("POST", "/key/delete", {"keys": [victim]}, key=admin_key)
assert deleted.status_code == 200, deleted.text
assert _verification_row(victim) == []
@pytest.mark.parametrize("model_endpoint", ["/v1/models", "/model/info"])
@pytest.mark.parametrize("access_level", ["key", "team"])
@pytest.mark.parametrize("model_access", ["all-team-models", "single-model"])
def test_key_model_list_follows_key_or_team_access(
gateway: Gateway, model_access: ModelAccess, access_level: AccessLevel, model_endpoint: ModelEndpoint
) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
granted: Final[list[JsonValue]] = [] if model_access == "all-team-models" else [model]
team: Final = scenario.team(models=granted if access_level == "team" else [])
key: Final = scenario.key(
team_id=team,
models=granted if access_level == "key" else [],
aliases={KEY_ALIAS: model},
)
data: Final = _data(gateway, model_endpoint, key)
if model_access == "all-team-models":
assert len(data) > 1
if model_endpoint == "/v1/models":
assert all(isinstance(object_value(entry)["id"], str) for entry in data)
assert model in {object_value(entry)["id"] for entry in data}
else:
assert model in {object_value(entry)["model_name"] for entry in data}
elif model_endpoint == "/v1/models":
assert {object_value(entry)["id"] for entry in data} == {model, KEY_ALIAS}
else:
assert [object_value(entry)["model_name"] for entry in data] == [model]
def test_internal_user_cannot_reassign_its_key_to_another_user(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
first: Final = scenario.user(user_role="internal_user")
second: Final = scenario.user(user_role="internal_user")
key: Final = string_value(_owned_key(gateway, scenario, {"user_id": first})["key"])
own_key: Final = string_value(_owned_key(gateway, scenario, {}, key=key)["key"])
assert _verification_row(own_key) == [{"user_id": first}]
update: Final = gateway.request("POST", "/key/update", {"key": own_key, "user_id": second}, key=key)
assert update.status_code == 403, update.text
assert _verification_row(own_key) == [{"user_id": first}]
def test_internal_user_cannot_delete_another_users_key(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
first: Final = scenario.user(user_role="internal_user")
second: Final = scenario.user(user_role="internal_user")
victim: Final = string_value(_owned_key(gateway, scenario, {"user_id": first})["key"])
attacker: Final = string_value(_owned_key(gateway, scenario, {"user_id": second})["key"])
deleted: Final = gateway.request("POST", "/key/delete", {"keys": [victim]}, key=attacker)
assert deleted.status_code == 403, deleted.text
assert _verification_row(victim) == [{"user_id": first}]

View file

@ -0,0 +1,138 @@
import uuid
from datetime import datetime, timezone
from typing import Final
from pydantic import JsonValue
from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time
from tests.integration._support.client import Gateway, Scenario, delete_key_if_present, object_value, string_value
from tests.integration._support.database import read_rows
def _utc(text: JsonValue) -> datetime:
parsed: Final = datetime.fromisoformat(string_value(text))
return parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=timezone.utc)
def _model_id(model_name: str) -> str:
rows: Final = read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_name = %s', (model_name,))
assert len(rows) == 1, rows
return string_value(rows[0]["model_id"])
def _data(gateway: Gateway, path: str, key: str, params: dict[str, str] | None = None) -> list[dict[str, JsonValue]]:
response: Final = gateway.request("GET", path, key=key, params=params)
assert response.status_code == 200, response.text
data: Final = object_value(response.json())["data"]
assert isinstance(data, list)
return [object_value(entry) for entry in data]
def _wildcard_model(gateway: Gateway, scenario: Scenario, prefix: str) -> None:
created: Final = gateway.post(
"/model/new",
{
"model_name": f"{prefix}/*",
"litellm_params": {
"model": "anthropic/*",
"api_key": "integration-provider-key",
"api_base": gateway.upstream_url,
},
},
)
scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"]))
def test_budget_duration_schedules_reset_at_the_next_standardized_boundary(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
budget_id: Final = scenario.budget(max_budget=10.0, budget_duration="1d")
rows: Final = read_rows(
'SELECT created_at::text AS created_at, budget_reset_at::text AS reset_at FROM "LiteLLM_BudgetTable" '
"WHERE budget_id = %s",
(budget_id,),
)
assert len(rows) == 1, rows
assert rows[0]["reset_at"] is not None, rows
expected: Final = get_next_standardized_reset_time("1d", _utc(rows[0]["created_at"]), "UTC")
assert abs((_utc(rows[0]["reset_at"]) - expected).total_seconds()) <= 3, (rows, expected)
def test_admin_health_counts_every_deployment_and_reports_a_live_one_healthy(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
model_id: Final = _model_id(model)
report: Final = gateway.get("/health", {"model": model})
assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report
healthy: Final = report["healthy_endpoints"]
assert isinstance(healthy, list) and len(healthy) == 1, report
assert object_value(healthy[0]).get("model_id") == model_id, report
def test_routes_listing_is_served_without_credentials(gateway: Gateway) -> None:
response: Final = gateway.client.get("/routes")
assert response.status_code == 200, response.text
routes: Final = object_value(response.json())["routes"]
assert isinstance(routes, list)
assert {"/routes", "/key/generate"} <= {object_value(route)["path"] for route in routes}
def test_unrestricted_key_lists_models_and_none_when_only_access_groups_are_requested(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
grouped: Final = scenario.model(model_info={"access_groups": [f"integration-{uuid.uuid4().hex}"]})
plain: Final = scenario.model()
key: Final = scenario.key()
listed: Final = {entry["id"] for entry in _data(gateway, "/models", key)}
assert {grouped, plain} <= listed
assert _data(gateway, "/models", key, {"only_model_access_groups": "True"}) == []
def test_model_info_by_id_matches_the_entry_in_the_keys_listing(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
scenario.model()
model_id: Final = _model_id(model)
key: Final = scenario.key(models=[model])
listing: Final = _data(gateway, "/model/info", key)
assert {entry["model_name"] for entry in listing} == {model}
listed: Final = [entry for entry in listing if object_value(entry["model_info"])["id"] == model_id]
assert len(listed) == 1
by_id: Final = _data(gateway, "/model/info", key, {"litellm_model_id": model_id})
assert by_id == listed
admin_by_id: Final = _data(gateway, "/model/info", gateway.key, {"litellm_model_id": model_id})
assert [object_value(entry["model_info"])["id"] for entry in admin_by_id] == [model_id]
def test_model_group_info_for_a_personal_user_key_lists_only_the_users_model(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
scenario.model()
created: Final = gateway.post("/user/new", {"user_id": f"integration-{uuid.uuid4().hex}", "models": [model]})
scenario.cleanups.callback(scenario.delete_user, string_value(created["user_id"]))
key: Final = string_value(created["key"])
scenario.cleanups.callback(delete_key_if_present, gateway, key)
groups: Final = [entry["model_group"] for entry in _data(gateway, "/model_group/info", key)]
assert groups == [model]
def test_model_group_info_expands_a_wildcard_deployment_into_concrete_groups(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
prefix: Final = f"integration{uuid.uuid4().hex}"
_wildcard_model(gateway, scenario, prefix)
groups: Final = {
string_value(entry["model_group"]) for entry in _data(gateway, "/model_group/info", gateway.key)
}
assert f"{prefix}/*" not in groups
assert any(group.startswith(f"{prefix}/claude") for group in groups), sorted(groups)[:20]
def test_azure_deployment_route_denies_a_model_outside_the_keys_list(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
allowed: Final = scenario.model()
other: Final = scenario.model()
key: Final = scenario.key(models=[allowed])
body: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": "integration control"}]}
served: Final = gateway.request("POST", f"/openai/deployments/{allowed}/chat/completions", body, key=key)
assert served.status_code == 200, served.text
denied: Final = gateway.request("POST", f"/openai/deployments/{other}/chat/completions", body, key=key)
assert denied.status_code == 403, denied.text
assert "is not available for this API key" in denied.text

View file

@ -0,0 +1,205 @@
import uuid
from concurrent.futures import ThreadPoolExecutor
from typing import Final, Literal
import pytest
from pydantic import JsonValue
from tests.integration._support.client import Gateway, Scenario, delete_key_if_present, object_value, string_value
from tests.integration._support.database import read_rows
MemberDimension = Literal["user_id", "user_email"]
UNCHANGED_BY_ALIAS_UPDATE_SKIP: Final = frozenset(
{
"team_alias",
"members_with_roles",
"created_at",
"updated_at",
"model_spend",
"model_max_budget",
"model_id",
"litellm_organization_table",
"object_permission_id",
"object_permission",
"litellm_model_table",
"policies",
"allow_team_guardrail_config",
"projects",
}
)
def _team_info(gateway: Gateway, team_id: str) -> dict[str, JsonValue]:
return object_value(gateway.get("/team/info", {"team_id": team_id})["team_info"])
def _member_ids(gateway: Gateway, team_id: str) -> list[str | None]:
members: Final = _team_info(gateway, team_id)["members_with_roles"]
assert isinstance(members, list)
return [
member_id if isinstance(member_id := object_value(member).get("user_id"), str) else None for member in members
]
def _user_team_ids(gateway: Gateway, user_id: str) -> list[str]:
teams: Final = gateway.get("/user/info", {"user_id": user_id})["teams"]
assert isinstance(teams, list)
return [string_value(object_value(team)["team_id"]) for team in teams]
def _delete_users_if_present(gateway: Gateway, user_ids: list[str]) -> None:
present: Final = [
user_id
for user_id in user_ids
if read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id = %s', (user_id,))
]
if present:
response: Final = gateway.request("POST", "/user/delete", {"user_ids": list[JsonValue](present)})
assert response.status_code == 200, response.text
def _member_delete(gateway: Gateway, team_id: str, user_id: str) -> int:
return gateway.request("POST", "/team/member_delete", {"team_id": team_id, "user_id": user_id}).status_code
def _owned_email_member(gateway: Gateway, scenario: Scenario) -> str:
user_id: Final = f"integration_{uuid.uuid4().hex}@example.com"
scenario.cleanups.callback(_delete_users_if_present, gateway, [user_id])
return user_id
def test_concurrent_team_creation_records_the_named_member(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
user: Final = scenario.user()
with ThreadPoolExecutor(max_workers=10) as pool:
teams: Final = list(
pool.map(lambda _: scenario.team(members_with_roles=[{"role": "user", "user_id": user}]), range(10))
)
assert len(set(teams)) == 10
for team in teams:
assert user in _member_ids(gateway, team)
def test_team_info_serves_admin_and_team_keys_and_denies_other_keys(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
assert _team_info(gateway, team)["team_id"] == team
team_key: Final = scenario.key(team_id=team)
as_team: Final = gateway.request("GET", "/team/info", params={"team_id": team}, key=team_key)
assert as_team.status_code == 200, as_team.text
assert object_value(object_value(as_team.json())["team_info"])["team_id"] == team
outsider: Final = scenario.key()
denied: Final = gateway.request("GET", "/team/info", params={"team_id": team}, key=outsider)
assert denied.status_code in {401, 403}, denied.text
def test_team_alias_update_keeps_every_other_team_field(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
admin: Final = scenario.user()
created: Final = gateway.post(
"/team/new",
{
"team_alias": f"integration-{uuid.uuid4().hex}",
"members_with_roles": [{"role": "admin", "user_id": admin}],
},
)
team: Final = string_value(created["team_id"])
scenario.cleanups.callback(scenario.delete_team, team)
initial_size: Final = len(_member_ids(gateway, team))
new_members: Final = [_owned_email_member(gateway, scenario) for _ in range(10)]
gateway.post(
"/team/member_add",
{"team_id": team, "member": [{"role": "user", "user_id": member} for member in new_members]},
)
members: Final = _member_ids(gateway, team)
assert len(members) == initial_size + 10
assert set(new_members) <= set(members)
new_alias: Final = f"integration-{uuid.uuid4().hex}"
updated: Final = object_value(gateway.post("/team/update", {"team_id": team, "team_alias": new_alias})["data"])
assert updated["team_alias"] == new_alias
updated_members: Final = updated["members_with_roles"]
assert isinstance(updated_members, list)
assert len(updated_members) == len(members)
compared: Final = (set(created) | set(updated)) - UNCHANGED_BY_ALIAS_UPDATE_SKIP
for field in compared:
assert updated.get(field) == created.get(field), field
assert compared
def test_member_added_by_email_sees_the_team_in_user_info(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
email: Final = f"integration-{uuid.uuid4().hex}@example.com"
user: Final = scenario.user(user_email=email)
team: Final = scenario.team()
gateway.post("/team/member_add", {"team_id": team, "member": {"role": "user", "user_email": email}})
assert team in _user_team_ids(gateway, user)
def test_team_delete_detaches_members_and_hides_the_team(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
first: Final = scenario.user()
second: Final = scenario.user()
created: Final = gateway.post(
"/team/new",
{
"team_alias": f"integration-{uuid.uuid4().hex}",
"members_with_roles": [{"role": "admin", "user_id": first}, {"role": "user", "user_id": second}],
},
)
team: Final = string_value(created["team_id"])
scenario.cleanups.callback(gateway.request, "POST", "/team/delete", {"team_ids": [team]})
team_key: Final = string_value(gateway.post("/key/generate", {"team_id": team, "user_id": second})["key"])
scenario.cleanups.callback(delete_key_if_present, gateway, team_key)
assert _user_team_ids(gateway, second) == [team]
gateway.post("/team/delete", {"team_ids": [team]})
assert _user_team_ids(gateway, second) == []
missing: Final = gateway.request("GET", "/team/info", params={"team_id": team})
assert missing.status_code == 404, missing.text
@pytest.mark.parametrize("dimension", ["user_id", "user_email"])
def test_member_delete_removes_the_member_by_id_or_email(gateway: Gateway, dimension: MemberDimension) -> None:
with gateway.scenario() as scenario:
email: Final = f"integration-{uuid.uuid4().hex}@example.com"
user: Final = scenario.user(user_email=email)
selector: Final[dict[str, JsonValue]] = {"user_id": user} if dimension == "user_id" else {"user_email": email}
team: Final = scenario.team(members_with_roles=[{"role": "user", **selector}])
assert user in _member_ids(gateway, team)
deleted: Final = gateway.request("POST", "/team/member_delete", {"team_id": team, **selector})
assert deleted.status_code == 200, deleted.text
assert user not in _member_ids(gateway, team)
assert team not in _user_team_ids(gateway, user)
def test_member_add_with_email_shaped_user_id_grows_team_by_one(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
initial: Final = _member_ids(gateway, team)
member: Final = _owned_email_member(gateway, scenario)
gateway.post("/team/member_add", {"team_id": team, "member": {"role": "user", "user_id": member}})
after: Final = _member_ids(gateway, team)
assert member in after
assert len(after) == len(initial) + 1
def test_member_delete_removes_once_and_rejects_repeats(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
member: Final = f"integration-{uuid.uuid4().hex}"
scenario.cleanups.callback(_delete_users_if_present, gateway, [member])
gateway.post("/team/member_add", {"team_id": team, "member": {"role": "user", "user_id": member}})
before: Final = _member_ids(gateway, team)
assert member in before
assert _member_delete(gateway, team, member) == 200
assert [_member_delete(gateway, team, member) for _ in range(4)] == [400, 400, 400, 400]
after: Final = _member_ids(gateway, team)
assert member not in after
assert len(after) == len(before) - 1
def test_member_delete_of_a_non_member_is_a_bad_request(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team()
outsider: Final = f"integration-{uuid.uuid4().hex}"
assert outsider not in _member_ids(gateway, team)
assert _member_delete(gateway, team, outsider) == 400

View file

@ -0,0 +1,34 @@
import uuid
from concurrent.futures import ThreadPoolExecutor
from typing import Final
from tests.integration._support.client import Gateway, object_value
from tests.integration._support.database import read_rows
def test_concurrent_user_creation_persists_models_and_aliases(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
alias: Final = f"alias-{uuid.uuid4().hex}"
with ThreadPoolExecutor(max_workers=10) as pool:
users: Final = list(pool.map(lambda _: scenario.user(models=[model], aliases={alias: model}), range(10)))
assert len(set(users)) == 10
rows: Final = [
read_rows('SELECT models FROM "LiteLLM_UserTable" WHERE user_id = %s', (user,)) for user in users
]
assert rows == [[{"models": [model]}]] * len(users)
def test_user_info_serves_admin_and_self_and_denies_other_users(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
user: Final = scenario.user(user_role="internal_user")
own_key: Final = scenario.key(user_id=user)
other: Final = scenario.user(user_role="internal_user")
other_key: Final = scenario.key(user_id=other)
admin: Final = object_value(gateway.get("/user/info", {"user_id": user})["user_info"])
assert admin["user_id"] == user
own: Final = gateway.request("GET", "/user/info", params={"user_id": user}, key=own_key)
assert own.status_code == 200, own.text
assert object_value(object_value(own.json())["user_info"])["user_id"] == user
denied: Final = gateway.request("GET", "/user/info", params={"user_id": user}, key=other_key)
assert denied.status_code == 403, denied.text

View file

@ -0,0 +1,291 @@
import threading
import time
import uuid
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from typing import Final, Literal
import httpx
import pytest
from pydantic import JsonValue, TypeAdapter
from tests.integration._support.database import read_rows
from tests.integration._support.client import (
JSON_OBJECT,
Gateway,
Scenario,
eventually,
object_value,
string_value,
)
from tests.integration._support.openai_wire import chat_reply
from tests.integration._support.wire import Reply, Request, wire_server
UNREACHABLE_API_BASE: Final = "http://127.0.0.1:9/v1"
END_USER_REQUESTS: Final = 10
CONCURRENT_REQUESTS: Final = 25
DISTRIBUTION_REQUESTS: Final = 20
SLOW_UPSTREAM_SECONDS: Final = 3
HELD_REPLY_SECONDS: Final = 30
END_USER_ROW_SECONDS: Final = 70
CUSTOM_FALLBACK_TEXT: Final = "custom fallback prompt"
Caller = Literal["virtual-key", "master-key"]
JSON_OBJECTS: Final = TypeAdapter(list[dict[str, JsonValue]])
def _messages(text: str = "integration control") -> list[JsonValue]:
return [{"role": "user", "content": text}]
def _chat(
gateway: Gateway,
body: dict[str, JsonValue],
*,
key: str | None = None,
headers: dict[str, str] | None = None,
) -> httpx.Response:
return gateway.request("POST", "/v1/chat/completions", {"messages": _messages(), **body}, key=key, headers=headers)
def _body(response: httpx.Response) -> dict[str, JsonValue]:
return JSON_OBJECT.validate_json(response.content)
def _content(response: httpx.Response) -> str:
choices: Final = _body(response)["choices"]
assert isinstance(choices, list) and choices, response.text
return string_value(object_value(object_value(choices[0])["message"])["content"])
def _scripted_model(gateway: Gateway, scenario: Scenario, statuses: list[int], num_retries: int) -> str:
upstream_model: Final = f"integration-{uuid.uuid4().hex}"
script_url: Final = f"{gateway.upstream_url}/__scripts/{upstream_model}"
configured: Final = httpx.post(script_url, json={"statuses": statuses})
assert configured.status_code == 200, configured.text
scenario.cleanups.callback(httpx.delete, script_url)
return scenario.model(model=f"openai/{upstream_model}", num_retries=num_retries)
@dataclass(frozen=True, slots=True)
class UniqueModel:
name: str
upstream: str
deployment_id: str
def _unique_model(scenario: Scenario) -> UniqueModel:
upstream_model: Final = f"integration-{uuid.uuid4().hex}"
deployment_id: Final = f"integration-{uuid.uuid4().hex}"
model_info: Final[dict[str, JsonValue]] = {"id": deployment_id}
name: Final = scenario.model(model=f"openai/{upstream_model}", model_info=model_info)
return UniqueModel(name, upstream_model, deployment_id)
def _slow_reply(_: Request) -> Reply:
time.sleep(SLOW_UPSTREAM_SECONDS)
return chat_reply("chatcmpl-slow", "gpt-4o-mini", "late", stream=False)
def _fallback_reply(_: Request) -> Reply:
return chat_reply("chatcmpl-fallback", "gpt-4o-mini", "served by fallback", stream=False)
def _held_reply(release: threading.Event) -> Callable[[Request], Reply]:
def respond(_: Request) -> Reply:
assert release.wait(timeout=HELD_REPLY_SECONDS)
return chat_reply("chatcmpl-held", "gpt-4o-mini", "held", stream=False)
return respond
def _delete_auto_created_end_user(gateway: Gateway, user_id: str) -> None:
eventually(
lambda: read_rows('SELECT user_id FROM "LiteLLM_EndUserTable" WHERE user_id = %s', (user_id,)),
lambda rows: len(rows) == 1,
seconds=END_USER_ROW_SECONDS,
)
gateway.post("/end_user/delete", {"user_ids": [user_id]})
assert read_rows('SELECT user_id FROM "LiteLLM_EndUserTable" WHERE user_id = %s', (user_id,)) == []
def _status(gateway: Gateway, model: str) -> int:
return _chat(gateway, {"model": model}).status_code
def _served_model_id(gateway: Gateway, model: str) -> str:
response: Final = _chat(gateway, {"model": model})
assert response.status_code == 200, response.text
return response.headers["x-litellm-model-id"]
def _concurrent_round(gateway: Gateway, bad: str, good: str) -> tuple[list[int], list[int]]:
with ThreadPoolExecutor(max_workers=CONCURRENT_REQUESTS * 2) as pool:
bad_calls: Final = [pool.submit(_status, gateway, bad) for _ in range(CONCURRENT_REQUESTS)]
good_calls: Final = [pool.submit(_status, gateway, good) for _ in range(CONCURRENT_REQUESTS)]
return [call.result() for call in bad_calls], [call.result() for call in good_calls]
def _deployment(gateway: Gateway, scenario: Scenario, model_name: str, model: str = "openai/gpt-4o-mini") -> str:
created: Final = gateway.post(
"/model/new",
{
"model_name": model_name,
"litellm_params": {
"model": model,
"api_key": "integration-provider-key",
"api_base": f"{gateway.upstream_url}/v1",
},
},
)
deployment_id: Final = string_value(object_value(created["model_info"])["id"])
scenario.cleanups.callback(scenario.delete_model, deployment_id)
return deployment_id
@pytest.mark.parametrize("caller", ["virtual-key", "master-key"])
def test_end_user_budget_tpm_limit_rate_limits_their_requests(gateway: Gateway, caller: Caller) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
budget: Final = scenario.budget(tpm_limit=2)
end_user: Final = f"integration-{uuid.uuid4().hex}"
gateway.post("/end_user/new", {"user_id": end_user, "budget_id": budget})
scenario.cleanups.callback(gateway.post, "/end_user/delete", {"user_ids": [end_user]})
key: Final = scenario.key(models=[model]) if caller == "virtual-key" else gateway.key
control_user: Final = f"integration-{uuid.uuid4().hex}"
scenario.cleanups.callback(_delete_auto_created_end_user, gateway, control_user)
control: Final = _chat(gateway, {"model": model, "user": control_user}, key=key)
assert control.status_code == 200, control.text
statuses: Final = [
_chat(gateway, {"model": model, "user": end_user}, key=key).status_code for _ in range(END_USER_REQUESTS)
]
assert statuses.count(200) < 5, statuses
assert set(statuses) <= {200, 429}, statuses
def test_client_fallbacks_reach_an_allowed_model_and_name_a_denied_one(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
primary: Final = scenario.model(api_base=UNREACHABLE_API_BASE)
fallback: Final = _unique_model(scenario)
body: Final[dict[str, JsonValue]] = {"model": primary, "fallbacks": [fallback.name]}
served: Final = _chat(gateway, body, key=scenario.key(models=[primary, fallback.name]))
assert served.status_code == 200, served.text
assert _body(served)["model"] == fallback.upstream
assert served.headers["x-litellm-model-id"] == fallback.deployment_id
assert _content(served)
denied: Final = _chat(gateway, body, key=scenario.key(models=[primary]))
assert denied.status_code == 403, denied.text
assert fallback.name in denied.text
def test_client_fallback_with_custom_messages_sends_them_to_the_fallback(gateway: Gateway) -> None:
custom: Final = _messages(CUSTOM_FALLBACK_TEXT)
with gateway.scenario() as scenario, wire_server(_fallback_reply) as wire:
primary: Final = scenario.model(api_base=UNREACHABLE_API_BASE)
fallback: Final = scenario.model(api_base=wire.url)
body: Final[dict[str, JsonValue]] = {
"model": primary,
"fallbacks": [{"model": fallback, "messages": custom}],
}
served: Final = _chat(gateway, body, key=scenario.key(models=[primary, fallback]))
assert served.status_code == 200, served.text
assert _content(served) == "served by fallback"
forwarded: Final = [JSON_OBJECT.validate_json(request.body)["messages"] for request in wire.drain()]
assert forwarded == [custom]
denied: Final = _chat(gateway, body, key=scenario.key(models=[primary]))
assert denied.status_code == 403, denied.text
assert fallback in denied.text
assert wire.drain() == ()
def test_rate_limited_deployment_is_retried_and_reports_retry_counts(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = _scripted_model(gateway, scenario, [429, 200], 50)
response: Final = _chat(gateway, {"model": model})
assert response.status_code == 200, response.text
assert response.headers["x-litellm-attempted-retries"] == "1"
assert response.headers["x-litellm-max-retries"] == "50"
def test_request_fallbacks_reroute_after_a_connection_failure(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
primary: Final = scenario.model(api_base=UNREACHABLE_API_BASE)
fallback: Final = _unique_model(scenario)
response: Final = _chat(gateway, {"model": primary, "fallbacks": [fallback.name]})
assert response.status_code == 200, response.text
assert _body(response)["model"] == fallback.upstream
assert response.headers["x-litellm-model-id"] == fallback.deployment_id
assert response.headers["x-litellm-attempted-fallbacks"] == "1"
def test_model_level_timeout_is_reported_on_a_timed_out_request(gateway: Gateway) -> None:
with gateway.scenario() as scenario, wire_server(_slow_reply) as wire:
response: Final = _chat(gateway, {"model": scenario.model(api_base=wire.url, timeout=1)})
assert response.status_code == 408, response.text
assert response.headers["x-litellm-timeout"] == "1.0"
def test_request_timeout_header_overrides_the_model_timeout(gateway: Gateway) -> None:
with gateway.scenario() as scenario, wire_server(_slow_reply) as wire:
response: Final = _chat(
gateway,
{"model": scenario.model(api_base=wire.url, timeout=1)},
headers={"x-litellm-timeout": "0.001"},
)
assert response.status_code == 408, response.text
assert response.headers["x-litellm-timeout"] == "0.001"
def test_failing_model_traffic_does_not_starve_concurrent_good_requests(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
bad: Final = scenario.model(api_base=UNREACHABLE_API_BASE)
good: Final = scenario.model()
for _ in range(2):
bad_calls, good_calls = _concurrent_round(gateway, bad, good)
assert good_calls == [200] * CONCURRENT_REQUESTS
assert 200 not in bad_calls
def test_rpm_limited_deployment_rejects_a_second_call_while_the_first_is_in_flight(gateway: Gateway) -> None:
release: Final = threading.Event()
with gateway.scenario() as scenario, wire_server(_held_reply(release)) as wire, ThreadPoolExecutor(1) as pool:
model: Final = scenario.model(api_base=wire.url, rpm=1)
first: Final = pool.submit(_chat, gateway, {"model": model})
try:
eventually(wire.received.qsize, lambda count: count == 1)
second: Final = _chat(gateway, {"model": model})
assert second.status_code == 429, second.text
assert wire.received.qsize() == 1
finally:
release.set()
assert first.result().status_code == 200
assert len(wire.drain()) == 1
def test_model_group_with_two_deployments_serves_from_both(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
first_id: Final = f"integration-{uuid.uuid4().hex}"
model: Final = scenario.model(model_info={"id": first_id})
second_id: Final = _deployment(gateway, scenario, model)
served_by: Final = {_served_model_id(gateway, model) for _ in range(DISTRIBUTION_REQUESTS)}
assert served_by == {first_id, second_id}
def test_unlisted_provider_model_resolves_through_a_wildcard_deployment(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
prefix: Final = f"integration{uuid.uuid4().hex}"
_deployment(gateway, scenario, f"{prefix}/*", "openai/*")
response: Final = _chat(gateway, {"model": f"{prefix}/gpt-4o-mini"})
assert response.status_code == 200, response.text
assert _content(response)
def test_comma_separated_models_fan_out_to_one_response_per_model(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
first: Final = _unique_model(scenario)
second: Final = _unique_model(scenario)
response: Final = _chat(gateway, {"model": f"{first.name},{second.name}"})
assert response.status_code == 200, response.text
replies: Final = JSON_OBJECTS.validate_json(response.content)
assert len(replies) == 2, replies
assert {string_value(reply["model"]) for reply in replies} == {first.upstream, second.upstream}

View file

@ -0,0 +1,102 @@
from contextlib import ExitStack
from itertools import chain
from typing import Final
from pydantic import JsonValue
from tests.integration._support.client import (
JSON_OBJECT,
Gateway,
Scenario,
delete_key_if_present,
eventually,
object_value,
string_value,
)
from tests.integration._support.database import read_rows
MEMBER_BUDGET: Final = 0.0000001
SPEND_LANDING_SECONDS: Final = 70
def test_spend_log_of_an_org_team_key_records_org_and_team(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
organization: Final = scenario.organization(models=[model])
team: Final = scenario.team(organization_id=organization, models=[model])
key: Final = scenario.key(team_id=team, models=[model])
request_id: Final = string_value(gateway.chat(model, key=key)["id"])
rows: Final = eventually(
lambda: read_rows(
'SELECT team_id, metadata::text AS metadata FROM "LiteLLM_SpendLogs" WHERE request_id = %s',
(request_id,),
),
lambda found: len(found) == 1,
seconds=SPEND_LANDING_SECONDS,
)
metadata: Final = JSON_OBJECT.validate_json(string_value(rows[0]["metadata"]))
assert metadata["user_api_key_org_id"] == organization
assert metadata["user_api_key_team_id"] == team
assert rows[0]["team_id"] == team
def _team_memberships(team: JsonValue, user_id: str) -> list[dict[str, JsonValue]]:
entries: Final = object_value(team).get("team_memberships") or []
assert isinstance(entries, list)
return [object_value(entry) for entry in entries if object_value(entry).get("user_id") == user_id]
def _membership(gateway: Gateway, user_id: str, team_id: str) -> dict[str, JsonValue]:
teams: Final = gateway.get("/user/info", {"user_id": user_id})["teams"]
assert isinstance(teams, list)
memberships: Final = list(
chain.from_iterable(
_team_memberships(team, user_id) for team in teams if object_value(team)["team_id"] == team_id
)
)
assert len(memberships) == 1, teams
return memberships[0]
def test_team_member_budget_blocks_the_member_after_spend_lands(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
team: Final = scenario.team()
created: Final = gateway.post(
"/user/new",
{"user_id": f"integration-{model}", "team_id": team, "models": [model], "max_budget": 10.0},
)
user: Final = string_value(created["user_id"])
key: Final = string_value(created["key"])
with ExitStack() as member_cleanups:
member_cleanups.callback(scenario.delete_user, user)
member_cleanups.callback(delete_key_if_present, gateway, key)
gateway.post("/team/member_update", {"team_id": team, "user_id": user, "max_budget_in_team": MEMBER_BUDGET})
membership: Final = _membership(gateway, user, team)
scenario.cleanups.callback(scenario.delete_budget, string_value(membership["budget_id"]))
scenario.cleanups.push(member_cleanups.pop_all())
assert object_value(membership["litellm_budget_table"])["max_budget"] == MEMBER_BUDGET
assert (
gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "first"}]},
key=key,
).status_code
== 200
)
eventually(
lambda: read_rows(
'SELECT spend FROM "LiteLLM_TeamMembership" WHERE user_id = %s AND team_id = %s', (user, team)
),
lambda rows: len(rows) == 1 and isinstance(spend := rows[0]["spend"], float) and spend >= MEMBER_BUDGET,
seconds=SPEND_LANDING_SECONDS,
)
blocked: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "second"}]},
key=key,
)
assert blocked.status_code != 200, blocked.text
assert "Budget has been exceeded" in blocked.text

View file

@ -63,4 +63,4 @@ tests/rust-python-harness/
- `shared/` contains reusable parity, tracing, reporting primitives, and unit-runner machinery
- Keep fixtures with their owning API and existing Python tests in their current locations
- Each strategy folder carries an `AGENTS.md` one-liner stating what it should be doing
- Run the harness's own checks with `uv run pytest -o consider_namespace_packages=true tests/rust-python-harness/shared tests/rust-python-harness/cli tests/rust-python-harness/strategies/trace_parity tests/rust-python-harness/strategies/unit_tests_parity tests/rust-python-harness/strategies/unit_tests_rust tests/test_rust_python_harness.py -q`
- Run the harness's own checks with `uv run pytest -o consider_namespace_packages=true tests/rust-python-harness/shared tests/rust-python-harness/cli tests/rust-python-harness/strategies/trace_parity tests/rust-python-harness/strategies/unit_tests_parity tests/rust-python-harness/strategies/unit_tests_rust -q`

View file

@ -464,3 +464,12 @@ def test_runner_interrupt_skips_the_completion_report(
assert exit_code == 130
assert "Rust <-> Python parity report" not in captured.out
assert captured.err == "Interrupted\n"
def test_strategy_subcommand_accepts_function_filter(capsys: pytest.CaptureFixture[str]) -> None:
exit_code: Final = main(["run", "unit_tests_rust", "--function", "messages"])
captured: Final = capsys.readouterr()
assert exit_code == 0
assert "- messages: not_implemented" in captured.out
assert "unit_tests_rust:messages: not_implemented" not in captured.out

View file

@ -0,0 +1,53 @@
from __future__ import annotations
from typing import Final
import pytest
from .models import CaseResult, Coverage, HarnessCase, RunStatus
from .strategy import ModuleCaseSpec, NotImplementedCaseSpec, SkippedCaseSpec
def _case(spec: ModuleCaseSpec | NotImplementedCaseSpec | SkippedCaseSpec) -> HarnessCase:
return HarnessCase(strategy_id="example", strategy_label="Example", sdk_function="messages", spec=spec)
def _runnable_case() -> HarnessCase:
return _case(ModuleCaseSpec(coverage=Coverage.COMPLETE, module="tests.example"))
def test_should_mark_not_implemented_and_skipped_cases_without_running() -> None:
not_implemented: Final = CaseResult(case=_case(NotImplementedCaseSpec(reason="No case is registered.")))
skipped: Final = CaseResult(case=_case(SkippedCaseSpec(reason="The surface does not apply.")))
not_implemented.set_initial_status()
skipped.set_initial_status()
assert not_implemented.status is RunStatus.NOT_IMPLEMENTED
assert skipped.status is RunStatus.SKIPPED
def test_should_finalize_a_fully_passing_case() -> None:
result: Final = CaseResult(case=_runnable_case())
result.set_initial_status()
result.collected.update({"one", "two"})
result.completed.update({"one", "two"})
result.passed = 2
result.finalize()
assert result.status is RunStatus.PASSED
def test_should_replace_a_pass_with_a_teardown_error() -> None:
result: Final = CaseResult(case=_runnable_case())
result.set_initial_status()
result.collected.add("one")
result.record("one", RunStatus.PASSED, 0.1)
result.record("one", RunStatus.ERROR, 0.2)
assert result.status is RunStatus.ERROR
assert result.passed == 0
assert result.errors == 1
assert result.duration == pytest.approx(0.3)

View file

@ -0,0 +1,26 @@
from __future__ import annotations
from typing import Final
from .models import Coverage, HarnessCase, HarnessRun, RunStatus
from .strategy import ModuleCaseSpec
from .ui import _format_duration, _summary
def test_should_format_developer_facing_run_context() -> None:
run: Final = HarnessRun.from_cases(
(
HarnessCase(
strategy_id="example",
strategy_label="Example",
sdk_function="messages",
spec=ModuleCaseSpec(coverage=Coverage.COMPLETE, module="tests.example"),
),
)
)
result: Final = next(iter(run.results.values()))
result.collected.add("tests/test_parity.py::test_one")
result.record("tests/test_parity.py::test_one", RunStatus.PASSED, 1.25)
assert _summary(run) == (1, 0, 0, 0)
assert _format_duration(1.25) == "1.2s"

View file

@ -0,0 +1,23 @@
from __future__ import annotations
from importlib import import_module
from typing import Final, cast
import pytest
from ..models import TraceSuite
@pytest.mark.parametrize(
"module",
[
"tests.rust-python-harness.strategies.trace_parity.sdk.messages.case",
"tests.rust-python-harness.strategies.trace_parity.sdk.chat_completions.case",
"tests.rust-python-harness.strategies.trace_parity.sdk.transcription.case",
],
)
def test_implemented_namespace_case_modules_remain_importable(module: str) -> None:
loaded: Final = import_module(module)
suite: Final = cast(object, getattr(loaded, "TRACE_SUITE"))
assert isinstance(suite, TraceSuite)
assert suite.scenarios

View file

@ -1,96 +0,0 @@
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
def test_anthropic_compaction_usage_calculation():
"""
Test that calculate_usage correctly sums tokens from the iterations array
as requested in Issue #27060.
"""
anthropic_config = AnthropicConfig()
# Mock usage object with compaction iterations
usage_object = {
"input_tokens": 100, # Top-level (excludes compaction)
"output_tokens": 50, # Top-level (excludes compaction)
"iterations": [
{
"iteration": 1,
"type": "compaction",
"input_tokens": 1000,
"output_tokens": 500,
},
{
"iteration": 2,
"type": "message",
"input_tokens": 100,
"output_tokens": 50,
},
],
}
usage = anthropic_config.calculate_usage(
usage_object=usage_object, reasoning_content=None
)
# Assertions
# Total prompt tokens should be 1000 + 100 = 1100
assert usage.prompt_tokens == 1100
# Total completion tokens should be 500 + 50 = 550
assert usage.completion_tokens == 550
# Total tokens should be 1650
assert usage.total_tokens == 1650
# Assert details
assert usage.prompt_tokens_details.text_tokens == 1100
# Assert iterations passthrough
assert usage.iterations is not None
assert len(usage.iterations) == 2
assert usage.iterations[0]["type"] == "compaction"
def test_anthropic_compaction_usage_with_iteration_cache():
"""
Test that calculate_usage correctly sums caching tokens FROM iterations.
This covers the specific case mentioned by JasonPan.
"""
anthropic_config = AnthropicConfig()
usage_object = {
"input_tokens": 100,
"output_tokens": 50,
"iterations": [
{
"type": "compaction",
"input_tokens": 500,
"output_tokens": 200,
"cache_creation_input_tokens": 50,
"cache_read_input_tokens": 17000,
},
{
"type": "message",
"input_tokens": 100,
"output_tokens": 50,
"cache_creation_input_tokens": 10,
"cache_read_input_tokens": 20,
},
],
}
usage = anthropic_config.calculate_usage(
usage_object=usage_object, reasoning_content=None
)
# input_tokens sum = 500 + 100 = 600
# cache_creation sum = 50 + 10 = 60
# cache_read sum = 17000 + 20 = 17020
# Total prompt tokens = 600 + 60 + 17020 = 17680
assert usage.prompt_tokens == 17680
assert usage.completion_tokens == 250
assert usage.prompt_tokens_details.cache_creation_tokens == 60
assert usage.prompt_tokens_details.cached_tokens == 17020
if __name__ == "__main__":
test_anthropic_compaction_usage_calculation()
test_anthropic_compaction_usage_with_iteration_cache()

View file

@ -1,102 +0,0 @@
import os
# What is this?
## Unit tests for the /budget/* endpoints
from litellm._uuid import uuid
from datetime import datetime, timezone
import aiohttp
import pytest
import pytest_asyncio
from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_timezone
def _parse_budget_api_datetime(value: str) -> datetime:
"""Parse ISO timestamps returned by the proxy JSON API."""
if value.endswith("Z"):
value = value[:-1] + "+00:00"
dt = datetime.fromisoformat(value)
if dt.tzinfo is None:
dt = dt.replace(tzinfo=timezone.utc)
return dt
async def delete_budget(session, budget_id):
url = "http://0.0.0.0:4000/budget/delete"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {"id": budget_id}
async with session.post(url, headers=headers, json=data) as response:
assert response.status == 200
print(f"Deleted Budget {budget_id}")
async def create_budget(session, data):
url = "http://0.0.0.0:4000/budget/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
async with session.post(url, headers=headers, json=data) as response:
assert response.status == 200
response_data = await response.json()
budget_id = response_data["budget_id"]
print(f"Created Budget {budget_id}")
return response_data
@pytest_asyncio.fixture
async def budget_setup():
"""
Fixture to create a budget for testing and clean it up afterward.
This fixture performs the following steps:
1. Opens an aiohttp ClientSession.
2. Generates a random budget_id and defines the budget data (duration: 1 day, max_budget: 0.02).
3. Calls create_budget to create the budget.
4. Yields the budget_response (a dict) for use in the test.
5. After the test completes, deletes the created budget by calling delete_budget.
Returns:
dict: The JSON response from create_budget, which includes the created budget's data.
"""
async with aiohttp.ClientSession() as session:
# Generate a unique budget_id and define the budget data.
budget_id = f"budget-{uuid.uuid4()}"
data = {"budget_id": budget_id, "budget_duration": "1d", "max_budget": 0.02}
budget_response = await create_budget(session, data)
# Yield the response so the test can use it.
yield budget_response
# After the test, delete the created budget to clean up.
await delete_budget(session, budget_id)
@pytest.mark.asyncio
async def test_create_budget_with_duration(budget_setup):
"""
Test creating a budget with a specified duration and verify that 'budget_reset_at'
matches the next standardized reset (see get_budget_reset_time / new_budget), not
necessarily created_at + wall-clock duration.
"""
assert (
budget_setup["budget_reset_at"] is not None
), "The budget_reset_at field should not be None"
created_at = _parse_budget_api_datetime(budget_setup["created_at"])
expected_reset_at = get_next_standardized_reset_time(
duration=budget_setup["budget_duration"],
current_time=created_at,
timezone_str=get_budget_reset_timezone(),
)
actual_reset_at = _parse_budget_api_datetime(budget_setup["budget_reset_at"])
tolerance_seconds = 3
time_difference = abs((actual_reset_at - expected_reset_at).total_seconds())
assert time_difference <= tolerance_seconds, (
f"Expected budget_reset_at to be within {tolerance_seconds} seconds of {expected_reset_at}, "
f"but the difference was {time_difference} seconds."
)

View file

@ -1,303 +0,0 @@
# What this tests ?
## Makes sure the number of callbacks on the proxy don't increase over time
## Num callbacks should be a fixed number at t=0 and t=10, t=20
"""
PROD TEST - DO NOT Delete this Test
"""
import pytest
import asyncio
import aiohttp
import os
import re
import dotenv
from collections import Counter
from dotenv import load_dotenv
load_dotenv()
# A *leak* is sustained, monotonic growth of one callback TYPE across the whole
# sampling window. A one-time bump that then plateaus is benign pollution from
# other tests sharing this proxy (this suite runs `pytest -n 4` against a single
# proxy container, so other workers legitimately add team/key-scoped callbacks
# while this test sleeps). We therefore sample N times and only flag a type
# whose normalized count never decreases, grows in >=2 distinct intervals, and
# nets >= LEAK_MIN_NET_GROWTH overall.
NUM_SAMPLES = 4
SAMPLE_INTERVAL_SECONDS = 20
LEAK_MIN_NET_GROWTH = 5
LEAK_MIN_GROWING_INTERVALS = 2
# A routing-strategy switch / alerting config is a *known, bounded, one-time*
# registration (CCI diagnostic 2026-05-16: total 85->95 on the first interval
# after switching to latency-based-routing, then flat at 95 for 2.5 min under
# load). We absorb that step by settling before the baseline sample, so only
# growth *after* the deliberate perturbation can count as a leak.
SETTLE_SECONDS = 30
# Strip instance-identity noise so N leaked instances of one class collapse to
# one rising counter instead of N opaque, unrelated-looking strings.
_ADDR_RE = re.compile(r" at 0x[0-9a-fA-F]+")
_OBJ_RE = re.compile(r"<([\w.]+) object")
def _normalize_callback(cb_str: str) -> str:
"""Reduce a callback's str() to a stable type key (drops 0x… addresses)."""
s = _ADDR_RE.sub("", cb_str)
m = _OBJ_RE.search(s)
if m:
return m.group(1).split(".")[-1]
# bound methods: "<bound method Cls.m of <... at 0x..>>" -> "Cls.m"
bm = re.search(r"bound method ([\w.]+)", s)
if bm:
return bm.group(1)
return s.strip()
def _summarize(all_litellm_callbacks) -> Counter:
return Counter(_normalize_callback(str(c)) for c in all_litellm_callbacks)
def _detect_leaks(samples):
"""
samples: list[Counter] taken in time order.
Returns {callback_type: [counts across samples]} for types that grew
monotonically (never decreased), in >=LEAK_MIN_GROWING_INTERVALS intervals,
and netted >=LEAK_MIN_NET_GROWTH overall — i.e. a real leak, not a one-shot
step from a parallel test.
"""
leaks = {}
all_types = set().union(*[set(s) for s in samples]) if samples else set()
for t in all_types:
series = [s.get(t, 0) for s in samples]
deltas = [b - a for a, b in zip(series, series[1:])]
net = series[-1] - series[0]
non_decreasing = all(d >= 0 for d in deltas)
growing_intervals = sum(1 for d in deltas if d > 0)
if (
non_decreasing
and net >= LEAK_MIN_NET_GROWTH
and growing_intervals >= LEAK_MIN_GROWING_INTERVALS
):
leaks[t] = series
return leaks
def _terminal_suspects(samples):
"""
Types whose net growth clears the threshold monotonically but is confined
to the *final* interval — `growing_intervals == 1` with that one growing
interval being the last. `_detect_leaks`' `>= 2` guard silently passes
these, so a real leak that accumulates entirely in the last sampled window
is indistinguishable from a one-time terminal step *without one more
sample*. Returns the set of such types so the caller can re-confirm.
"""
suspects = set()
all_types = set().union(*[set(s) for s in samples]) if samples else set()
for t in all_types:
series = [s.get(t, 0) for s in samples]
deltas = [b - a for a, b in zip(series, series[1:])]
if not deltas:
continue
net = series[-1] - series[0]
non_decreasing = all(d >= 0 for d in deltas)
growing = [i for i, d in enumerate(deltas) if d > 0]
if (
non_decreasing
and net >= LEAK_MIN_NET_GROWTH
and growing == [len(deltas) - 1]
):
suspects.add(t)
return suspects
async def _detect_leaks_confirmed(session, samples):
"""
`_detect_leaks`, plus a single confirmation sample when growth is confined
to the final interval (see `_terminal_suspects`). A genuine ongoing leak
keeps climbing -> now grows in >= 2 intervals -> flagged; a one-time
terminal registration plateaus -> still 1 growing interval -> ignored.
Returns `(leaks, samples)` (samples may have one extra entry appended).
"""
leaks = _detect_leaks(samples)
if not leaks and _terminal_suspects(samples):
await asyncio.sleep(SAMPLE_INTERVAL_SECONDS)
_, _, all_cb = await get_active_callbacks(session=session)
samples = samples + [_summarize(all_cb)]
leaks = _detect_leaks(samples)
return leaks, samples
def _format_report(samples, leaks) -> str:
lines = ["Callback count per type across samples (time order):"]
all_types = sorted(set().union(*[set(s) for s in samples]))
for t in all_types:
series = [s.get(t, 0) for s in samples]
marker = " <-- LEAK" if t in leaks else ""
lines.append(f" {t}: {series}{marker}")
totals = [sum(s.values()) for s in samples]
lines.append(f"TOTAL callbacks per sample: {totals}")
if leaks:
lines.append(
"Leaking callback types (sustained monotonic growth): "
+ ", ".join(sorted(leaks))
)
return "\n".join(lines)
async def _sample_callbacks(session, num_samples, interval):
"""Take `num_samples` callback snapshots `interval`s apart."""
samples = []
alerts = []
for i in range(num_samples):
if i > 0:
await asyncio.sleep(interval)
num_cb, num_alert, all_cb = await get_active_callbacks(session=session)
samples.append(_summarize(all_cb))
alerts.append(num_alert)
return samples, alerts
async def config_update(session, routing_strategy=None):
url = "http://0.0.0.0:4000/config/update"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
print("routing_strategy: ", routing_strategy)
data = {
"router_settings": {
"routing_strategy": routing_strategy,
},
"general_settings": {
"alert_to_webhook_url": {"llm_exceptions": "example-slack-webhook-url"},
"alert_types": ["llm_exceptions", "db_exceptions"],
},
}
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_active_callbacks(session):
url = "http://0.0.0.0:4000/active/callbacks"
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
}
async with session.get(url, headers=headers) as response:
status = response.status
response_text = await response.text()
print("response from /active/callbacks")
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
_json_response = await response.json()
_num_callbacks = _json_response["num_callbacks"]
_num_alerts = _json_response["num_alerting"]
all_litellm_callbacks = _json_response["all_litellm_callbacks"]
print("current number of callbacks: ", _num_callbacks)
print("current number of alerts: ", _num_alerts)
return _num_callbacks, _num_alerts, all_litellm_callbacks
async def get_current_routing_strategy(session):
url = "http://0.0.0.0:4000/get/config/callbacks"
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
}
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}")
_json_response = await response.json()
print("JSON response: ", _json_response)
router_settings = _json_response["router_settings"]
print("Router settings: ", router_settings)
routing_strategy = router_settings["routing_strategy"]
return routing_strategy
@pytest.mark.asyncio
@pytest.mark.order1
@pytest.mark.flaky(reruns=2, reruns_delay=5)
async def test_check_num_callbacks():
"""
PROD invariant: no callback TYPE should grow without bound over time.
This suite runs `pytest -n 4` against one shared proxy, so the raw count is
noisy — other workers legitimately add team/key-scoped callbacks that then
plateau. We settle first, then sample several times, and only fail on
*sustained, monotonic* per-type growth (a genuine leak), naming the type.
"""
async with aiohttp.ClientSession() as session:
# Absorb proxy warmup / in-flight parallel registration before baseline.
await asyncio.sleep(SETTLE_SECONDS)
samples, _ = await _sample_callbacks(
session, NUM_SAMPLES, SAMPLE_INTERVAL_SECONDS
)
assert sum(samples[0].values()) > 0, "expected some callbacks registered"
leaks, samples = await _detect_leaks_confirmed(session, samples)
report = _format_report(samples, leaks)
print(report)
assert not leaks, f"Callback leak detected.\n{report}"
@pytest.mark.asyncio
@pytest.mark.order2
@pytest.mark.flaky(reruns=2, reruns_delay=5)
async def test_check_num_callbacks_on_lowest_latency():
"""
Same PROD invariant as test_check_num_callbacks, but after switching the
router to latency-based-routing. That switch is a *known, bounded* one-time
registration (it adds the latency strategy handler + Slack alerting); we
settle past it before baselining so only post-switch growth counts as a
leak. Also asserts the alerting count is stable.
"""
async with aiohttp.ClientSession() as session:
await asyncio.sleep(30)
original_routing_strategy = await get_current_routing_strategy(session=session)
await config_update(session=session, routing_strategy="latency-based-routing")
try:
# Absorb the deliberate one-time config/update registration step.
await asyncio.sleep(SETTLE_SECONDS)
samples, alerts = await _sample_callbacks(
session, NUM_SAMPLES, SAMPLE_INTERVAL_SECONDS
)
leaks, samples = await _detect_leaks_confirmed(session, samples)
report = _format_report(samples, leaks)
print(report)
assert not leaks, f"Callback leak detected.\n{report}"
assert (
len(set(alerts)) == 1
), f"alerting count changed across samples: {alerts}"
finally:
await config_update(
session=session, routing_strategy=original_routing_strategy
)

View file

@ -1,55 +0,0 @@
import importlib
import os
from unittest.mock import MagicMock, patch
import litellm.litellm_core_utils.default_encoding as default_encoding
def _reload_default_encoding(monkeypatch, **env_overrides):
"""
Helper to reload default_encoding with a clean TIKTOKEN_CACHE_DIR and
specific environment overrides.
"""
monkeypatch.delenv("TIKTOKEN_CACHE_DIR", raising=False)
monkeypatch.delenv("CUSTOM_TIKTOKEN_CACHE_DIR", raising=False)
for key, value in env_overrides.items():
monkeypatch.setenv(key, value)
importlib.reload(default_encoding)
def test_default_encoding_uses_bundled_tokenizers_by_default(monkeypatch):
"""
TIKTOKEN_CACHE_DIR should point at the bundled tokenizers directory
when no CUSTOM_TIKTOKEN_CACHE_DIR is set, even in non-root environments.
"""
_reload_default_encoding(monkeypatch, LITELLM_NON_ROOT="true")
assert "TIKTOKEN_CACHE_DIR" in os.environ
cache_dir = os.environ["TIKTOKEN_CACHE_DIR"]
assert "tokenizers" in cache_dir
def test_custom_tiktoken_cache_dir_override(monkeypatch, tmp_path):
"""
CUSTOM_TIKTOKEN_CACHE_DIR must override the default bundled directory
and the directory should be created if it does not exist.
Reload with an empty custom dir would otherwise trigger tiktoken to
download the vocab; we patch get_encoding so the test is offline-safe
and does not depend on tiktoken's in-memory cache state.
"""
custom_dir = tmp_path / "tiktoken_cache"
with patch(
"litellm.litellm_core_utils.default_encoding.tiktoken.get_encoding",
return_value=MagicMock(),
):
_reload_default_encoding(monkeypatch, CUSTOM_TIKTOKEN_CACHE_DIR=str(custom_dir))
cache_dir = os.environ.get("TIKTOKEN_CACHE_DIR")
assert cache_dir == str(custom_dir)
assert os.path.isdir(cache_dir)
# Restore module to a clean state so default_encoding.encoding is a real
# tiktoken Encoding, not the MagicMock, for any test that runs after this.
monkeypatch.delenv("TIKTOKEN_CACHE_DIR", raising=False)
monkeypatch.delenv("CUSTOM_TIKTOKEN_CACHE_DIR", raising=False)
importlib.reload(default_encoding)

View file

@ -1,201 +0,0 @@
import os
# What is this?
## Unit tests for the /end_users/* endpoints
import pytest
import asyncio
import aiohttp
import time
from litellm._uuid import uuid
from openai import AsyncOpenAI
from typing import Optional
"""
- `/end_user/new`
- `/end_user/info`
"""
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,
calling_key=os.environ["LITELLM_MASTER_KEY"],
):
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,
}
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()
async def new_end_user(
session,
i,
user_id=str(uuid.uuid4()),
model_region=None,
default_model=None,
budget_id=None,
):
url = "http://0.0.0.0:4000/end_user/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {
"user_id": user_id,
"allowed_model_region": model_region,
"default_model": default_model,
}
if budget_id is not None:
data["budget_id"] = budget_id
print("end user data: {}".format(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()
async def new_budget(session, i, budget_id=None):
url = "http://0.0.0.0:4000/budget/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {
"budget_id": budget_id,
"tpm_limit": 2,
}
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()
@pytest.mark.asyncio
async def test_enduser_tpm_limits_non_master_key():
"""
1. budget_id = Create Budget with tpm_limit = 10
2. create end_user with budget_id
3. Make /chat/completions calls
4. Sleep 1 second
4. Make /chat/completions call -> expect this to fail because rate limit hit
"""
async with aiohttp.ClientSession() as session:
# create a budget with budget_id = "free-tier"
budget_id = f"free-tier-{uuid.uuid4()}"
await new_budget(session, 0, budget_id=budget_id)
await asyncio.sleep(2)
end_user_id = str(uuid.uuid4())
await new_end_user(
session=session, i=0, user_id=end_user_id, budget_id=budget_id
)
## MAKE CALL ##
key_gen = await generate_key(session=session, i=0, models=[])
key = key_gen["key"]
# chat completion 1
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000", max_retries=0)
# chat completion 2
passed = 0
for _ in range(10):
try:
result = await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": "Hey!"}],
user=end_user_id,
)
passed += 1
except Exception:
pass
print("Passed requests=", passed)
assert (
passed < 5
), f"Sent 10 requests and end-user has tpm_limit of 2. Number requests passed: {passed}. Expected less than 5 to pass"
@pytest.mark.asyncio
async def test_enduser_tpm_limits_with_master_key():
"""
1. budget_id = Create Budget with tpm_limit = 10
2. create end_user with budget_id
3. Make /chat/completions calls
4. Sleep 1 second
4. Make /chat/completions call -> expect this to fail because rate limit hit
"""
async with aiohttp.ClientSession() as session:
# create a budget with budget_id = "free-tier"
budget_id = f"free-tier-{uuid.uuid4()}"
await new_budget(session, 0, budget_id=budget_id)
end_user_id = str(uuid.uuid4())
await new_end_user(
session=session, i=0, user_id=end_user_id, budget_id=budget_id
)
# chat completion 1
client = AsyncOpenAI(
api_key=os.environ["LITELLM_MASTER_KEY"], base_url="http://0.0.0.0:4000", max_retries=0
)
# chat completion 2
passed = 0
for _ in range(10):
try:
result = await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": "Hey!"}],
user=end_user_id,
)
passed += 1
except Exception:
pass
print("Passed requests=", passed)
assert (
passed < 5
), f"Sent 10 requests and end-user has tpm_limit of 2. Number requests passed: {passed}. Expected less than 5 to pass"

View file

@ -1,325 +0,0 @@
import os
from typing import Final
# What is this?
## This tests if the proxy fallbacks work as expected
import pytest
import asyncio
import aiohttp
from tests.large_text import text
import time
from typing import Optional
from openai import AsyncOpenAI, PermissionDeniedError
PROXY_BASE_URL: Final = os.environ.get("LITELLM_PROXY_BASE_URL", "http://0.0.0.0:4000")
async def generate_key(
session,
i,
models: list,
calling_key=os.environ["LITELLM_MASTER_KEY"],
):
url: Final = f"{PROXY_BASE_URL}/key/generate"
headers = {
"Authorization": f"Bearer {calling_key}",
"Content-Type": "application/json",
}
data = {
"models": models,
}
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()
async def chat_completion(
session,
key: str,
model: str,
messages: list,
return_headers: bool = False,
extra_headers: Optional[dict] = None,
**kwargs,
):
url: Final = f"{PROXY_BASE_URL}/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
if extra_headers is not None:
headers.update(extra_headers)
data = {"model": model, "messages": messages, **kwargs}
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:
if return_headers:
return None, response.headers
else:
raise Exception(f"Request did not return a 200 status code: {status}")
if return_headers:
return await response.json(), response.headers
else:
return await response.json()
@pytest.mark.parametrize("has_access", [True, False])
@pytest.mark.asyncio
async def test_chat_completion_client_fallbacks(has_access: bool) -> None:
models: Final = ["gpt-3.5-turbo", "gpt-6-luna"] if has_access else ["gpt-3.5-turbo"]
async with aiohttp.ClientSession() as session:
generated_key: Final = await generate_key(session=session, i=0, models=models)
async with AsyncOpenAI(api_key=generated_key["key"], base_url=PROXY_BASE_URL, max_retries=0) as client:
request: Final = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Who was Alexander?"}],
"max_tokens": 32,
"temperature": 0,
"extra_body": {
"mock_testing_fallbacks": True,
"fallbacks": ["gpt-6-luna"],
},
}
if not has_access:
with pytest.raises(PermissionDeniedError) as denied:
await client.chat.completions.create(**request)
assert denied.value.status_code == 403
assert "gpt-6-luna" in str(denied.value)
return
response: Final = await client.chat.completions.create(**request)
assert response.model == "gpt-6-luna"
assert response.choices[0].message.content
@pytest.mark.asyncio
async def test_chat_completion_with_retries():
"""
make chat completion call with prompt > context window. expect it to work with fallback
"""
async with aiohttp.ClientSession() as session:
model = "fake-openai-endpoint-4"
messages = [
{"role": "system", "content": text},
{"role": "user", "content": "Who was Alexander?"},
]
response, headers = await chat_completion(
session=session,
key=os.environ["LITELLM_MASTER_KEY"],
model=model,
messages=messages,
mock_testing_rate_limit_error=True,
return_headers=True,
)
print(f"headers: {headers}")
assert headers["x-litellm-attempted-retries"] == "1"
assert headers["x-litellm-max-retries"] == "50"
@pytest.mark.asyncio
async def test_chat_completion_with_fallbacks():
"""
make chat completion call with prompt > context window. expect it to work with fallback
"""
async with aiohttp.ClientSession() as session:
model = "badly-configured-openai-endpoint"
messages = [
{"role": "system", "content": text},
{"role": "user", "content": "Who was Alexander?"},
]
response, headers = await chat_completion(
session=session,
key=os.environ["LITELLM_MASTER_KEY"],
model=model,
messages=messages,
fallbacks=["fake-openai-endpoint-5"],
return_headers=True,
)
print(f"headers: {headers}")
assert headers["x-litellm-attempted-fallbacks"] == "1"
@pytest.mark.asyncio
async def test_chat_completion_with_timeout():
"""
make chat completion call with low timeout and `mock_timeout`: true. Expect it to fail and correct timeout to be set in headers.
"""
async with aiohttp.ClientSession() as session:
model = "fake-openai-endpoint-5"
messages = [
{"role": "system", "content": text},
{"role": "user", "content": "Who was Alexander?"},
]
start_time = time.time()
response, headers = await chat_completion(
session=session,
key=os.environ["LITELLM_MASTER_KEY"],
model=model,
messages=messages,
num_retries=0,
mock_timeout=True,
return_headers=True,
)
end_time = time.time()
print(f"headers: {headers}")
assert (
headers["x-litellm-timeout"] == "1.0"
) # assert model-specific timeout used
@pytest.mark.asyncio
async def test_chat_completion_with_timeout_from_request():
"""
make chat completion call with low timeout and `mock_timeout`: true. Expect it to fail and correct timeout to be set in headers.
"""
async with aiohttp.ClientSession() as session:
model = "fake-openai-endpoint-5"
messages = [
{"role": "system", "content": text},
{"role": "user", "content": "Who was Alexander?"},
]
extra_headers = {
"x-litellm-timeout": "0.001",
}
start_time = time.time()
response, headers = await chat_completion(
session=session,
key=os.environ["LITELLM_MASTER_KEY"],
model=model,
messages=messages,
num_retries=0,
mock_timeout=True,
extra_headers=extra_headers,
return_headers=True,
)
end_time = time.time()
print(f"headers: {headers}")
assert (
headers["x-litellm-timeout"] == "0.001"
) # assert model-specific timeout used
@pytest.mark.parametrize("has_access", [True, False])
@pytest.mark.asyncio
async def test_chat_completion_client_fallbacks_with_custom_message(has_access: bool) -> None:
original_messages: Final = [{"role": "user", "content": "Who was Alexander?"}]
custom_messages: Final = [
{
"role": "user",
"content": (
"Describe the weather in a coastal city during winter, including the usual temperature, rain, wind, "
"and the clothing a visitor should bring."
),
}
]
models: Final = ["gpt-3.5-turbo", "gpt-6-luna"] if has_access else ["gpt-3.5-turbo"]
async with aiohttp.ClientSession() as session:
generated_key: Final = await generate_key(session=session, i=0, models=models)
async with AsyncOpenAI(api_key=generated_key["key"], base_url=PROXY_BASE_URL, max_retries=0) as client:
request: Final = {
"model": "gpt-3.5-turbo",
"messages": original_messages,
"max_tokens": 32,
"temperature": 0,
"extra_body": {
"mock_testing_fallbacks": True,
"fallbacks": [
{
"model": "gpt-6-luna",
"messages": custom_messages,
}
],
},
}
if not has_access:
with pytest.raises(PermissionDeniedError) as denied:
await client.chat.completions.create(**request)
assert denied.value.status_code == 403
assert "gpt-6-luna" in str(denied.value)
return
response: Final = await client.chat.completions.create(**request)
assert response.model == "gpt-6-luna"
assert response.choices[0].message.content
custom_control: Final = await client.chat.completions.create(
model="gpt-6-luna",
messages=custom_messages,
max_tokens=32,
temperature=0,
)
original_control: Final = await client.chat.completions.create(
model="gpt-6-luna",
messages=original_messages,
max_tokens=32,
temperature=0,
)
assert response.usage is not None
assert custom_control.usage is not None
assert original_control.usage is not None
assert custom_control.usage.completion_tokens > 0
assert original_control.usage.completion_tokens > 0
assert custom_control.usage.prompt_tokens != original_control.usage.prompt_tokens
assert response.usage.prompt_tokens == custom_control.usage.prompt_tokens
from typing import List
async def make_request(client: AsyncOpenAI, model: str) -> bool:
try:
await client.chat.completions.create(
model=model,
messages=[{"role": "user", "content": "Who was Alexander?"}],
)
return True
except Exception as e:
print(f"Error with {model}: {str(e)}")
return False
async def run_good_model_test(client: AsyncOpenAI, num_requests: int) -> bool:
tasks = [make_request(client, "good-model") for _ in range(num_requests)]
good_results = await asyncio.gather(*tasks)
return all(good_results)
@pytest.mark.asyncio
async def test_chat_completion_bad_and_good_model():
"""
Prod test - ensure even if bad model is down, good model is still working.
"""
client = AsyncOpenAI(api_key=os.environ["LITELLM_MASTER_KEY"], base_url="http://0.0.0.0:4000")
num_requests = 100
num_iterations = 3
for iteration in range(num_iterations):
print(f"\nIteration {iteration + 1}/{num_iterations}")
start_time = time.time()
# Fire and forget bad model requests
for _ in range(num_requests):
asyncio.create_task(make_request(client, "bad-model"))
# Wait only for good model requests
success = await run_good_model_test(client, num_requests)
print(
f"Iteration {iteration + 1}: {'✓' if success else '✗'} ({time.time() - start_time:.2f}s)"
)
assert success, "Not all good model requests succeeded"

View file

@ -1,102 +0,0 @@
"""
Test that Azure GPT-5 models support temperature parameter in Responses API.
"""
import pytest
from litellm.utils import ProviderConfigManager
from litellm.types.utils import LlmProviders
def test_azure_gpt5_supports_temperature():
"""Test that Azure GPT-5 uses the correct config that supports temperature."""
config = ProviderConfigManager.get_provider_responses_api_config(
provider=LlmProviders.AZURE, model="gpt-5"
)
# Should use AzureOpenAIResponsesAPIConfig, not AzureOpenAIOSeriesResponsesAPIConfig
assert type(config).__name__ == "AzureOpenAIResponsesAPIConfig"
# Should support temperature parameter
supported_params = config.get_supported_openai_params("gpt-5")
assert (
"temperature" in supported_params
), "Azure GPT-5 should support temperature parameter"
def test_azure_o_series_does_not_support_temperature():
"""Test that Azure O-series models still use the correct O-series config."""
test_models = ["o1", "o3"]
for model in test_models:
config = ProviderConfigManager.get_provider_responses_api_config(
provider=LlmProviders.AZURE, model=model
)
# Should use AzureOpenAIOSeriesResponsesAPIConfig
assert (
type(config).__name__ == "AzureOpenAIOSeriesResponsesAPIConfig"
), f"Azure {model} should use O-series config"
# Should NOT support temperature parameter
supported_params = config.get_supported_openai_params(model)
assert (
"temperature" not in supported_params
), f"Azure {model} should NOT support temperature parameter"
def test_openai_gpt5_supports_temperature():
"""Test that OpenAI GPT-5 supports temperature parameter."""
config = ProviderConfigManager.get_provider_responses_api_config(
provider=LlmProviders.OPENAI, model="gpt-5"
)
# Should use OpenAIResponsesAPIConfig
assert type(config).__name__ == "OpenAIResponsesAPIConfig"
# Should support temperature parameter
supported_params = config.get_supported_openai_params("gpt-5")
assert (
"temperature" in supported_params
), "OpenAI GPT-5 should support temperature parameter"
def test_azure_gpt5_variants_support_temperature():
"""Test that various GPT-5 model name variants support temperature."""
gpt5_variants = ["gpt-5", "gpt-5-turbo", "GPT-5", "azure/gpt-5"]
for model in gpt5_variants:
config = ProviderConfigManager.get_provider_responses_api_config(
provider=LlmProviders.AZURE, model=model
)
# All GPT-5 variants should use the base config, not O-series config
assert (
type(config).__name__ == "AzureOpenAIResponsesAPIConfig"
), f"Model '{model}' should not use O-series config"
# All should support temperature
supported_params = config.get_supported_openai_params(model)
assert (
"temperature" in supported_params
), f"Model '{model}' should support temperature parameter"
def test_azure_gpt_models_support_temperature():
"""Test that all GPT models (gpt-3.5, gpt-4, gpt-5, etc.) support temperature."""
gpt_models = ["gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "gpt-4o", "gpt-5"]
for model in gpt_models:
config = ProviderConfigManager.get_provider_responses_api_config(
provider=LlmProviders.AZURE, model=model
)
# All GPT models should use the base config, not O-series config
assert (
type(config).__name__ == "AzureOpenAIResponsesAPIConfig"
), f"Model '{model}' should not use O-series config"
# All should support temperature
supported_params = config.get_supported_openai_params(model)
assert (
"temperature" in supported_params
), f"Model '{model}' should support temperature parameter"

View file

@ -1,84 +0,0 @@
import os
# What this tests?
## Tests /health + /routes endpoints.
import pytest
import asyncio
import aiohttp
async def health(session, call_key):
url = "http://0.0.0.0:4000/health"
headers = {
"Authorization": f"Bearer {call_key}",
"Content-Type": "application/json",
}
async with session.get(url, headers=headers) as response:
status = response.status
response_text = await response.text()
print(f"Response (Status code: {status}):")
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 generate_key(session):
url = "http://0.0.0.0:4000/key/generate"
headers = {
"Authorization": "Bearer " + os.environ["LITELLM_MASTER_KEY"],
"Content-Type": "application/json",
}
data = {
"models": ["gpt-4", "text-embedding-ada-002", "gpt-image-1"],
"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()
@pytest.mark.asyncio
async def test_health():
"""
- Call /health
"""
async with aiohttp.ClientSession() as session:
# as admin #
all_healthy_models = await health(session=session, call_key=os.environ["LITELLM_MASTER_KEY"])
total_model_count = (
all_healthy_models["healthy_count"] + all_healthy_models["unhealthy_count"]
)
assert total_model_count > 0
@pytest.mark.asyncio
async def test_routes():
"""
Check if 200
"""
async with aiohttp.ClientSession() as session:
url = "http://0.0.0.0:4000/routes"
async with session.get(url) 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}")

View file

@ -2,57 +2,11 @@
## Tests /key endpoints.
import pytest
import asyncio, uuid
import asyncio
import aiohttp
from openai import AsyncOpenAI
import sys, os
import os
from typing import Optional
import litellm
from litellm.proxy._types import LitellmUserRoles
async def generate_team(
session, models: Optional[list] = None, team_id: Optional[str] = None
):
url = "http://0.0.0.0:4000/team/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
if team_id is None:
team_id = "litellm-dashboard"
data = {"team_id": team_id, **({"models": models} if models is not None else {})}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response (Status code: {status}):")
print(response_text)
print()
_json_response = await response.json()
return _json_response
async def generate_user(
session,
user_role="app_owner",
):
url = "http://0.0.0.0:4000/user/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {
"user_role": user_role,
"team_id": "litellm-dashboard",
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"Response (Status code: {status}):")
print(response_text)
print()
_json_response = await response.json()
return _json_response
async def generate_key(
session,
@ -99,25 +53,6 @@ async def generate_key(
return await response.json()
@pytest.mark.asyncio
async def test_key_gen():
async with aiohttp.ClientSession() as session:
tasks = [generate_key(session, i) for i in range(1, 11)]
await asyncio.gather(*tasks)
@pytest.mark.asyncio
async def test_simple_key_gen():
async with aiohttp.ClientSession() as session:
key_data = await generate_key(session, i=0)
key = key_data["key"]
assert key_data["token"] is not None
assert key_data["token"] != key
assert key_data["token_id"] is not None
assert key_data["created_at"] is not None
assert key_data["updated_at"] is not None
@pytest.mark.asyncio
async def test_key_gen_bad_key():
"""
@ -147,64 +82,6 @@ async def test_key_gen_bad_key():
pass
async def chat_completion(session, key, model="gpt-4"):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": model,
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
],
}
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(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 delete_key(session, get_key, auth_key=os.environ["LITELLM_MASTER_KEY"]):
"""
Delete key
"""
url = "http://0.0.0.0:4000/key/delete"
headers = {
"Authorization": f"Bearer {auth_key}",
"Content-Type": "application/json",
}
data = {"keys": [get_key]}
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_key_info(session, call_key, get_key=None):
"""
Make sure only models user has access to are returned
@ -235,111 +112,6 @@ async def get_key_info(session, call_key, get_key=None):
return await response.json()
async def get_model_list(session, call_key, endpoint: str = "/v1/models"):
"""
Make sure only models user has access to are returned
"""
url = "http://0.0.0.0:4000" + endpoint
headers = {
"Authorization": f"Bearer {call_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}. Responses {response_text}"
)
return await response.json()
async def get_model_info(session, call_key):
"""
Make sure only models user has access to are returned
"""
url = "http://0.0.0.0:4000/model/info"
headers = {
"Authorization": f"Bearer {call_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}. Responses {response_text}"
)
return await response.json()
@pytest.mark.asyncio
async def test_key_info():
"""
Get key info
- as admin -> 200
- as key itself -> 200
- as non existent key -> 404
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, i=0)
key = key_gen["key"]
# as admin #
await get_key_info(session=session, get_key=key, call_key=os.environ["LITELLM_MASTER_KEY"])
# as key itself #
await get_key_info(session=session, get_key=key, call_key=key)
# as key itself, use the auth param, and no query key needed
await get_key_info(session=session, call_key=key)
# as random key #
random_key = f"sk-{uuid.uuid4()}"
status = await get_key_info(session=session, get_key=random_key, call_key=key)
assert status == 404
@pytest.mark.asyncio
async def test_model_info():
"""
Get model info for models key has access to
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, i=0)
key = key_gen["key"]
# as admin #
admin_models = await get_model_info(session=session, call_key=os.environ["LITELLM_MASTER_KEY"])
admin_models = admin_models["data"]
# as key itself #
user_models = await get_model_info(session=session, call_key=key)
user_models = user_models["data"]
assert len(admin_models) > len(user_models)
assert len(user_models) > 0
async def get_spend_logs(session, request_id):
url = f"http://0.0.0.0:4000/spend/logs?request_id={request_id}"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_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.skip(reason="Frequent check on ci/cd leads to read timeout issue.")
@pytest.mark.asyncio
async def test_key_with_budgets():
@ -387,88 +159,3 @@ async def test_key_with_budgets():
# assert rounded_response_cost == rounded_key_info_spend
@pytest.mark.asyncio
async def test_key_delete_ui():
"""
Admin UI flow - DO NOT DELETE
-> Create a key with user_id = "ishaan"
-> Log on Admin UI, delete the key for user "ishaan"
-> This should work, since we're on the admin UI and role == "proxy_admin
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, i=0, user_id="ishaan-smart")
key = key_gen["key"]
# generate a admin UI key
team = await generate_team(session=session)
admin_ui_key = await generate_user(
session=session, user_role=LitellmUserRoles.PROXY_ADMIN.value
)
print(
"trying to delete key=",
key,
"using key=",
admin_ui_key["key"],
" to auth in",
)
await delete_key(
session=session,
get_key=key,
auth_key=admin_ui_key["key"],
)
@pytest.mark.parametrize("model_access", ["all-team-models", "gpt-3.5-turbo"])
@pytest.mark.parametrize("model_access_level", ["key", "team"])
@pytest.mark.parametrize("model_endpoint", ["/v1/models", "/model/info"])
@pytest.mark.asyncio
async def test_key_model_list(model_access, model_access_level, model_endpoint):
"""
Test if `/v1/models` works as expected.
"""
async with aiohttp.ClientSession() as session:
_models = [] if model_access == "all-team-models" else [model_access]
team_id = "litellm_dashboard_{}".format(uuid.uuid4())
new_team = await generate_team(
session=session,
models=_models if model_access_level == "team" else None,
team_id=team_id,
)
assert new_team["team_id"] == team_id
key_gen = await generate_key(
session=session,
i=0,
team_id=team_id,
models=_models if model_access_level == "key" else [],
)
key = key_gen["key"]
print(f"key: {key}")
model_list = await get_model_list(
session=session, call_key=key, endpoint=model_endpoint
)
print(f"model_list: {model_list}")
if model_access == "all-team-models":
if model_endpoint == "/v1/models":
assert not isinstance(model_list["data"][0]["id"], list)
assert isinstance(model_list["data"][0]["id"], str)
elif model_endpoint == "/model/info":
assert isinstance(model_list["data"], list)
assert len(model_list["data"]) > 0
if model_access == "gpt-3.5-turbo":
if model_endpoint == "/v1/models":
assert {entry["id"] for entry in model_list["data"]} == {
model_access,
"mistral-7b",
}, "generate_key sets alias mistral-7b -> gpt-3.5-turbo, so /v1/models lists both; model_access={}, model_access_level={}".format(
model_access, model_access_level
)
elif model_endpoint == "/model/info":
assert isinstance(model_list["data"], list)
assert len(model_list["data"]) == 1

View file

@ -2,61 +2,7 @@
Unit test for LiteLLM Proxy Responses API configuration.
"""
import pytest
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
def test_litellm_proxy_responses_api_config():
"""Test that litellm_proxy provider returns correct Responses API config"""
from litellm.llms.litellm_proxy.responses.transformation import (
LiteLLMProxyResponsesAPIConfig,
)
config = ProviderConfigManager.get_provider_responses_api_config(
model="litellm_proxy/gpt-5.5",
provider=LlmProviders.LITELLM_PROXY,
)
print(f"config: {config}")
assert config is not None, "Config should not be None for litellm_proxy provider"
assert isinstance(
config, LiteLLMProxyResponsesAPIConfig
), f"Expected LiteLLMProxyResponsesAPIConfig, got {type(config)}"
assert (
config.custom_llm_provider == LlmProviders.LITELLM_PROXY
), "custom_llm_provider should be LITELLM_PROXY"
def test_litellm_proxy_responses_api_config_get_complete_url():
"""Test that get_complete_url works correctly"""
import os
from litellm.llms.litellm_proxy.responses.transformation import (
LiteLLMProxyResponsesAPIConfig,
)
config = LiteLLMProxyResponsesAPIConfig()
# Test with explicit api_base
url = config.get_complete_url(
api_base="https://my-proxy.example.com",
litellm_params={},
)
assert url == "https://my-proxy.example.com/responses"
# Test with trailing slash
url = config.get_complete_url(
api_base="https://my-proxy.example.com/",
litellm_params={},
)
assert url == "https://my-proxy.example.com/responses"
# Test that it raises error when api_base is None and env var is not set
if "LITELLM_PROXY_API_BASE" in os.environ:
del os.environ["LITELLM_PROXY_API_BASE"]
with pytest.raises(ValueError, match="api_base not set"):
config.get_complete_url(api_base=None, litellm_params={})
def test_litellm_proxy_responses_api_config_inherits_from_openai():
@ -70,15 +16,6 @@ def test_litellm_proxy_responses_api_config_inherits_from_openai():
config = LiteLLMProxyResponsesAPIConfig()
# Should inherit from OpenAI config
assert isinstance(config, OpenAIResponsesAPIConfig)
# Should have the correct provider set
assert config.custom_llm_provider == LlmProviders.LITELLM_PROXY
if __name__ == "__main__":
test_litellm_proxy_responses_api_config()
test_litellm_proxy_responses_api_config_get_complete_url()
test_litellm_proxy_responses_api_config_inherits_from_openai()
print("All tests passed!")

View file

@ -5,8 +5,6 @@ import pytest
import asyncio
import aiohttp
import os
import dotenv
from typing import Final
from dotenv import load_dotenv
load_dotenv()
@ -32,46 +30,6 @@ async def generate_key(session, models=[]):
return await response.json()
async def get_models(session, key, only_model_access_groups=False):
url = "http://0.0.0.0:4000/models"
if only_model_access_groups:
url += "?only_model_access_groups=True"
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 /models")
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_get_models_multiple_tests():
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session)
key = key_gen["key"]
models = await get_models(session=session, key=key)
print(f"\n\nmodels: {models}")
assert len(models["data"]) > 0
## Test only_model_access_groups
new_response = await get_models(
session=session, key=key, only_model_access_groups=True
)
print(f"\n\nnew_response: {new_response}")
assert (
len(new_response["data"]) == 0
) # no model access groups set on config.yaml
async def add_models(
session, model_id="123", model_name="azure-gpt-3.5", key=os.environ["LITELLM_MASTER_KEY"], team_id=None
):
@ -106,48 +64,6 @@ async def add_models(
return response_json
async def get_model_info(session, key, litellm_model_id=None):
"""
Make sure only models user has access to are returned
"""
if litellm_model_id:
url = f"http://0.0.0.0:4000/model/info?litellm_model_id={litellm_model_id}"
else:
url = "http://0.0.0.0:4000/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_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
async def get_model_group_info(session, key):
url = "http://0.0.0.0:4000/model_group/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_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="azure-gpt-3.5"):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
@ -173,49 +89,6 @@ async def chat_completion(session, key, model="azure-gpt-3.5"):
raise Exception(f"Request did not return a 200 status code: {status}")
@pytest.mark.asyncio
async def test_get_models():
"""
Get models user has access to
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, models=["gpt-4"])
key = key_gen["key"]
response = await get_model_info(session=session, key=key)
models = [m["model_name"] for m in response["data"]]
for m in models:
assert m == "gpt-4"
@pytest.mark.asyncio
async def test_get_specific_model():
"""
Return specific model info
Ensure value of model_info is same as on `/model/info` (no id set)
"""
async with aiohttp.ClientSession() as session:
key_gen = await generate_key(session=session, models=["gpt-4"])
key = key_gen["key"]
response = await get_model_info(session=session, key=key)
models = [m["model_name"] for m in response["data"]]
model_specific_info = None
for idx, m in enumerate(models):
assert m == "gpt-4"
litellm_model_id = response["data"][idx]["model_info"]["id"]
model_specific_info = response["data"][idx]
assert litellm_model_id is not None
response = await get_model_info(
session=session, key=key, litellm_model_id=litellm_model_id
)
assert response["data"][0]["model_info"]["id"] == litellm_model_id
assert (
response["data"][0] == model_specific_info
), "Model info is not the same. Got={}, Expected={}".format(
response["data"][0], model_specific_info
)
async def delete_model(session, model_id="123", key=os.environ["LITELLM_MASTER_KEY"]):
"""
Make sure only models user has access to are returned
@ -270,45 +143,3 @@ async def test_add_and_delete_models():
pass
@pytest.mark.asyncio
async def test_get_personal_models_for_user():
"""
Test /models endpoint with team
"""
from tests.test_users import new_user
async with aiohttp.ClientSession() as session:
# Creat a user
user_data = await new_user(session=session, i=0, models=["gpt-3.5-turbo"])
user_id = user_data["user_id"]
user_api_key = user_data["key"]
model_group_info = await get_model_group_info(session=session, key=user_api_key)
print(model_group_info)
assert len(model_group_info["data"]) == 1
assert model_group_info["data"][0]["model_group"] == "gpt-3.5-turbo"
@pytest.mark.asyncio
async def test_model_group_info_e2e():
"""
Test /model/group/info endpoint
"""
async with aiohttp.ClientSession() as session:
models = await get_models(session=session, key=os.environ["LITELLM_MASTER_KEY"])
print(models)
model_group_info = await get_model_group_info(session=session, key=os.environ["LITELLM_MASTER_KEY"])
print(model_group_info)
model_groups: Final = [m["model_group"] for m in model_group_info["data"]]
assert "anthropic/*" not in model_groups, (
f"Expected 'anthropic/*' to be expanded, but it was returned verbatim: {model_groups}"
)
assert any(m.startswith("anthropic/") for m in model_groups), (
f"Expected concrete anthropic models from the 'anthropic/*' config entry, got: {model_groups}"
)

View file

@ -4,43 +4,12 @@ Tests both basic functionality and complex scenarios including target_model_name
"""
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, patch
import pytest
import litellm
from litellm.proxy._types import UserAPIKeyAuth
@pytest.mark.asyncio
async def test_vector_store_update_basic():
"""Test basic vector store update functionality."""
mock_response = {
"id": "vs_test123",
"object": "vector_store",
"created_at": 1699061776,
"name": "Updated Name",
"metadata": {"key": "value"},
"status": "completed",
}
with patch(
"litellm.vector_stores.main.aupdate",
new=AsyncMock(return_value=mock_response),
) as mock_update:
router = litellm.Router(model_list=[])
result = await router.avector_store_update(
vector_store_id="vs_test123",
name="Updated Name",
metadata={"key": "value"},
custom_llm_provider="openai",
)
assert result["id"] == "vs_test123"
assert result["name"] == "Updated Name"
assert result["metadata"]["key"] == "value"
mock_update.assert_called_once()
@pytest.mark.asyncio
@ -66,71 +35,6 @@ async def test_async_vector_store_update():
mock_aupdate.assert_called_once()
@pytest.mark.asyncio
async def test_vector_store_list_with_pagination():
"""Test vector store list with pagination parameters."""
mock_response = {
"object": "list",
"data": [{"id": f"vs_{i}"} for i in range(5)],
"has_more": True,
"first_id": "vs_0",
"last_id": "vs_4",
}
with patch(
"litellm.vector_stores.main.list",
return_value=mock_response,
) as mock_list:
router = litellm.Router(model_list=[])
result = router.vector_store_list(
limit=5,
after="vs_previous",
order="asc",
custom_llm_provider="openai",
)
assert result["has_more"] is True
assert len(result["data"]) == 5
# Verify pagination params were passed
call_kwargs = mock_list.call_args.kwargs
assert call_kwargs["limit"] == 5
assert call_kwargs["after"] == "vs_previous"
assert call_kwargs["order"] == "asc"
@pytest.mark.asyncio
async def test_vector_store_update_with_expires_after():
"""Test vector store update with expiration policy."""
expires_after = {
"anchor": "last_active_at",
"days": 7,
}
mock_response = {
"id": "vs_test123",
"expires_after": expires_after,
"expires_at": 1699668576,
}
with patch(
"litellm.vector_stores.main.update",
return_value=mock_response,
) as mock_update:
router = litellm.Router(model_list=[])
result = router.vector_store_update(
vector_store_id="vs_test123",
expires_after=expires_after,
custom_llm_provider="openai",
)
assert result["expires_after"]["days"] == 7
assert result["expires_at"] is not None
call_kwargs = mock_update.call_args.kwargs
assert call_kwargs["expires_after"] == expires_after
def test_router_initializes_new_endpoints():
"""Test that router properly initializes the new vector store endpoints."""
router = litellm.Router(model_list=[])
@ -165,11 +69,6 @@ if __name__ == "__main__":
test_router_initializes_new_endpoints()
print("✓ Router initialization successful")
# Test basic sync operations
print("✓ Testing basic sync operations...")
asyncio.run(test_vector_store_update_basic())
print("✓ Basic sync operations successful")
# Test async operations
print("✓ Testing async operations...")
asyncio.run(test_async_vector_store_update())

View file

@ -3,10 +3,7 @@ from typing import Final
# 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, AzureOpenAI, AsyncAzureOpenAI
from typing import Optional, List, Union
from openai import AsyncOpenAI
LITELLM_MASTER_KEY = os.environ["LITELLM_MASTER_KEY"]
@ -19,64 +16,6 @@ def response_header_check(response):
assert headers_size < 4096, "Response headers exceed the 4kb limit"
async def generate_key(
session,
models=[
"gpt-4",
"text-embedding-ada-002",
"gpt-image-1",
"fake-openai-endpoint-2",
"mistral-embed",
],
):
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "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}")
response_header_check(
response
) # calling the function to check response headers
return await response.json()
async def new_user(session):
url = "http://0.0.0.0:4000/user/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {
"models": ["gpt-4", "text-embedding-ada-002", "gpt-image-1"],
"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}")
response_header_check(
response
) # calling the function to check response headers
return await response.json()
async def moderation(session, key):
url = "http://0.0.0.0:4000/moderations"
headers = {
@ -98,141 +37,6 @@ async def moderation(session, key):
return await response.json()
async def chat_completion(session, key, model: Union[str, List] = "gpt-4"):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": model,
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
],
}
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}, response text={response_text}"
)
response_header_check(
response
) # calling the function to check response headers
return await response.json()
async def queue_chat_completion(
session, key, priority: int, model: Union[str, List] = "gpt-4"
):
url = "http://0.0.0.0:4000/queue/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": model,
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
],
"priority": priority,
}
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 response.raw_headers
async def chat_completion_with_headers(session, key, model="gpt-4"):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": model,
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
],
}
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}")
response_header_check(
response
) # calling the function to check response headers
raw_headers = response.raw_headers
raw_headers_json = {}
for (
item
) in (
response.raw_headers
): # ((b'date', b'Fri, 19 Apr 2024 21:17:29 GMT'), (), )
raw_headers_json[item[0].decode("utf-8")] = item[1].decode("utf-8")
return raw_headers_json
async def chat_completion_with_model_from_route(session, key, route):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
async def completion(session, key):
url = "http://0.0.0.0:4000/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {"model": "gpt-4", "prompt": "Hello!"}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
response_header_check(
response
) # calling the function to check response headers
response = await response.json()
return response
async def embeddings(session, key, model="text-embedding-ada-002"):
url = "http://0.0.0.0:4000/embeddings"
headers = {
@ -258,120 +62,6 @@ async def embeddings(session, key, model="text-embedding-ada-002"):
) # calling the function to check response headers
async def image_generation(session, key):
url = "http://0.0.0.0:4000/images/generations"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": "gpt-image-1",
"prompt": "A cute baby sea otter",
}
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:
if (
"Connection error" in response_text
): # OpenAI endpoint returns a connection error
return
raise Exception(f"Request did not return a 200 status code: {status}")
response_header_check(
response
) # calling the function to check response headers
@pytest.mark.asyncio
async def test_chat_completion():
"""
- 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, models=["gpt-3.5-turbo"])
azure_client = AsyncAzureOpenAI(
azure_endpoint="http://0.0.0.0:4000",
azure_deployment="random-model",
api_key=key_gen["key"],
api_version="2024-02-15-preview",
)
with pytest.raises(openai.PermissionDeniedError) as e:
response = await azure_client.chat.completions.create(
model="gpt-4",
messages=[{"role": "user", "content": "Hello!"}],
)
assert "is not available for this API key" in str(e.value)
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=1)
@pytest.mark.skip(reason="Flaky test, this works locally but not on CI")
async def test_chat_completion_ratelimit():
"""
- call model with rpm 1
- make 2 parallel calls
- make sure 1 fails
"""
async with aiohttp.ClientSession() as session:
# key_gen = await generate_key(session=session)
key = os.environ["LITELLM_MASTER_KEY"]
tasks = []
tasks.append(
chat_completion(session=session, key=key, model="fake-openai-endpoint-2")
)
tasks.append(
chat_completion(session=session, key=key, model="fake-openai-endpoint-2")
)
try:
await asyncio.gather(*tasks)
pytest.fail("Expected at least 1 call to fail")
except Exception as e:
if "Request did not return a 200 status code: 429" in str(e):
pass
else:
pytest.fail(f"Wrong error received - {str(e)}")
@pytest.mark.asyncio
@pytest.mark.skip(reason="Flaky test")
async def test_chat_completion_different_deployments():
"""
- call model group with 2 deployments
- make 5 calls
- expect 2 unique deployments
"""
async with aiohttp.ClientSession() as session:
# key_gen = await generate_key(session=session)
key = os.environ["LITELLM_MASTER_KEY"]
results = []
for _ in range(20):
results.append(
await chat_completion_with_headers(
session=session, key=key, model="fake-openai-endpoint-3"
)
)
try:
print(f"results: {results}")
init_model_id = results[0]["x-litellm-model-id"]
deployments_shuffled = False
for result in results[1:]:
if init_model_id != result["x-litellm-model-id"]:
deployments_shuffled = True
if deployments_shuffled == False:
pytest.fail("Expected at least 1 shuffled call")
except Exception as e:
pass
@pytest.mark.asyncio
async def test_chat_completion_streaming():
"""
@ -425,46 +115,3 @@ async def test_completion_streaming_usage_metrics():
assert last_chunk.usage.total_tokens > 0, "Total tokens should be greater than 0"
@pytest.mark.asyncio
async def test_proxy_all_models():
"""
- proxy_server_config.yaml has model = * / *
- Make chat completion call
- groq is NOT defined on /models
"""
async with aiohttp.ClientSession() as session:
# 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=LITELLM_MASTER_KEY, model="groq/openai/gpt-oss-120b"
)
await chat_completion(
session=session,
key=LITELLM_MASTER_KEY,
model="anthropic/claude-sonnet-4-5-20250929",
)
@pytest.mark.asyncio
async def test_batch_chat_completions():
"""
- Make chat completion call using
"""
async with aiohttp.ClientSession() as session:
# call chat/completions with a model that the key was not created for + the model is not on the config.yaml
response = await chat_completion(
session=session,
key=os.environ["LITELLM_MASTER_KEY"],
model="gpt-3.5-turbo,fake-openai-endpoint",
)
print(f"response: {response}")
assert len(response) == 2
assert isinstance(response, list)

View file

@ -1,91 +0,0 @@
import sys
import os
import threading
import time
import pytest
# Add the project root to the path
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
from litellm.types.utils import StandardCallbackDynamicParams
def get_thread_count() -> int:
"""Helper to get active thread count"""
return threading.active_count()
@pytest.fixture
def otel_logger():
"""Fixture to provide a clean OTEL logger for each test"""
config = OpenTelemetryConfig(
exporter="console", enable_metrics=False, service_name="litellm-unit-test"
)
return OpenTelemetry(config=config)
def test_otel_thread_leak_dynamic_headers(otel_logger):
"""
Unit test to verify that calling get_tracer_to_use_for_request with
dynamic headers doesn't cause a linear thread leak.
This test reproduces the issue where each unique team/key credential
set causes a new TracerProvider (and its background threads) to be
spawned but never closed.
"""
# 1. Setup dynamic header simulation (monkey-patch)
# This simulates what LangfuseOtelLogger does for per-team keys
def mock_construct_dynamic_headers(standard_callback_dynamic_params):
if standard_callback_dynamic_params:
return {"Authorization": "Bearer fake_token"}
return None
otel_logger.construct_dynamic_otel_headers = mock_construct_dynamic_headers
# 2. Establish Baseline
initial_threads = get_thread_count()
# 3. Simulate requests
num_requests = 10
latencies = []
print("\n🚀 Simulating requests with dynamic headers:")
for i in range(num_requests):
kwargs = {
"standard_callback_dynamic_params": StandardCallbackDynamicParams(
langfuse_public_key=f"key_{i}",
langfuse_secret_key=f"secret_{i}",
)
}
# Measure latency
start_time = time.perf_counter()
tracer = otel_logger.get_tracer_to_use_for_request(kwargs)
end_time = time.perf_counter()
latency_ms = (end_time - start_time) * 1000
latencies.append(latency_ms)
print(f" Request {i+1:2d}: Latency = {latency_ms:6.2f} ms")
# Verify a tracer was actually returned
assert tracer is not None
avg_latency = sum(latencies) / len(latencies)
print(f"\n📊 Average Latency: {avg_latency:.2f} ms")
# 4. Check for leaks
# Allow for a small constant increase (OTEL might start a few shared threads)
# but a linear leak would result in +10 or more threads here.
final_threads = get_thread_count()
thread_delta = final_threads - initial_threads
print(f"\nThread growth: {thread_delta} threads across {num_requests} requests")
# ASSERTION: The growth should be significantly less than 1 thread per request.
# If the bug exists, thread_delta will be >= num_requests.
assert thread_delta < (num_requests / 2), (
f"Thread leak detected! Threads grew by {thread_delta} over {num_requests} requests. "
"Each request with dynamic headers appears to be leaking background threads."
)

View file

@ -1,77 +0,0 @@
import asyncio
import aiohttp
import pytest
from unittest.mock import MagicMock, patch
from litellm.proxy.guardrails.guardrail_hooks.presidio import (
OPTIONAL_PresidioPIIMasking,
)
@pytest.mark.asyncio
async def test_sanity_presidio_session_reuse_main_thread():
"""
SANITY CHECK:
Verify that Presidio guardrail reuses sessions in the main thread.
This ensures we don't break existing session pooling functionality.
"""
presidio = OPTIONAL_PresidioPIIMasking(
mock_testing=True,
presidio_analyzer_api_base="http://mock-analyzer",
presidio_anonymizer_api_base="http://mock-anonymizer",
)
session_creations = 0
original_init = aiohttp.ClientSession.__init__
def mocked_init(self, *args, **kwargs):
nonlocal session_creations
session_creations += 1
original_init(self, *args, **kwargs)
with patch.object(aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True):
for _ in range(10):
async with presidio._get_session_iterator() as session:
pass
# Expected: Only 1 session created for all 10 calls.
assert session_creations == 1
await presidio._close_http_session()
@pytest.mark.asyncio
async def test_bug_presidio_session_explosion_background_thread_causes_latency():
"""
BUG REPRODUCTION:
Verify that background threads (like logging hooks) REUSE sessions.
Previously, each call in a background loop created a NEW ephemeral session,
leading to socket exhaustion and the reported 97s latency spike.
"""
import threading
presidio = OPTIONAL_PresidioPIIMasking(
mock_testing=True,
presidio_analyzer_api_base="http://mock-analyzer",
presidio_anonymizer_api_base="http://mock-anonymizer",
)
# Force the code to think it's in a background thread
presidio._main_thread_id = threading.get_ident() + 1
session_creations = 0
original_init = aiohttp.ClientSession.__init__
def mocked_init(self, *args, **kwargs):
nonlocal session_creations
session_creations += 1
original_init(self, *args, **kwargs)
with patch.object(aiohttp.ClientSession, "__init__", side_effect=mocked_init, autospec=True):
for _ in range(10):
async with presidio._get_session_iterator() as session:
pass
# FIX VERIFICATION: Should now be 1 session (reused) instead of 10.
assert session_creations == 1
await presidio._close_http_session()

View file

@ -1,170 +0,0 @@
# %%
import asyncio
import os
import pytest
import random
from typing import Any
from dotenv import load_dotenv
load_dotenv()
import litellm
from pydantic import BaseModel
from litellm import utils, Router
COMPLETION_TOKENS = 5
base_model_list = [
{
"model_name": "gpt-5-mini",
"litellm_params": {
"model": "gpt-5-mini",
"api_key": os.getenv("OPENAI_API_KEY"),
"max_tokens": COMPLETION_TOKENS,
},
}
]
class RouterConfig(BaseModel):
rpm: int
tpm: int
@pytest.fixture(scope="function")
def router_factory():
def create_router(rpm, tpm, routing_strategy):
model_list = base_model_list.copy()
model_list[0]["rpm"] = rpm
model_list[0]["tpm"] = tpm
return Router(
model_list=model_list,
routing_strategy=routing_strategy,
enable_pre_call_checks=True,
debug_level="DEBUG",
)
return create_router
def generate_list_of_messages(num_messages):
"""
create num_messages new chat conversations
"""
return [
[{"role": "user", "content": f"{i}. Hey, how's it going? {random.random()}"}]
for i in range(num_messages)
]
def calculate_limits(list_of_messages):
"""
Return the min rpm and tpm level that would let all messages in list_of_messages be sent this minute
"""
rpm = len(list_of_messages)
tpm = sum(
(utils.token_counter(messages=m) + COMPLETION_TOKENS for m in list_of_messages)
)
return rpm, tpm
async def async_call(router: Router, list_of_messages) -> Any:
tasks = [
router.acompletion(model="gpt-5-mini", messages=m) for m in list_of_messages
]
return await asyncio.gather(*tasks)
def sync_call(router: Router, list_of_messages) -> Any:
return [
router.completion(model="gpt-5-mini", messages=m) for m in list_of_messages
]
class ExpectNoException(Exception):
pass
@pytest.mark.parametrize(
"num_try_send, num_allowed_send",
[
(2, 3), # sending as many as allowed, ExpectNoException
# (10, 10), # sending as many as allowed, ExpectNoException
(3, 2), # Sending more than allowed, ValueError
# (10, 9), # Sending more than allowed, ValueError
],
)
@pytest.mark.parametrize(
"sync_mode", [True, False]
) # Use parametrization for sync/async
@pytest.mark.parametrize(
"routing_strategy",
[
"usage-based-routing",
# "simple-shuffle", # dont expect to rate limit
# "least-busy", # dont expect to rate limit
# "latency-based-routing",
],
)
def test_async_rate_limit(
router_factory, num_try_send, num_allowed_send, sync_mode, routing_strategy
):
"""
Check if router.completion and router.acompletion can send more messages than they've been limited to.
Args:
router_factory: makes new router object, without any shared Global state
num_try_send (int): number of messages to try to send
num_allowed_send (int): max number of messages allowed to be sent in 1 minute
sync_mode (bool): if making sync (router.completion) or async (router.acompletion)
Raises:
ValueError: Error router throws when it hits rate limits
ExpectNoException: Signfies that no other error has happened. A NOP
"""
# Can send more messages then we're going to; so don't expect a rate limit error
litellm.logging_callback_manager._reset_all_callbacks()
args = locals()
print(f"args: {args}")
expected_exception = (
ExpectNoException if num_try_send <= num_allowed_send else ValueError
)
# usage-based-routing tracks RPM in log_success_event which runs in a
# background ThreadPoolExecutor. The cache update races with the next
# call's routing check, so over-limit detection is non-deterministic in
# both sync tight-loops and async concurrent gathers.
if num_try_send > num_allowed_send:
pytest.skip(
"RPM tracking via background thread is racy; "
"RPM over-limit rejection is tested for usage-based-routing-v2 in "
"tests/unit/router_strategy/test_router_routing_groups.py"
)
list_of_messages = generate_list_of_messages(max(num_try_send, num_allowed_send))
rpm, tpm = calculate_limits(list_of_messages[:num_allowed_send])
list_of_messages = list_of_messages[:num_try_send]
router: Router = router_factory(rpm, tpm, routing_strategy)
print(f"router: {router.model_list}")
received = []
def _send_and_check():
results = (
sync_call(router, list_of_messages)
if sync_mode
else asyncio.run(async_call(router, list_of_messages))
)
received.extend(results)
print(results)
if len([i for i in results if i is not None]) != num_try_send:
# since not all results got returned, raise rate limit error
raise ValueError("No deployments available for selected model")
raise ExpectNoException
with pytest.raises(expected_exception) as excinfo: # asserts correct type raised
_send_and_check()
print(expected_exception, excinfo)
if expected_exception is ValueError:
assert "No deployments available for selected model" in str(excinfo.value)
else:
assert len([i for i in received if i is not None]) == num_try_send

View file

@ -1,117 +0,0 @@
"""
Test that async HTTP clients are properly cleaned up to prevent resource leaks.
Issue: https://github.com/BerriAI/litellm/issues/12107
"""
import asyncio
import os
import warnings
import pytest
import litellm
@pytest.mark.asyncio
async def test_acompletion_resource_cleanup():
"""Test that acompletion doesn't leave unclosed client sessions."""
# Suppress warnings to check for them later
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
# Make an async completion call
response = await litellm.acompletion(
model="gemini/gemini-2.0-flash-lite-001",
messages=[{"role": "user", "content": "Hello"}],
mock_response="Hi there! How can I help you today?",
)
# Check that response was received
assert (
response.choices[0].message.content == "Hi there! How can I help you today?"
)
# Manually close async clients
await litellm.close_litellm_async_clients()
# Give a small delay for any warnings to appear
await asyncio.sleep(0.1)
# Check for resource warnings
resource_warnings = [
warning
for warning in w
if "Unclosed" in str(warning.message)
and (
"client session" in str(warning.message)
or "connector" in str(warning.message)
)
]
# Should be no unclosed resource warnings
assert (
len(resource_warnings) == 0
), f"Found unclosed resources: {[str(w.message) for w in resource_warnings]}"
@pytest.mark.asyncio
async def test_multiple_acompletion_calls_cleanup():
"""Test that multiple acompletion calls reuse clients and don't leak resources."""
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
# Make multiple async completion calls
for i in range(3):
response = await litellm.acompletion(
model="gemini/gemini-2.0-flash-lite-001",
messages=[{"role": "user", "content": f"Hello {i}"}],
mock_response=f"Response {i}",
)
assert response.choices[0].message.content == f"Response {i}"
# Clean up
await litellm.close_litellm_async_clients()
# Give a small delay for any warnings to appear
await asyncio.sleep(0.1)
# Check for resource warnings
resource_warnings = [
warning
for warning in w
if "Unclosed" in str(warning.message)
and (
"client session" in str(warning.message)
or "connector" in str(warning.message)
)
]
assert (
len(resource_warnings) == 0
), f"Found unclosed resources: {[str(w.message) for w in resource_warnings]}"
@pytest.mark.asyncio
async def test_cleanup_function_is_safe_to_call_multiple_times():
"""Test that the cleanup function can be called multiple times safely."""
# This should not raise any errors
await litellm.close_litellm_async_clients()
await litellm.close_litellm_async_clients()
await litellm.close_litellm_async_clients()
# Should still work after multiple cleanups
response = await litellm.acompletion(
model="gemini/gemini-2.0-flash-lite-001",
messages=[{"role": "user", "content": "Hello"}],
mock_response="Hi!",
)
assert response.choices[0].message.content == "Hi!"
# Clean up again
await litellm.close_litellm_async_clients()
if __name__ == "__main__":
# Run the test
asyncio.run(test_acompletion_resource_cleanup())
print("✅ All tests passed!")

View file

@ -1,119 +1,11 @@
from __future__ import annotations
import importlib
from typing import Final
import pytest
models = importlib.import_module("tests.rust-python-harness.shared.reporting.models")
strategy_module = importlib.import_module("tests.rust-python-harness.shared.reporting.strategy")
ui = importlib.import_module("tests.rust-python-harness.shared.reporting.ui")
contracts = importlib.import_module("tests.rust-python-harness.shared.unit_runners.contracts")
cli = importlib.import_module("tests.rust-python-harness.cli")
UNIT_TEST_CONTRACTS = contracts.UNIT_TEST_CONTRACTS
CaseResult = models.CaseResult
Coverage = models.Coverage
HarnessCase = models.HarnessCase
HarnessRun = models.HarnessRun
RunStatus = models.RunStatus
ModuleCaseSpec = strategy_module.ModuleCaseSpec
NotImplementedCaseSpec = strategy_module.NotImplementedCaseSpec
SkippedCaseSpec = strategy_module.SkippedCaseSpec
_format_duration = ui._format_duration
_summary = ui._summary
def _case(module: str = "tests.example") -> HarnessCase:
return HarnessCase(
strategy_id="example",
strategy_label="Example",
sdk_function="messages",
spec=ModuleCaseSpec(coverage=Coverage.COMPLETE, module=module),
)
@pytest.mark.parametrize(
"module",
[
"tests.rust-python-harness.strategies.trace_parity.sdk.messages.case",
"tests.rust-python-harness.strategies.trace_parity.sdk.chat_completions.case",
"tests.rust-python-harness.strategies.trace_parity.sdk.transcription.case",
],
)
def test_implemented_namespace_case_modules_remain_importable(module: str) -> None:
assert importlib.import_module(module)
def test_should_mark_not_implemented_and_skipped_cases_without_running() -> None:
not_implemented: Final = CaseResult(
case=HarnessCase(
strategy_id="example",
strategy_label="Example",
sdk_function="messages",
spec=NotImplementedCaseSpec(reason="No case is registered."),
)
)
skipped: Final = CaseResult(
case=HarnessCase(
strategy_id="example",
strategy_label="Example",
sdk_function="messages",
spec=SkippedCaseSpec(reason="The surface does not apply."),
)
)
not_implemented.set_initial_status()
skipped.set_initial_status()
assert not_implemented.status is RunStatus.NOT_IMPLEMENTED
assert skipped.status is RunStatus.SKIPPED
def test_should_finalize_a_fully_passing_case() -> None:
result = CaseResult(case=_case())
result.set_initial_status()
result.collected.update({"one", "two"})
result.completed.update({"one", "two"})
result.passed = 2
result.finalize()
assert result.status is RunStatus.PASSED
def test_should_replace_a_pass_with_a_teardown_error() -> None:
result = CaseResult(case=_case())
result.set_initial_status()
result.collected.add("one")
result.record("one", RunStatus.PASSED, 0.1)
result.record("one", RunStatus.ERROR, 0.2)
assert result.status is RunStatus.ERROR
assert result.passed == 0
assert result.errors == 1
assert result.duration == pytest.approx(0.3)
def test_should_format_developer_facing_run_context() -> None:
run = HarnessRun.from_cases((_case(),))
result = next(iter(run.results.values()))
result.collected.add("tests/test_parity.py::test_one")
result.record("tests/test_parity.py::test_one", RunStatus.PASSED, 1.25)
assert _summary(run) == (1, 0, 0, 0)
assert _format_duration(1.25) == "1.2s"
def test_should_leave_functions_without_unit_test_contracts_unimplemented() -> None:
assert "messages" not in UNIT_TEST_CONTRACTS
def test_strategy_subcommand_accepts_function_filter(capsys: pytest.CaptureFixture[str]) -> None:
exit_code: Final = cli.main(["run", "unit_tests_rust", "--function", "messages"])
captured: Final = capsys.readouterr()
assert exit_code == 0
assert "- messages: not_implemented" in captured.out
assert "unit_tests_rust:messages: not_implemented" not in captured.out

View file

@ -1,18 +1,13 @@
import os
import sys
import unittest
from datetime import datetime
from unittest.mock import patch, AsyncMock, MagicMock
from unittest.mock import patch
# Add the project root to sys.path
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))
import litellm
from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger
from litellm.integrations.opentelemetry import OpenTelemetry
from litellm.types.services import ServiceTypes
from litellm._service_logger import ServiceLogging
from litellm.types.utils import StandardCallbackDynamicParams
class TestServiceLoggerOTEL(unittest.IsolatedAsyncioTestCase):
@ -43,110 +38,6 @@ class TestServiceLoggerOTEL(unittest.IsolatedAsyncioTestCase):
"LangfuseOtelLogger.async_service_failure_hook",
)
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_tracing")
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_metrics")
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_logs")
async def test_langfuse_otel_does_not_create_proxy_request_span(
self, mock_logs, mock_metrics, mock_tracing
):
"""
Test that LangfuseOtelLogger returns None for create_litellm_proxy_request_started_span.
This prevents empty proxy request spans from being sent to Langfuse when
requests don't result in actual LLM calls (e.g., auth failures, health checks).
"""
logger = LangfuseOtelLogger()
# Verify the method is overridden
self.assertEqual(
logger.create_litellm_proxy_request_started_span.__qualname__,
"LangfuseOtelLogger.create_litellm_proxy_request_started_span",
)
# Verify it returns None
result = logger.create_litellm_proxy_request_started_span(
start_time=datetime.now(),
headers={"Authorization": "Bearer test"},
)
self.assertIsNone(result)
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_tracing")
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_metrics")
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_logs")
async def test_service_logging_shadowing_fix(
self, mock_logs, mock_metrics, mock_tracing
):
"""
Test the architectural fix: multiple OTEL loggers should receive logs independently.
"""
# 1. Initialize two loggers
langfuse_logger = LangfuseOtelLogger()
otel_logger = OpenTelemetry()
# 2. Setup service_callback list
litellm.service_callback = [langfuse_logger, otel_logger]
service_logging = ServiceLogging()
# 3. Mock the base OpenTelemetry hook
with patch.object(
OpenTelemetry, "async_service_success_hook", new_callable=AsyncMock
) as mock_base_hook:
# Trigger a service event
await service_logging.async_service_success_hook(
service=ServiceTypes.DB,
call_type="success",
duration=0.1,
parent_otel_span=MagicMock(),
start_time=0.0,
end_time=1.0,
)
# The architectural fix ensures we call each correctly.
self.assertEqual(
mock_base_hook.call_count,
1,
"Generic OTEL logger should have received the log exactly once.",
)
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_tracing")
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_metrics")
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_logs")
async def test_langfuse_otel_env_config_includes_v4_ingestion_header(
self, mock_logs, mock_metrics, mock_tracing
):
logger = LangfuseOtelLogger()
headers = OpenTelemetry._get_headers_dictionary(logger.config.headers)
self.assertEqual(
headers["x-langfuse-ingestion-version"],
"4",
)
self.assertTrue(headers["Authorization"].startswith("Basic "))
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_tracing")
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_metrics")
@patch("litellm.integrations.opentelemetry.OpenTelemetry._init_logs")
async def test_langfuse_otel_dynamic_headers_include_v4_ingestion_header(
self, mock_logs, mock_metrics, mock_tracing
):
logger = LangfuseOtelLogger()
headers = logger.construct_dynamic_otel_headers(
StandardCallbackDynamicParams(
langfuse_public_key="pk-lf-dynamic",
langfuse_secret_key="sk-lf-dynamic",
)
)
self.assertIsNotNone(headers)
self.assertEqual(
headers["x-langfuse-ingestion-version"],
"4",
)
self.assertTrue(headers["Authorization"].startswith("Basic "))
if __name__ == "__main__":
unittest.main()

View file

@ -2,157 +2,10 @@ import os
# What this tests?
## Tests /spend endpoints.
import pytest, uuid, json
import asyncio
import pytest
import aiohttp
async def generate_key(session, models=[], team_id=None):
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {
"models": models,
"duration": None,
}
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(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-3.5-turbo"):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": model,
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": f"Hello! {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_spend_logs(session, request_id=None, api_key=None):
if api_key is not None:
url = f"http://0.0.0.0:4000/spend/logs?api_key={api_key}"
else:
url = f"http://0.0.0.0:4000/spend/logs?request_id={request_id}"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_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()
async def generate_org(session: aiohttp.ClientSession) -> dict:
"""
Generate a new organization using the API.
Args:
session: aiohttp client session
Returns:
dict: Response containing org_id
"""
url = "http://0.0.0.0:4000/organization/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
request_body = {
"organization_alias": f"test-org-{uuid.uuid4()}",
}
async with session.post(url, headers=headers, json=request_body) as response:
return await response.json()
async def generate_team(session: aiohttp.ClientSession, org_id: str) -> dict:
"""
Generate a new team within an organization using the API.
Args:
session: aiohttp client session
org_id: Organization ID to create the team in
Returns:
dict: Response containing team_id
"""
url = "http://0.0.0.0:4000/team/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {"organization_id": org_id}
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
@pytest.mark.skip(
reason="Flaky in CI: /spend/logs?request_id=... returns 500 even after a 20s wait for the spend log to be written. Spend-log accuracy is covered by tests/unit/proxy/spend_tracking/ and tests/integration/spend/."
)
@pytest.mark.asyncio
async def test_spend_logs_with_org_id():
"""
- Create Organization
- Create Team in organization
- Create Key in organization
- Make call (makes sure it's in spend logs)
- Get request id from logs
- Assert spend logs have correct org_id and team_id
"""
async with aiohttp.ClientSession() as session:
org_gen = await generate_org(session=session)
print("org_gen: ", json.dumps(org_gen, indent=4, default=str))
org_id = org_gen["organization_id"]
team_gen = await generate_team(session=session, org_id=org_id)
print("team_gen: ", json.dumps(team_gen, indent=4, default=str))
team_id = team_gen["team_id"]
key_gen = await generate_key(session=session, team_id=team_id)
print("key_gen: ", json.dumps(key_gen, indent=4, default=str))
key = key_gen["key"]
response = await chat_completion(session=session, key=key)
await asyncio.sleep(20)
spend_logs_response = await get_spend_logs(
session=session, request_id=response["id"]
)
print(
"spend_logs_response: ",
json.dumps(spend_logs_response, indent=4, default=str),
)
spend_logs_response = spend_logs_response[0]
assert spend_logs_response["metadata"]["user_api_key_org_id"] == org_id
assert spend_logs_response["metadata"]["user_api_key_team_id"] == team_id
assert spend_logs_response["team_id"] == team_id
async def get_predict_spend_logs(session):
url = "http://0.0.0.0:4000/global/predict/spend/logs"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}

View file

@ -1,912 +0,0 @@
import os
# What this tests ?
## Tests /team endpoints.
import pytest
import asyncio
import aiohttp
import time, uuid
from openai import AsyncOpenAI
from typing import Optional
import openai
from unittest.mock import MagicMock, patch
async def get_user_info(session, get_user, call_user, view_all: Optional[bool] = None):
"""
Make sure only models user has access to are returned
"""
if view_all is True:
url = "http://localhost:4000/user/info"
else:
url = f"http://localhost:4000/user/info?user_id={get_user}"
headers = {
"Authorization": f"Bearer {call_user}",
"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:
if call_user != get_user:
return status
else:
print(f"call_user: {call_user}; get_user: {get_user}")
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
async def wait_for_team_member_spend_update(
session, user_id, team_id, expected_min_spend, max_wait=10
):
"""
Wait for the team member spend update to be committed to the database.
Polls the user info endpoint until the spend is updated.
This is needed because spend updates are queued asynchronously and committed periodically.
"""
start_time = time.time()
initial_spend = None
while time.time() - start_time < max_wait:
try:
user_info = await get_user_info(session, user_id, call_user=os.environ["LITELLM_MASTER_KEY"])
if user_info.get("teams"):
for team in user_info["teams"]:
if team.get("team_id") == team_id:
for membership in team.get("team_memberships", []):
spend = membership.get("spend", 0.0)
if initial_spend is None:
initial_spend = spend
print(f"Initial team member spend: {spend}")
if spend >= expected_min_spend:
print(
f"[OK] Team member spend updated: {spend} >= {expected_min_spend}"
)
return True
print(
f"[WAITING] Team member spend: {spend}, expected >= {expected_min_spend}, elapsed: {time.time() - start_time:.1f}s"
)
await asyncio.sleep(0.5)
except Exception as e:
print(f"Error checking team member spend: {e}")
await asyncio.sleep(0.5)
print(
f"[TIMEOUT] Timeout waiting for team member spend update (expected >= {expected_min_spend})"
)
return False
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://localhost:4000/user/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "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 add_member(
session, i, team_id, user_id=None, user_email=None, max_budget=None, members=None
):
url = "http://localhost:4000/team/member_add"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {"team_id": team_id, "member": {"role": "user"}}
if user_email is not None:
data["member"]["user_email"] = user_email
elif user_id is not None:
data["member"]["user_id"] = user_id
elif members is not None:
data["member"] = members
if max_budget is not None:
data["max_budget_in_team"] = max_budget
print("sent data: {}".format(data))
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"ADD MEMBER 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(
session,
i,
team_id,
user_id=None,
user_email=None,
max_budget=None,
):
url = "http://localhost:4000/team/member_update"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {"team_id": team_id}
if user_id is not None:
data["user_id"] = user_id
elif user_email is not None:
data["user_email"] = user_email
if max_budget is not None:
data["max_budget_in_team"] = max_budget
print("sent data: {}".format(data))
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(f"ADD MEMBER 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 delete_member(session, i, team_id, user_id=None, user_email=None):
url = "http://localhost:4000/team/member_delete"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {"team_id": team_id}
if user_id is not None:
data["user_id"] = user_id
elif user_email is not None:
data["user_email"] = user_email
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 generate_key(
session,
i,
budget=None,
budget_duration=None,
models=["azure-models", "gpt-4", "dall-e-3"],
team_id=None,
):
url = "http://localhost:4000/key/generate"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {
"models": models,
"duration": None,
"max_budget": budget,
"budget_duration": budget_duration,
}
if team_id is not None:
data["team_id"] = team_id
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()
async def chat_completion(session, key, model="gpt-4"):
url = "http://localhost:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": model,
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
],
}
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(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 new_team(session, i, user_id=None, member_list=None, model_aliases=None):
import json
url = "http://localhost:4000/team/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {"team_alias": "my-new-team"}
if user_id is not None:
data["members_with_roles"] = [{"role": "user", "user_id": user_id}]
elif member_list is not None:
data["members_with_roles"] = member_list
if model_aliases is not None:
data["model_aliases"] = model_aliases
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()
async def update_team(session, i, team_id, user_id=None, member_list=None, **kwargs):
url = "http://localhost:4000/team/update"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {"team_id": team_id, **kwargs}
if user_id is not None:
data["members_with_roles"] = [{"role": "user", "user_id": user_id}]
elif member_list is not None:
data["members_with_roles"] = member_list
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 delete_team(
session,
i,
team_id,
):
url = "http://localhost:4000/team/delete"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {
"team_ids": [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}")
return await response.json()
async def list_teams(
session,
i,
):
url = "http://localhost:4000/team/list"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
async with session.get(url, headers=headers) as response:
status = response.status
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_team_new():
"""
Make 20 parallel calls to /user/new. Assert all worked.
"""
user_id = f"{uuid.uuid4()}"
async with aiohttp.ClientSession() as session:
new_user(session=session, i=0, user_id=user_id)
tasks = [new_team(session, i, user_id=user_id) for i in range(1, 11)]
await asyncio.gather(*tasks)
async def get_team_info(session, get_team, call_key):
url = f"http://localhost:4000/team/info?team_id={get_team}"
headers = {
"Authorization": f"Bearer {call_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 == 404:
raise openai.NotFoundError(
message="404 received", response=MagicMock(), body=None
)
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_team_info():
"""
Scenario 1:
- test with admin key -> expect to work
Scenario 2:
- test with team key -> expect to work
Scenario 3:
- test with non-team key -> expect to fail
"""
async with aiohttp.ClientSession() as session:
"""
Scenario 1 - as admin
"""
new_team_data = await new_team(
session,
0,
)
team_id = new_team_data["team_id"]
## as admin ##
await get_team_info(session=session, get_team=team_id, call_key=os.environ["LITELLM_MASTER_KEY"])
"""
Scenario 2 - as team key
"""
key_gen = await generate_key(session=session, i=0, team_id=team_id)
key = key_gen["key"]
await get_team_info(session=session, get_team=team_id, call_key=key)
"""
Scenario 3 - as non-team key
"""
key_gen = await generate_key(session=session, i=0)
key = key_gen["key"]
try:
await get_team_info(session=session, get_team=team_id, call_key=key)
pytest.fail("Expected call to fail")
except Exception as e:
pass
"""
- Create team
- Add user (user exists in db)
- Update team
- Check if it works
"""
"""
- Create team
- Add user (user doesn't exist in db)
- Update team
- Check if it works
"""
@pytest.mark.asyncio
async def test_team_update_sc_2():
"""
- Create team
- Add 3 users (doesn't exist in db)
- Change team alias
- Check if it works
- Assert team object unchanged besides team alias
"""
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 team with 1 admin and 1 user
member_list = [
{"role": "admin", "user_id": admin_user},
]
team_data = await new_team(session=session, i=0, member_list=member_list)
## Create 10 normal users
members = [
{"role": "user", "user_id": f"krrish_{uuid.uuid4()}@berri.ai"}
for _ in range(10)
]
await add_member(
session=session, i=0, team_id=team_data["team_id"], members=members
)
## ASSERT TEAM SIZE
team_info = await get_team_info(
session=session, get_team=team_data["team_id"], call_key=os.environ["LITELLM_MASTER_KEY"]
)
assert len(team_info["team_info"]["members_with_roles"]) == 12
## CHANGE TEAM ALIAS
new_team_data = await update_team(
session=session, i=0, team_id=team_data["team_id"], team_alias="my-new-team"
)
assert new_team_data["data"]["team_alias"] == "my-new-team"
print(f"team_data: {team_data}")
## assert rest of object is the same
for k, v in new_team_data["data"].items():
if k == "members_with_roles":
assert len(new_team_data["data"][k]) == len(
team_info["team_info"]["members_with_roles"]
)
elif (
k == "created_at"
or k == "updated_at"
or k == "model_spend"
or k == "model_max_budget"
or k == "model_id"
or k == "litellm_organization_table"
or k == "object_permission_id"
or k == "object_permission"
or k == "litellm_model_table"
or k == "policies"
or k == "allow_team_guardrail_config"
or k == "projects"
):
pass
else:
assert new_team_data["data"][k] == team_data[k]
@pytest.mark.asyncio
async def test_team_member_add_email():
from tests.test_users import get_user_info
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 team with 1 admin and 1 user
member_list = [
{"role": "admin", "user_id": admin_user},
]
team_data = await new_team(session=session, i=0, member_list=member_list)
## Add 1 user via email
user_email = "krrish{}@berri.ai".format(uuid.uuid4())
new_user_info = await new_user(session=session, i=0, user_email=user_email)
new_member = {"role": "user", "user_email": user_email}
await add_member(
session=session, i=0, team_id=team_data["team_id"], members=[new_member]
)
## check user info to confirm user is in team
updated_user_info = await get_user_info(
session=session, get_user=new_user_info["user_id"], call_user=os.environ["LITELLM_MASTER_KEY"]
)
print(updated_user_info)
## check if team in user table
is_team_in_list: bool = False
for team in updated_user_info["teams"]:
if team_data["team_id"] == team["team_id"]:
is_team_in_list = True
assert is_team_in_list
@pytest.mark.asyncio
async def test_team_delete():
"""
- Create team
- Create key for team
- Check if key works
- Delete team
"""
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)
## ASSERT USER MEMBERSHIP IS CREATED
user_info = await get_user_info(
session=session, get_user=normal_user, call_user=os.environ["LITELLM_MASTER_KEY"]
)
assert len(user_info["teams"]) == 1
## Create key
key_gen = await generate_key(session=session, i=0, team_id=team_data["team_id"])
key = key_gen["key"]
## Test key
# response = await chat_completion(session=session, key=key)
## Delete team
await delete_team(session=session, i=0, team_id=team_data["team_id"])
## ASSERT USER MEMBERSHIP IS DELETED
user_info = await get_user_info(
session=session, get_user=normal_user, call_user=os.environ["LITELLM_MASTER_KEY"]
)
assert len(user_info["teams"]) == 0
## ASSERT TEAM INFO NOW RETURNS A 404
with pytest.raises(openai.NotFoundError):
await get_team_info(
session=session, get_team=team_data["team_id"], call_key=os.environ["LITELLM_MASTER_KEY"]
)
@pytest.mark.parametrize("dimension", ["user_id", "user_email"])
@pytest.mark.asyncio
async def test_member_delete(dimension):
"""
- Create team
- Add member
- Get team info (check if member in team)
- Delete member
- Get team info (check if member in team)
"""
async with aiohttp.ClientSession() as session:
# Create Team
## 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()}"
normal_user_email = "{}@berri.ai".format(normal_user)
print(f"normal_user: {normal_user}")
await new_user(
session=session, i=0, user_id=normal_user, user_email=normal_user_email
)
## Create team with 1 admin and 1 user
member_list = [
{"role": "admin", "user_id": admin_user},
]
if dimension == "user_id":
member_list.append({"role": "user", "user_id": normal_user})
elif dimension == "user_email":
member_list.append({"role": "user", "user_email": normal_user_email})
team_data = await new_team(session=session, i=0, member_list=member_list)
user_in_team = False
for member in team_data["members_with_roles"]:
if dimension == "user_id" and member["user_id"] == normal_user:
user_in_team = True
elif (
dimension == "user_email" and member["user_email"] == normal_user_email
):
user_in_team = True
assert (
user_in_team is True
), "User not in team. Team list={}, User details - id={}, email={}. Dimension={}".format(
team_data["members_with_roles"], normal_user, normal_user_email, dimension
)
# Delete member
if dimension == "user_id":
updated_team_data = await delete_member(
session=session, i=0, team_id=team_data["team_id"], user_id=normal_user
)
elif dimension == "user_email":
updated_team_data = await delete_member(
session=session,
i=0,
team_id=team_data["team_id"],
user_email=normal_user_email,
)
print(f"updated_team_data: {updated_team_data}")
user_in_team = False
for member in team_data["members_with_roles"]:
if dimension == "user_id" and member["user_id"] == normal_user:
user_in_team = True
elif (
dimension == "user_email" and member["user_email"] == normal_user_email
):
user_in_team = True
assert user_in_team is True
@pytest.mark.asyncio
async def test_users_in_team_budget():
"""
- Create User
- Create Team with User
- Add User to team with budget = 0.0000001
- Make Call 1 -> pass
- Make Call 2 -> fail
"""
get_user = f"krrish_{time.time()}@berri.ai"
async with aiohttp.ClientSession() as session:
# IMPORTANT: Create team first, then create user with team_id.
# This order is critical for the test to work correctly:
# - When a user is created with team_id, the API key gets team_id set from the start
# - This ensures spend tracking and budget enforcement work correctly
# - If we create the user first (without team_id) and then add them to a team,
# the key's team_id remains None, breaking team budget tracking
# DO NOT change this order - it's testing the intended flow where keys are
# associated with teams at creation time.
team = await new_team(session, 0, user_id=None)
print(f"[DEBUG] Created team: {team['team_id']}")
print(f"[DEBUG] Full team data: {team}")
# Create user with team_id so the key is associated with the team from the start
key_gen = await new_user(
session,
0,
user_id=get_user,
budget=10,
budget_duration="5s",
models=["fake-openai-endpoint"],
team_id=team["team_id"],
)
key = key_gen["key"]
print(f"[DEBUG] Created user '{get_user}' with key: {key}")
print(f"[DEBUG] User budget: 10, budget_duration: 5s")
print(f"[DEBUG] Key team_id: {team['team_id']}")
# Check user info BEFORE updating member budget
user_info_before = await get_user_info(session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"])
print(f"[DEBUG] User info BEFORE update_member:")
print(f" - User budget: {user_info_before.get('max_budget')}")
print(f" - User spend: {user_info_before.get('spend')}")
if user_info_before.get("teams"):
for team_info in user_info_before["teams"]:
if team_info.get("team_id") == team["team_id"]:
print(f" - Team memberships: {team_info.get('team_memberships')}")
# update user to have budget = 0.0000001
update_result = await update_member(
session, 0, team_id=team["team_id"], user_id=get_user, max_budget=0.0000001
)
print(f"[DEBUG] Updated member budget to 0.0000001")
print(f"[DEBUG] Update result: {update_result}")
# Check user info AFTER updating member budget
user_info_after = await get_user_info(session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"])
print(f"[DEBUG] User info AFTER update_member:")
print(f" - User budget: {user_info_after.get('max_budget')}")
print(f" - User spend: {user_info_after.get('spend')}")
if user_info_after.get("teams"):
for team_info in user_info_after["teams"]:
if team_info.get("team_id") == team["team_id"]:
print(f" - Team: {team_info.get('team_id')}")
for membership in team_info.get("team_memberships", []):
print(f" - Membership: {membership}")
if "litellm_budget_table" in membership:
budget_table = membership["litellm_budget_table"]
print(f" - Max budget: {budget_table.get('max_budget')}")
print(f" - Current spend: {membership.get('spend', 0)}")
# Call 1
print("\n[DEBUG] ===== Making Call 1 =====")
result = await chat_completion(session, key, model="fake-openai-endpoint")
print(f"[DEBUG] Call 1 PASSED (expected)")
print(f"[DEBUG] Call 1 result: {result}")
# Extract cost from result if available
if isinstance(result, dict):
usage = result.get("usage", {})
print(f"[DEBUG] Call 1 usage: {usage}")
# Wait for spend to be committed to database before checking budget
# Spend updates are queued asynchronously and committed periodically (every minute),
# so we need to wait for the spend from Call 1 to be persisted
print("\n[DEBUG] ===== Waiting for spend to be committed =====")
print("Waiting for team member spend to be committed to database...")
print(
"Note: Spend updates are flushed periodically, this may take up to 90 seconds..."
)
spend_updated = await wait_for_team_member_spend_update(
session, get_user, team["team_id"], 0.0000001, max_wait=90
)
if not spend_updated:
pytest.fail(
"Team member spend was not updated within 90s. "
"The spend update queue may not have flushed, or the model may have 0 cost."
)
# Check user info BEFORE Call 2
user_info_before_call2 = await get_user_info(
session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"]
)
print(f"\n[DEBUG] User info BEFORE Call 2:")
print(f" - User budget: {user_info_before_call2.get('max_budget')}")
print(f" - User spend: {user_info_before_call2.get('spend')}")
if user_info_before_call2.get("teams"):
for team_info in user_info_before_call2["teams"]:
if team_info.get("team_id") == team["team_id"]:
print(f" - Team: {team_info.get('team_id')}")
for membership in team_info.get("team_memberships", []):
if "litellm_budget_table" in membership:
budget_table = membership["litellm_budget_table"]
current_spend = membership.get("spend", 0)
max_budget = budget_table.get("max_budget")
print(f" - Max budget in team: {max_budget}")
print(f" - Current spend in team: {current_spend}")
print(
f" - Budget remaining: {max_budget - current_spend}"
)
print(f" - Should fail?: {current_spend >= max_budget}")
# Call 2
print("\n[DEBUG] ===== Making Call 2 =====")
call2_failed = False
call2_error = None
call2_status = None
try:
# Capture the response to check status code
url = "http://localhost:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": "fake-openai-endpoint",
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello!"},
],
}
async with session.post(url, headers=headers, json=data) as response:
call2_status = response.status
response_text = await response.text()
print(f"[DEBUG] Call 2 status code: {call2_status}")
print(f"[DEBUG] Call 2 response: {response_text}")
if call2_status != 200:
call2_failed = True
call2_error = f"Status {call2_status}: {response_text}"
raise Exception(call2_error)
else:
# Call succeeded when it should have failed
print(f"[ERROR] Call 2 PASSED when it should have FAILED!")
print(f"[ERROR] Response was 200 OK")
except Exception as e:
if call2_failed:
print(f"[DEBUG] Call 2 FAILED (expected): {e}")
print(f"[DEBUG] Checking if error message indicates budget exceeded...")
else:
call2_error = str(e)
print(f"[DEBUG] Call 2 raised exception: {e}")
# Check user info AFTER Call 2
user_info_after_call2 = await get_user_info(
session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"]
)
print(f"\n[DEBUG] User info AFTER Call 2:")
print(f" - User budget: {user_info_after_call2.get('max_budget')}")
print(f" - User spend: {user_info_after_call2.get('spend')}")
if user_info_after_call2.get("teams"):
for team_info in user_info_after_call2["teams"]:
if team_info.get("team_id") == team["team_id"]:
print(f" - Team: {team_info.get('team_id')}")
for membership in team_info.get("team_memberships", []):
if "litellm_budget_table" in membership:
budget_table = membership["litellm_budget_table"]
print(f" - Max budget: {budget_table.get('max_budget')}")
print(f" - Current spend: {membership.get('spend', 0)}")
# Assert Call 2 failed
if not call2_failed:
error_msg = (
f"\n[FAILURE] Call 2 should have failed but it passed!\n"
f"Expected: Budget enforcement to block the call\n"
f"Actual: Call returned status {call2_status}\n"
f"Team member budget: 0.0000001\n"
f"User budget: {user_info_before_call2.get('max_budget')}\n"
f"User spend before call: {user_info_before_call2.get('spend')}\n"
)
# Add team member info if available
if user_info_before_call2.get("teams"):
for team_info in user_info_before_call2["teams"]:
if team_info.get("team_id") == team["team_id"]:
for membership in team_info.get("team_memberships", []):
if "litellm_budget_table" in membership:
error_msg += f"Team member spend before call: {membership.get('spend', 0)}\n"
error_msg += f"Team member max budget: {membership['litellm_budget_table'].get('max_budget')}\n"
pytest.fail(error_msg)
# Check the error message contains budget exceeded
if call2_error and "Budget has been exceeded" not in call2_error:
pytest.fail(
f"Call 2 failed but not with expected error message.\n"
f"Expected error to contain: 'Budget has been exceeded'\n"
f"Actual error: {call2_error}"
)
print("[DEBUG] Call 2 failed as expected with budget exceeded error")
## Check user info
user_info = await get_user_info(session, get_user, call_user=os.environ["LITELLM_MASTER_KEY"])
assert (
user_info["teams"][0]["team_memberships"][0]["litellm_budget_table"][
"max_budget"
]
== 0.0000001
)

View file

@ -1,190 +0,0 @@
import os
import pytest
import requests
import time
from typing import Dict, List
import logging
from litellm._uuid import uuid
# Configure logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class TeamAPI:
def __init__(self, base_url: str, auth_token: str):
self.base_url = base_url
self.headers = {
"Authorization": f"Bearer {auth_token}",
"Content-Type": "application/json",
}
def create_team(self, team_alias: str, models: List[str] = None) -> Dict:
"""Create a new team"""
# Generate a unique team_id using uuid
team_id = f"test_team_{uuid.uuid4().hex[:8]}"
data = {
"team_id": team_id,
"team_alias": team_alias,
"models": models or ["o3-mini"],
}
response = requests.post(
f"{self.base_url}/team/new", headers=self.headers, json=data
)
response.raise_for_status()
logger.info(f"Created new team: {team_id}")
return response.json(), team_id
def get_team_info(self, team_id: str) -> Dict:
"""Get current team information"""
response = requests.get(
f"{self.base_url}/team/info",
headers=self.headers,
params={"team_id": team_id},
)
response.raise_for_status()
return response.json()
def add_team_member(self, team_id: str, user_email: str, role: str) -> Dict:
"""Add a single team member"""
data = {"team_id": team_id, "member": [{"role": role, "user_id": user_email}]}
response = requests.post(
f"{self.base_url}/team/member_add", headers=self.headers, json=data
)
response.raise_for_status()
return response.json()
def delete_team_member(self, team_id: str, user_id: str) -> Dict:
"""Delete a team member
Args:
team_id (str): ID of the team
user_id (str): User ID to remove from team
Returns:
Dict: Response from the API
"""
data = {"team_id": team_id, "user_id": user_id}
response = requests.post(
f"{self.base_url}/team/member_delete", headers=self.headers, json=data
)
response.raise_for_status()
return response.json()
@pytest.fixture
def api_client():
"""Fixture for TeamAPI client"""
base_url = "http://localhost:4000"
auth_token = os.environ["LITELLM_MASTER_KEY"] # Replace with your token
return TeamAPI(base_url, auth_token)
@pytest.fixture
def new_team(api_client):
"""Fixture that creates a new team for each test"""
team_alias = f"Test Team {uuid.uuid4().hex[:6]}"
team_response, team_id = api_client.create_team(team_alias)
logger.info(f"Created test team: {team_id} ({team_alias})")
return team_id
def verify_member_in_team(team_info: Dict, user_email: str) -> bool:
"""Verify if a member exists in team"""
return any(
member["user_id"] == user_email
for member in team_info["team_info"]["members_with_roles"]
)
def test_add_single_member(api_client, new_team):
"""Test adding a single member to a new team"""
# Get initial team info
initial_info = api_client.get_team_info(new_team)
initial_size = len(initial_info["team_info"]["members_with_roles"])
# Add new member
test_email = f"pytest_user_{uuid.uuid4().hex[:6]}@mycompany.com"
api_client.add_team_member(new_team, test_email, "user")
# Allow time for system to process
time.sleep(1)
# Verify addition
updated_info = api_client.get_team_info(new_team)
updated_size = len(updated_info["team_info"]["members_with_roles"])
# Assertions
assert verify_member_in_team(
updated_info, test_email
), f"Member {test_email} not found in team"
assert (
updated_size == initial_size + 1
), f"Team size did not increase by 1 (was {initial_size}, now {updated_size})"
def test_member_deletion(api_client, new_team):
"""Test that member deletion works correctly and removes all instances of a user"""
# Add a test user
user_id = f"pytest_user_{uuid.uuid4().hex[:6]}"
api_client.add_team_member(new_team, user_id, "user")
time.sleep(1)
# Verify user was added
team_info_before = api_client.get_team_info(new_team)
assert verify_member_in_team(
team_info_before, user_id
), "User was not added successfully"
initial_size = len(team_info_before["team_info"]["members_with_roles"])
# Attempt to delete the same user multiple times (5 times)
for i in range(5):
logger.info(f"Attempting deletion {i+1}/5")
if i == 0:
# First deletion should succeed
api_client.delete_team_member(new_team, user_id)
time.sleep(1)
else:
# Subsequent deletions should raise an error
try:
api_client.delete_team_member(new_team, user_id)
pytest.fail("Expected HTTPError for duplicate deletion")
except requests.exceptions.HTTPError as e:
logger.info(
f"Expected error received on deletion attempt {i+1}: {str(e)}"
)
# Verify final state
final_info = api_client.get_team_info(new_team)
final_size = len(final_info["team_info"]["members_with_roles"])
# Verify user is completely removed
assert not verify_member_in_team(
final_info, user_id
), "User still exists in team after deletion"
# Verify only one member was removed
assert (
final_size == initial_size - 1
), f"Team size changed unexpectedly (was {initial_size}, now {final_size})"
def test_delete_nonexistent_member(api_client, new_team):
"""Test that attempting to delete a nonexistent member raises appropriate error"""
nonexistent_user = f"nonexistent_{uuid.uuid4().hex[:6]}"
# Verify user doesn't exist first
team_info = api_client.get_team_info(new_team)
assert not verify_member_in_team(
team_info, nonexistent_user
), "Test setup error: nonexistent user somehow exists"
# Attempt to delete nonexistent user
with pytest.raises(requests.exceptions.HTTPError) as exc_info:
api_client.delete_team_member(new_team, nonexistent_user)
e = exc_info.value
logger.info(f"Expected error received: {str(e)}")
assert e.response.status_code == 400

View file

@ -5,10 +5,7 @@ import pytest
import asyncio
import aiohttp
import time
from openai import AsyncOpenAI
from tests.test_team import list_teams
from typing import Optional
from fastapi import HTTPException
async def new_user(
@ -41,16 +38,6 @@ async def new_user(
return await response.json()
@pytest.mark.asyncio
async def test_user_new():
"""
Make 20 parallel calls to /user/new. Assert all worked.
"""
async with aiohttp.ClientSession() as session:
tasks = [new_user(session, i) for i in range(1, 11)]
await asyncio.gather(*tasks)
async def get_user_info(session, get_user, call_user, view_all: Optional[bool] = None):
"""
Make sure only models user has access to are returned
@ -79,37 +66,6 @@ async def get_user_info(session, get_user, call_user, view_all: Optional[bool] =
return await response.json()
@pytest.mark.asyncio
async def test_user_info():
"""
Get user info
- as admin
- as user themself
- as random
"""
get_user = f"krrish_{time.time()}@berri.ai"
async with aiohttp.ClientSession() as session:
key_gen = await new_user(session, 0, user_id=get_user)
key = key_gen["key"]
## as admin ##
resp = await get_user_info(
session=session, get_user=get_user, call_user=os.environ["LITELLM_MASTER_KEY"]
)
assert isinstance(resp["user_info"], dict)
assert len(resp["user_info"]) > 0
## as user themself ##
resp = await get_user_info(session=session, get_user=get_user, call_user=key)
assert isinstance(resp["user_info"], dict)
assert len(resp["user_info"]) > 0
# as random user #
key_gen = await new_user(session=session, i=0)
random_key = key_gen["key"]
status = await get_user_info(
session=session, get_user=get_user, call_user=random_key
)
assert status == 403
@pytest.mark.skip(reason="Frequent check on ci/cd leads to read timeout issue.")
@pytest.mark.asyncio
async def test_users_budgets_reset():
@ -145,168 +101,6 @@ async def test_users_budgets_reset():
assert reset_at_init_value != reset_at_new_value
async def chat_completion(session, key, model="gpt-4"):
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000")
messages = [
{"role": "system", "content": "You are a helpful assistant"},
{"role": "user", "content": f"Hello! {time.time()}"},
]
data = {
"model": model,
"messages": messages,
}
response = await client.chat.completions.create(**data)
async def chat_completion_streaming(session, key, model="gpt-4"):
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000")
messages = [
{"role": "system", "content": "You are a helpful assistant"},
{"role": "user", "content": f"Hello! {time.time()}"},
]
data = {"model": model, "messages": messages, "stream": True}
response = await client.chat.completions.create(**data)
async for chunk in response:
continue
import json
from litellm._uuid import uuid
import pytest
from typing import Dict, Tuple
async def setup_test_users(session: aiohttp.ClientSession) -> Tuple[Dict, Dict]:
"""
Create two test users and an additional key for the first user.
Returns tuple of (user1_data, user2_data) where each contains user info and keys.
"""
# Create two test users
user1 = await new_user(
session=session,
i=0,
budget=100,
budget_duration="30d",
models=["anthropic.claude-haiku-4-5-20251001-v1:0"],
)
user2 = await new_user(
session=session,
i=1,
budget=100,
budget_duration="30d",
models=["anthropic.claude-haiku-4-5-20251001-v1:0"],
)
print("\nCreated two test users:")
print(f"User 1 ID: {user1['user_id']}")
print(f"User 2 ID: {user2['user_id']}")
# Create an additional key for user1
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {user1['key']}",
}
key_payload = {
"user_id": user1["user_id"],
"duration": "7d",
"key_alias": f"test_key_{uuid.uuid4()}",
"models": ["anthropic.claude-haiku-4-5-20251001-v1:0"],
}
print("\nGenerating additional key for user1...")
key_response = await session.post(
f"http://0.0.0.0:4000/key/generate", headers=headers, json=key_payload
)
assert key_response.status == 200, "Failed to generate additional key for user1"
user1_additional_key = await key_response.json()
print(f"\nGenerated key details:")
print(json.dumps(user1_additional_key, indent=2))
# Return both users' data including the additional key
return {
"user_data": user1,
"additional_key": user1_additional_key,
"headers": headers,
}, {
"user_data": user2,
"headers": {
"Content-Type": "application/json",
"Authorization": f"Bearer {user2['key']}",
},
}
async def print_response_details(response: aiohttp.ClientResponse) -> None:
"""Helper function to print response details"""
print("\nResponse Details:")
print(f"Status Code: {response.status}")
print("\nResponse Content:")
try:
formatted_json = json.dumps(await response.json(), indent=2)
print(formatted_json)
except json.JSONDecodeError:
print(await response.text())
@pytest.mark.asyncio
async def test_key_update_user_isolation():
"""Test that a user cannot update a key that belongs to another user"""
async with aiohttp.ClientSession() as session:
user1_data, user2_data = await setup_test_users(session)
# Try to update the key to belong to user2
update_payload = {
"key": user1_data["additional_key"]["key"],
"user_id": user2_data["user_data"][
"user_id"
], # Attempting to change ownership
"metadata": {"purpose": "testing_user_isolation", "environment": "test"},
}
print("\nAttempting to update key ownership to user2...")
update_response = await session.post(
f"http://0.0.0.0:4000/key/update",
headers=user1_data["headers"], # Using user1's headers
json=update_payload,
)
await print_response_details(update_response)
# Verify update attempt was rejected
assert (
update_response.status == 403
), "Request should have been rejected with 403 status code"
@pytest.mark.asyncio
async def test_key_delete_user_isolation():
"""Test that a user cannot delete a key that belongs to another user"""
async with aiohttp.ClientSession() as session:
user1_data, user2_data = await setup_test_users(session)
# Try to delete user1's additional key using user2's credentials
delete_payload = {
"keys": [user1_data["additional_key"]["key"]],
}
print("\nAttempting to delete user1's key using user2's credentials...")
delete_response = await session.post(
f"http://0.0.0.0:4000/key/delete",
headers=user2_data["headers"],
json=delete_payload,
)
await print_response_details(delete_response)
# Verify delete attempt was rejected
assert (
delete_response.status == 403
), "Request should have been rejected with 403 status code"

View file

@ -1,12 +1,17 @@
import base64
import json
import os
from datetime import datetime, timezone
from typing import Final
from unittest.mock import MagicMock, patch
import pytest
from opentelemetry.sdk.trace import TracerProvider
from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger
from litellm.integrations.opentelemetry import OpenTelemetryConfig
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import StandardCallbackDynamicParams
class TestLangfuseOtelIntegration:
@ -1006,5 +1011,58 @@ class TestLangfuseOtelResponsesAPI:
assert output_data[0]["arguments"] == {}
if __name__ == "__main__":
pytest.main([__file__])
LANGFUSE_ENV_ONLY: Final = {
key: value
for key, value in os.environ.items()
if not key.startswith(("LANGFUSE_", "OTEL_"))
}
def _decoded_basic_auth(header: str) -> str:
assert header.startswith("Basic ")
return base64.b64decode(header.removeprefix("Basic ")).decode()
def test_langfuse_otel_env_config_headers_carry_v4_ingestion_and_basic_auth() -> None:
with patch.dict(
os.environ, {**LANGFUSE_ENV_ONLY, "LANGFUSE_PUBLIC_KEY": "pk-lf-123", "LANGFUSE_SECRET_KEY": "sk-lf-123"}, clear=True
):
logger: Final = LangfuseOtelLogger()
headers: Final = OpenTelemetry._get_headers_dictionary(logger.config.headers)
assert headers["x-langfuse-ingestion-version"] == "4"
assert _decoded_basic_auth(headers["Authorization"]) == "pk-lf-123:sk-lf-123"
def test_langfuse_otel_dynamic_headers_carry_v4_ingestion_and_basic_auth() -> None:
with patch.dict(os.environ, LANGFUSE_ENV_ONLY, clear=True):
logger: Final = LangfuseOtelLogger()
headers: Final = logger.construct_dynamic_otel_headers(
StandardCallbackDynamicParams(langfuse_public_key="pk-lf-dynamic", langfuse_secret_key="sk-lf-dynamic")
)
assert headers is not None
assert headers["x-langfuse-ingestion-version"] == "4"
assert _decoded_basic_auth(headers["Authorization"]) == "pk-lf-dynamic:sk-lf-dynamic"
def test_langfuse_otel_does_not_start_proxy_request_span() -> None:
langfuse_provider: Final = TracerProvider()
generic_provider: Final = TracerProvider()
with patch.dict(os.environ, LANGFUSE_ENV_ONLY, clear=True):
langfuse_logger: Final = LangfuseOtelLogger(tracer_provider=langfuse_provider)
generic_logger: Final = OpenTelemetry(
config=OpenTelemetryConfig(exporter="console", skip_set_global=True), tracer_provider=generic_provider
)
started_at: Final = datetime(2026, 1, 1, tzinfo=timezone.utc)
request_headers: Final = {"Authorization": "Bearer test"}
try:
assert (
langfuse_logger.create_litellm_proxy_request_started_span(start_time=started_at, headers=request_headers)
is None
)
assert (
generic_logger.create_litellm_proxy_request_started_span(start_time=started_at, headers=request_headers)
is not None
)
finally:
langfuse_provider.shutdown()
generic_provider.shutdown()

View file

@ -0,0 +1,40 @@
import importlib
import os
from collections.abc import Iterator
from pathlib import Path
from typing import Final
import pytest
import litellm.litellm_core_utils.default_encoding as default_encoding
BUNDLED_TOKENIZERS: Final = Path(default_encoding.filename)
@pytest.fixture
def reload_default_encoding(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
monkeypatch.delenv("TIKTOKEN_CACHE_DIR", raising=False)
monkeypatch.delenv("CUSTOM_TIKTOKEN_CACHE_DIR", raising=False)
yield
monkeypatch.delenv("TIKTOKEN_CACHE_DIR", raising=False)
monkeypatch.delenv("CUSTOM_TIKTOKEN_CACHE_DIR", raising=False)
importlib.reload(default_encoding)
@pytest.mark.usefixtures("reload_default_encoding")
def test_tiktoken_cache_dir_defaults_to_bundled_tokenizers_for_non_root(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_NON_ROOT", "true")
importlib.reload(default_encoding)
assert Path(os.environ["TIKTOKEN_CACHE_DIR"]) == BUNDLED_TOKENIZERS
assert BUNDLED_TOKENIZERS.name == "tokenizers"
assert default_encoding.encoding.name == "cl100k_base"
assert default_encoding.encoding.decode(default_encoding.encoding.encode("hello world")) == "hello world"
@pytest.mark.usefixtures("reload_default_encoding")
def test_custom_tiktoken_cache_dir_overrides_and_is_created(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
custom_dir: Final = tmp_path / "tiktoken_cache"
monkeypatch.setenv("CUSTOM_TIKTOKEN_CACHE_DIR", str(custom_dir))
importlib.reload(default_encoding)
assert os.environ["TIKTOKEN_CACHE_DIR"] == str(custom_dir)
assert custom_dir.is_dir()

View file

@ -7917,3 +7917,54 @@ def test_is_prompt_caching_enabled(anthropic_messages):
custom_llm_provider="anthropic",
model="anthropic/claude-sonnet-4-5-20250929",
)
def test_calculate_usage_sums_compaction_and_message_iterations():
usage: Final = AnthropicConfig().calculate_usage(
usage_object={
"input_tokens": 100,
"output_tokens": 50,
"iterations": [
{"iteration": 1, "type": "compaction", "input_tokens": 1000, "output_tokens": 500},
{"iteration": 2, "type": "message", "input_tokens": 100, "output_tokens": 50},
],
},
reasoning_content=None,
)
assert usage.prompt_tokens == 1100
assert usage.completion_tokens == 550
assert usage.total_tokens == 1650
assert usage.prompt_tokens_details.text_tokens == 1100
assert usage.iterations is not None
assert len(usage.iterations) == 2
assert usage.iterations[0]["type"] == "compaction"
def test_calculate_usage_sums_cache_tokens_across_compaction_iterations():
usage: Final = AnthropicConfig().calculate_usage(
usage_object={
"input_tokens": 100,
"output_tokens": 50,
"iterations": [
{
"type": "compaction",
"input_tokens": 500,
"output_tokens": 200,
"cache_creation_input_tokens": 50,
"cache_read_input_tokens": 17000,
},
{
"type": "message",
"input_tokens": 100,
"output_tokens": 50,
"cache_creation_input_tokens": 10,
"cache_read_input_tokens": 20,
},
],
},
reasoning_content=None,
)
assert usage.prompt_tokens == 17680
assert usage.completion_tokens == 250
assert usage.prompt_tokens_details.cache_creation_tokens == 60
assert usage.prompt_tokens_details.cached_tokens == 17020

View file

@ -12,8 +12,11 @@ from litellm.llms.azure.responses.o_series_transformation import (
AzureOpenAIOSeriesResponsesAPIConfig,
)
from litellm.llms.azure.responses.transformation import AzureOpenAIResponsesAPIConfig
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
@pytest.mark.serial
@ -920,3 +923,25 @@ async def test_azure_responses_api_headers_with_llm_provider_prefix():
# Also verify openai-compatible headers are included
assert "x-ratelimit-limit-tokens" in headers
assert "x-ratelimit-remaining-tokens" in headers
@pytest.mark.parametrize(
"model", ["gpt-5", "gpt-5-turbo", "GPT-5", "azure/gpt-5", "gpt-3.5-turbo", "gpt-4", "gpt-4-turbo", "gpt-4o"]
)
def test_azure_gpt_models_resolve_to_responses_config_with_temperature(model: str) -> None:
config: Final = ProviderConfigManager.get_provider_responses_api_config(provider=LlmProviders.AZURE, model=model)
assert type(config) is AzureOpenAIResponsesAPIConfig
assert "temperature" in config.get_supported_openai_params(model)
@pytest.mark.parametrize("model", ["o1", "o3"])
def test_azure_o_series_resolves_to_o_series_config_without_temperature(model: str) -> None:
config: Final = ProviderConfigManager.get_provider_responses_api_config(provider=LlmProviders.AZURE, model=model)
assert type(config) is AzureOpenAIOSeriesResponsesAPIConfig
assert "temperature" not in config.get_supported_openai_params(model)
def test_openai_gpt5_resolves_to_responses_config_with_temperature() -> None:
config: Final = ProviderConfigManager.get_provider_responses_api_config(provider=LlmProviders.OPENAI, model="gpt-5")
assert type(config) is OpenAIResponsesAPIConfig
assert "temperature" in config.get_supported_openai_params("gpt-5")

View file

@ -1,4 +1,10 @@
import re
from collections.abc import Iterator
from typing import Final
import httpx
import pytest
import respx
import litellm
from litellm.llms.custom_httpx.async_client_cleanup import close_litellm_async_clients
@ -19,3 +25,97 @@ async def test_second_cleanup_pass_does_not_resurrect_owned_client():
litellm.in_memory_llm_clients_cache.cache_dict.pop(cache_key, None)
assert handler._client is original_client
GEMINI_GENERATE_URL: Final = re.compile(r"https://generativelanguage\.googleapis\.com/.*:generateContent.*")
def _gemini_reply(text: str) -> httpx.Response:
return httpx.Response(
200,
json={
"candidates": [
{"content": {"parts": [{"text": text}], "role": "model"}, "finishReason": "STOP", "index": 0}
],
"usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2},
},
)
def _cached_async_handlers() -> list[AsyncHTTPHandler]:
return [
handler
for handler in litellm.in_memory_llm_clients_cache.cache_dict.values()
if isinstance(handler, AsyncHTTPHandler)
]
@pytest.fixture
def gemini_httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setenv("GEMINI_API_KEY", "gemini-cleanup-test")
litellm.in_memory_llm_clients_cache.flush_cache()
yield
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.mark.asyncio
@pytest.mark.usefixtures("gemini_httpx_transport")
async def test_acompletion_client_is_closed_by_cleanup() -> None:
with respx.mock(assert_all_called=True) as mock:
mock.post(GEMINI_GENERATE_URL).mock(return_value=_gemini_reply("Hi there!"))
response: Final = await litellm.acompletion(
model="gemini/gemini-2.0-flash-lite-001",
messages=[{"role": "user", "content": "Hello"}],
)
assert response.choices[0].message.content == "Hi there!"
clients: Final = [handler._client for handler in _cached_async_handlers()]
assert clients
assert not any(client.is_closed for client in clients)
await close_litellm_async_clients()
assert all(client.is_closed for client in clients)
@pytest.mark.asyncio
@pytest.mark.usefixtures("gemini_httpx_transport")
async def test_repeated_acompletion_calls_reuse_one_client_that_cleanup_closes() -> None:
with respx.mock(assert_all_called=True) as mock:
route: Final = mock.post(GEMINI_GENERATE_URL).mock(
side_effect=[_gemini_reply(f"Response {index}") for index in range(3)]
)
for index in range(3):
response = await litellm.acompletion(
model="gemini/gemini-2.0-flash-lite-001",
messages=[{"role": "user", "content": f"Hello {index}"}],
)
assert response.choices[0].message.content == f"Response {index}"
assert route.call_count == 3
handlers: Final = _cached_async_handlers()
assert len(handlers) == 1
client: Final = handlers[0]._client
await close_litellm_async_clients()
assert client.is_closed
@pytest.mark.asyncio
@pytest.mark.usefixtures("gemini_httpx_transport")
async def test_cleanup_is_idempotent_and_acompletion_works_afterwards() -> None:
with respx.mock(assert_all_called=True) as mock:
route: Final = mock.post(GEMINI_GENERATE_URL).mock(side_effect=[_gemini_reply("Hello!"), _gemini_reply("Hi!")])
await litellm.acompletion(
model="gemini/gemini-2.0-flash-lite-001",
messages=[{"role": "user", "content": "Hello"}],
)
for _ in range(3):
await close_litellm_async_clients()
response: Final = await litellm.acompletion(
model="gemini/gemini-2.0-flash-lite-001",
messages=[{"role": "user", "content": "Hello"}],
)
assert response.choices[0].message.content == "Hi!"
assert route.call_count == 2
await close_litellm_async_clients()

View file

@ -0,0 +1,29 @@
from typing import Final
import pytest
from litellm.llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
def test_provider_config_manager_returns_litellm_proxy_responses_config() -> None:
config: Final = ProviderConfigManager.get_provider_responses_api_config(
model="litellm_proxy/gpt-5.5", provider=LlmProviders.LITELLM_PROXY
)
assert isinstance(config, LiteLLMProxyResponsesAPIConfig)
assert config.custom_llm_provider == LlmProviders.LITELLM_PROXY
@pytest.mark.parametrize("api_base", ["https://my-proxy.example.com", "https://my-proxy.example.com/"])
def test_get_complete_url_appends_responses_path(api_base: str) -> None:
assert (
LiteLLMProxyResponsesAPIConfig().get_complete_url(api_base=api_base, litellm_params={})
== "https://my-proxy.example.com/responses"
)
def test_get_complete_url_requires_api_base(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("LITELLM_PROXY_API_BASE", raising=False)
with pytest.raises(ValueError, match="api_base not set"):
LiteLLMProxyResponsesAPIConfig().get_complete_url(api_base=None, litellm_params={})

View file

@ -9,6 +9,7 @@ import json
import os
import re
from collections.abc import Iterable, Sequence
from concurrent.futures import ThreadPoolExecutor
from contextlib import asynccontextmanager
from typing import Final, Literal
from unittest.mock import MagicMock, patch
@ -4799,3 +4800,33 @@ async def test_legacy_pii_masking_config_registers_logging_only_guardrail(monkey
assert pii_masking_obj.should_run_guardrail(
data={}, event_type=GuardrailEventHooks.logging_only
)
async def _one_session(guardrail: OPTIONAL_PresidioPIIMasking) -> aiohttp.ClientSession:
async with guardrail._get_session_iterator() as session:
return session
@pytest.mark.asyncio
async def test_get_session_iterator_reuses_one_session_on_main_thread(
presidio_guardrail: OPTIONAL_PresidioPIIMasking,
) -> None:
sessions: Final = tuple([await _one_session(presidio_guardrail) for _ in range(10)])
assert all(session is sessions[0] for session in sessions)
assert sessions[0] is presidio_guardrail._http_session
await presidio_guardrail._close_http_session()
def test_get_session_iterator_reuses_one_session_per_background_loop(
presidio_guardrail: OPTIONAL_PresidioPIIMasking,
) -> None:
async def collect_and_close() -> tuple[aiohttp.ClientSession, ...]:
collected: Final = tuple([await _one_session(presidio_guardrail) for _ in range(10)])
await collected[0].close()
return collected
with ThreadPoolExecutor(max_workers=1) as pool:
sessions: Final = pool.submit(asyncio.run, collect_and_close()).result()
assert len(sessions) == 10
assert all(session is sessions[0] for session in sessions)
assert presidio_guardrail._http_session is None

View file

@ -516,12 +516,10 @@ def test_e2e_proxy_config_opts_in_to_the_mock_params_its_suite_sends():
for param in GATED_MOCK_PARAM_NAMES
if f"{param}=" in source or f'"{param}"' in source
)
assert senders, "expected the E2E suite to still exercise the gated mock testing params"
config = yaml.safe_load((repo_root / "proxy_server_config.yaml").read_text(encoding="utf-8"))
general_settings = config.get("general_settings") or {}
assert general_settings.get(MOCK_TESTING_CONFIG_KEY) is True, (
assert not senders or general_settings.get(MOCK_TESTING_CONFIG_KEY) is True, (
f"proxy_server_config.yaml must set general_settings.{MOCK_TESTING_CONFIG_KEY}: true — "
f"the E2E suite sends gated mock testing params ({', '.join(sorted(senders))}) "
"and the proxy rejects them with a 400 otherwise"

View file

@ -1,13 +1,14 @@
import asyncio
from datetime import datetime, timedelta
from typing import Final
from unittest.mock import AsyncMock
import pytest
from litellm import Router
from litellm import Router, token_counter
from litellm.caching.dual_cache import DualCache
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage
from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict
from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict, RouterErrors
MODEL_GROUP: Final = "lowest-tpm-router"
HIGH_USAGE_DEPLOYMENT_ID: Final = "highest-usage"
@ -117,3 +118,69 @@ async def test_v2_subclass_overriding_async_get_available_deployments_with_the_o
f"from {HIGH_USAGE_DEPLOYMENT_ID}",
f"from {LOW_USAGE_DEPLOYMENT_ID}",
}
def _rate_limited_router(num_allowed_send: int) -> tuple[Router, tuple[list[dict[str, str]], ...]]:
conversations: Final = tuple(
[{"role": "user", "content": f"{index}. Hey, how's it going?"}] for index in range(num_allowed_send)
)
tpm: Final = sum(token_counter(model="gpt-4o", messages=messages) + 5 for messages in conversations)
deployment: Final = _deployment(LOW_USAGE_DEPLOYMENT_ID)
router: Final = Router(
model_list=[{**deployment, "rpm": num_allowed_send, "tpm": tpm}],
routing_strategy="usage-based-routing",
enable_pre_call_checks=True,
num_retries=0,
)
return router, conversations
def test_usage_based_routing_v1_serves_sync_calls_within_rpm_and_tpm() -> None:
router, conversations = _rate_limited_router(num_allowed_send=3)
responses: Final = [router.completion(model=MODEL_GROUP, messages=messages) for messages in conversations[:2]]
assert [response.choices[0].message.content for response in responses] == [f"from {LOW_USAGE_DEPLOYMENT_ID}"] * 2
@pytest.mark.asyncio
async def test_usage_based_routing_v1_serves_async_calls_within_rpm_and_tpm() -> None:
router, conversations = _rate_limited_router(num_allowed_send=3)
responses: Final = await asyncio.gather(
*(router.acompletion(model=MODEL_GROUP, messages=messages) for messages in conversations[:2])
)
assert [response.choices[0].message.content for response in responses] == [f"from {LOW_USAGE_DEPLOYMENT_ID}"] * 2
RPM_LIMIT: Final = 3
def _router_with_recorded_rpm(recorded: int, enable_pre_call_checks: bool) -> Router:
router: Final = Router(
model_list=[{**_deployment(LOW_USAGE_DEPLOYMENT_ID), "rpm": RPM_LIMIT}],
routing_strategy="usage-based-routing",
enable_pre_call_checks=enable_pre_call_checks,
num_retries=0,
)
now: Final = datetime.now()
for offset in range(-1, 2):
router.cache.set_cache(
key=f"{MODEL_GROUP}:rpm:{(now + timedelta(minutes=offset)).strftime('%H-%M')}",
value={LOW_USAGE_DEPLOYMENT_ID: recorded},
ttl=float("inf"),
)
return router
@pytest.mark.parametrize("enable_pre_call_checks", [True, False])
def test_usage_based_routing_v1_serves_a_deployment_below_its_recorded_rpm_limit(enable_pre_call_checks: bool) -> None:
router: Final = _router_with_recorded_rpm(recorded=1, enable_pre_call_checks=enable_pre_call_checks)
response: Final = router.completion(model=MODEL_GROUP, messages=[{"role": "user", "content": "hello"}])
assert response.choices[0].message.content == f"from {LOW_USAGE_DEPLOYMENT_ID}"
@pytest.mark.parametrize("enable_pre_call_checks", [True, False])
def test_usage_based_routing_v1_rejects_a_deployment_that_reached_its_recorded_rpm_limit(
enable_pre_call_checks: bool,
) -> None:
router: Final = _router_with_recorded_rpm(recorded=RPM_LIMIT, enable_pre_call_checks=enable_pre_call_checks)
with pytest.raises(ValueError, match=RouterErrors.no_deployments_available.value):
router.completion(model=MODEL_GROUP, messages=[{"role": "user", "content": "hello"}])

View file

@ -0,0 +1,101 @@
import json
from collections.abc import Iterator
from typing import Final
import httpx
import litellm
import pytest
import respx
from litellm import Router
VECTOR_STORES_URL: Final = "https://api.openai.com/v1/vector_stores"
@pytest.fixture(autouse=True)
def httpx_transport(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setenv("OPENAI_API_KEY", "sk-vector-store-test")
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
litellm.in_memory_llm_clients_cache.flush_cache()
yield
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.mark.asyncio
async def test_avector_store_update_sends_name_and_metadata_to_provider() -> None:
with respx.mock(assert_all_called=True) as mock:
route: Final = mock.post(f"{VECTOR_STORES_URL}/vs_test123").mock(
return_value=httpx.Response(
200,
json={
"id": "vs_test123",
"object": "vector_store",
"created_at": 1699061776,
"name": "Updated Name",
"metadata": {"key": "value"},
"status": "completed",
},
)
)
result: Final = await Router(model_list=[]).avector_store_update(
vector_store_id="vs_test123",
name="Updated Name",
metadata={"key": "value"},
custom_llm_provider="openai",
)
sent: Final = json.loads(route.calls.last.request.content)
assert route.call_count == 1
assert sent["name"] == "Updated Name"
assert sent["metadata"] == {"key": "value"}
assert result["id"] == "vs_test123"
assert result["name"] == "Updated Name"
assert result["metadata"]["key"] == "value"
def test_vector_store_list_forwards_pagination_params_to_provider() -> None:
with respx.mock(assert_all_called=True) as mock:
route: Final = mock.get(VECTOR_STORES_URL).mock(
return_value=httpx.Response(
200,
json={
"object": "list",
"data": [{"id": f"vs_{index}", "object": "vector_store"} for index in range(5)],
"has_more": True,
"first_id": "vs_0",
"last_id": "vs_4",
},
)
)
result: Final = Router(model_list=[]).vector_store_list(
limit=5, after="vs_previous", order="asc", custom_llm_provider="openai"
)
params: Final = route.calls.last.request.url.params
assert params["limit"] == "5"
assert params["after"] == "vs_previous"
assert params["order"] == "asc"
assert result["has_more"] is True
assert len(result["data"]) == 5
def test_vector_store_update_forwards_expires_after_to_provider() -> None:
expires_after: Final = {"anchor": "last_active_at", "days": 7}
with respx.mock(assert_all_called=True) as mock:
route: Final = mock.post(f"{VECTOR_STORES_URL}/vs_test123").mock(
return_value=httpx.Response(
200,
json={
"id": "vs_test123",
"object": "vector_store",
"expires_after": expires_after,
"expires_at": 1699668576,
},
)
)
result: Final = Router(model_list=[]).vector_store_update(
vector_store_id="vs_test123", expires_after=expires_after, custom_llm_provider="openai"
)
sent: Final = json.loads(route.calls.last.request.content)
assert sent["expires_after"] == expires_after
assert result["expires_after"]["days"] == 7
assert result["expires_at"] == 1699668576

View file

@ -7,10 +7,17 @@ is called without call_type in kwargs (e.g. from batch polling callbacks).
import pytest
from datetime import datetime
from typing import Final
from unittest.mock import AsyncMock, patch
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
import litellm
from litellm._service_logger import ServiceLogging
from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
from litellm.types.services import ServiceTypes
@ -316,3 +323,47 @@ async def test_only_redis_service_spans_carry_the_ambient_key_family(monkeypatch
"redis.get router_session_pins": "router_session_pins",
"batch_write_to_db _PROXY_track_cost_callback": None,
}
def _in_memory_provider(exporter: InMemorySpanExporter) -> TracerProvider:
provider: Final = TracerProvider()
provider.add_span_processor(SimpleSpanProcessor(exporter))
return provider
@pytest.mark.asyncio
async def test_generic_otel_logger_receives_service_event_once_beside_langfuse_otel(
monkeypatch: pytest.MonkeyPatch,
) -> None:
langfuse_exporter: Final = InMemorySpanExporter()
generic_exporter: Final = InMemorySpanExporter()
langfuse_provider: Final = _in_memory_provider(langfuse_exporter)
generic_provider: Final = _in_memory_provider(generic_exporter)
langfuse_logger: Final = LangfuseOtelLogger(
config=OpenTelemetryConfig(exporter="console", skip_set_global=True), tracer_provider=langfuse_provider
)
generic_logger: Final = OpenTelemetry(
config=OpenTelemetryConfig(exporter="console", skip_set_global=True), tracer_provider=generic_provider
)
monkeypatch.setattr(litellm, "service_callback", [langfuse_logger, generic_logger])
parent: Final = generic_logger.tracer.start_span("parent")
try:
await ServiceLogging().async_service_success_hook(
service=ServiceTypes.DB,
call_type="success",
duration=0.1,
parent_otel_span=parent,
start_time=0.0,
end_time=1.0,
)
finally:
langfuse_provider.shutdown()
generic_provider.shutdown()
service_spans: Final = [
span for span in generic_exporter.get_finished_spans() if span.attributes.get("service") == ServiceTypes.DB.value
]
assert len(service_spans) == 1
assert service_spans[0].attributes.get("call_type") == "success"
assert langfuse_exporter.get_finished_spans() == ()