mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
a81f507bca
commit
d5173e3d70
47 changed files with 1746 additions and 4409 deletions
147
tests/integration/configuration/test_callback_leak_contracts.py
Normal file
147
tests/integration/configuration/test_callback_leak_contracts.py
Normal 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
|
||||
152
tests/integration/management/test_key_route_contracts.py
Normal file
152
tests/integration/management/test_key_route_contracts.py
Normal 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}]
|
||||
|
|
@ -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
|
||||
205
tests/integration/management/test_team_member_routes.py
Normal file
205
tests/integration/management/test_team_member_routes.py
Normal 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
|
||||
34
tests/integration/management/test_user_routes.py
Normal file
34
tests/integration/management/test_user_routes.py
Normal 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
|
||||
|
|
@ -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}
|
||||
102
tests/integration/spend/test_spend_attribution_contracts.py
Normal file
102
tests/integration/spend/test_spend_attribution_contracts.py
Normal 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
|
||||
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
53
tests/rust-python-harness/shared/reporting/test_models.py
Normal file
53
tests/rust-python-harness/shared/reporting/test_models.py
Normal 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)
|
||||
26
tests/rust-python-harness/shared/reporting/test_ui.py
Normal file
26
tests/rust-python-harness/shared/reporting/test_ui.py
Normal 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"
|
||||
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
|
|
@ -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."
|
||||
)
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
|
|
@ -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}")
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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!")
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
@ -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!")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
40
tests/unit/litellm_core_utils/test_default_encoding.py
Normal file
40
tests/unit/litellm_core_utils/test_default_encoding.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
0
tests/unit/llms/litellm_proxy/responses/__init__.py
Normal file
0
tests/unit/llms/litellm_proxy/responses/__init__.py
Normal 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={})
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"}])
|
||||
|
|
|
|||
101
tests/unit/test_router/test_router_vector_store_endpoints.py
Normal file
101
tests/unit/test_router/test_router_vector_store_endpoints.py
Normal 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
|
||||
|
|
@ -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() == ()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue