mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
test(integration): move legacy proxy, router and Redis tests into tests/integration (#44128)
* test(integration): move legacy proxy, router and Redis tests into tests/integration Port 39 legacy tests to the integration tier that owns them, running against the scripted upstream, local Postgres and Redis, test-owned wire peers and owned proxies. Delete 5 legacy tests whose contract is already owned by an existing integration test, and remove the legacy functions, files and helpers left unused. * test(integration): cover recovery of a spent key after its budget is raised
This commit is contained in:
parent
826b21aab6
commit
a93c396a5c
43 changed files with 1830 additions and 3418 deletions
|
|
@ -14,78 +14,6 @@ from litellm.caching.caching import DualCache
|
|||
from litellm.exceptions import BlockedPiiEntityError
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_presidio_with_entities_config():
|
||||
"""Test for Presidio guardrail with entities config - requires actual Presidio API"""
|
||||
# Setup the guardrail with specific entities config
|
||||
litellm._turn_on_debug()
|
||||
pii_entities_config = {
|
||||
PiiEntityType.CREDIT_CARD: PiiAction.MASK,
|
||||
PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK,
|
||||
}
|
||||
|
||||
presidio_guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
pii_entities_config=pii_entities_config,
|
||||
presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"),
|
||||
presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"),
|
||||
)
|
||||
|
||||
# Test text with different PII types
|
||||
test_text = "My credit card number is 4111-1111-1111-1111, my email is test@example.com, and my phone is 555-123-4567"
|
||||
|
||||
# Test the analyze request configuration
|
||||
analyze_request = presidio_guardrail._get_presidio_analyze_request_payload(
|
||||
text=test_text, presidio_config=None, request_data={}
|
||||
)
|
||||
|
||||
# Verify entities were passed correctly
|
||||
assert "entities" in analyze_request
|
||||
assert set(analyze_request["entities"]) == set(pii_entities_config.keys())
|
||||
|
||||
# Test the check_pii method - this will call the actual Presidio API
|
||||
redacted_text = await presidio_guardrail.check_pii(
|
||||
text=test_text, output_parse_pii=True, presidio_config=None, request_data={}
|
||||
)
|
||||
|
||||
# Verify PII has been masked/replaced/redacted in the result
|
||||
assert "4111-1111-1111-1111" not in redacted_text
|
||||
assert "test@example.com" not in redacted_text
|
||||
|
||||
# Since this entity is not in the config, it should not be masked
|
||||
assert "555-123-4567" in redacted_text
|
||||
|
||||
# The specific replacements will vary based on Presidio's implementation
|
||||
print(f"Redacted text: {redacted_text}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_presidio_apply_guardrail():
|
||||
"""Test for Presidio guardrail apply guardrail - requires actual Presidio API"""
|
||||
litellm._turn_on_debug()
|
||||
presidio_guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
pii_entities_config={},
|
||||
presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"),
|
||||
presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"),
|
||||
)
|
||||
|
||||
test_text = (
|
||||
"My credit card number is 4111-1111-1111-1111 and my email is test@example.com"
|
||||
)
|
||||
response = await presidio_guardrail.apply_guardrail(
|
||||
inputs={"texts": [test_text]},
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
print("response from apply guardrail for presidio: ", response)
|
||||
|
||||
# Extract the modified text from the response
|
||||
modified_text = response["texts"][0] if response.get("texts") else ""
|
||||
|
||||
# assert the default config masks the credit card and email
|
||||
assert "4111-1111-1111-1111" not in modified_text
|
||||
assert "test@example.com" not in modified_text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_presidio_with_blocked_entities():
|
||||
"""Test for Presidio guardrail with blocked entities - requires actual Presidio API"""
|
||||
|
|
@ -174,58 +102,6 @@ async def test_presidio_pre_call_hook_with_blocked_entities():
|
|||
assert excinfo.value.guardrail_name == presidio_guardrail.guardrail_name
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("call_type", ["completion", "acompletion"])
|
||||
async def test_presidio_pre_call_hook_with_different_call_types(call_type):
|
||||
"""Test for Presidio guardrail pre-call hook with both completion and acompletion call types"""
|
||||
# Setup the guardrail with specific entities config
|
||||
pii_entities_config = {
|
||||
PiiEntityType.CREDIT_CARD: PiiAction.MASK,
|
||||
PiiEntityType.EMAIL_ADDRESS: PiiAction.MASK,
|
||||
}
|
||||
|
||||
presidio_guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
pii_entities_config=pii_entities_config,
|
||||
presidio_analyzer_api_base=os.environ.get("PRESIDIO_ANALYZER_API_BASE"),
|
||||
presidio_anonymizer_api_base=os.environ.get("PRESIDIO_ANONYMIZER_API_BASE"),
|
||||
)
|
||||
|
||||
# Create a sample request with PII data
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{
|
||||
"role": "user",
|
||||
"content": "My credit card is 4111-1111-1111-1111 and my email is test@example.com. My phone number is 555-123-4567",
|
||||
},
|
||||
],
|
||||
"model": "gpt-5-mini",
|
||||
}
|
||||
|
||||
# Mock objects needed for the pre-call hook
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test_key")
|
||||
cache = DualCache()
|
||||
|
||||
# Call the pre-call hook with the specified call type
|
||||
modified_data = await presidio_guardrail.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict, cache=cache, data=data, call_type=call_type
|
||||
)
|
||||
|
||||
# Verify the messages have been modified to mask PII
|
||||
assert (
|
||||
modified_data["messages"][0]["content"] == "You are a helpful assistant."
|
||||
) # System prompt should be unchanged
|
||||
|
||||
user_message = modified_data["messages"][1]["content"]
|
||||
assert "4111-1111-1111-1111" not in user_message
|
||||
assert "test@example.com" not in user_message
|
||||
|
||||
# Since this entity is not in the config, it should not be masked
|
||||
assert "555-123-4567" in user_message
|
||||
|
||||
print(f"Modified user message for call_type={call_type}: {user_message}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"base_url",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -0,0 +1,33 @@
|
|||
import uuid
|
||||
from hashlib import sha256
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
|
||||
|
||||
def test_a_key_bound_to_a_user_id_with_no_user_row_serves_and_attributes_spend_to_that_id(gateway: Gateway) -> None:
|
||||
user_id: Final = f"integration-absent-{uuid.uuid4().hex}"
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
key: Final = scenario.key(user_id=user_id, models=[model])
|
||||
assert read_rows('SELECT user_id FROM "LiteLLM_UserTable" WHERE user_id=%s', (user_id,)) == []
|
||||
response: Final = gateway.chat(model, key=key, text=f"unknown user {uuid.uuid4().hex}")
|
||||
assert object_value(response["usage"])["total_tokens"] == 40
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT "user", spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (str(response["id"]),)
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
assert rows[0]["user"] == user_id
|
||||
assert float(str(rows[0]["spend"])) == pytest.approx(0.06)
|
||||
eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (sha256(key.encode()).hexdigest(),)
|
||||
),
|
||||
lambda values: len(values) == 1 and float(str(values[0]["spend"])) == pytest.approx(0.06),
|
||||
seconds=70,
|
||||
)
|
||||
107
tests/integration/authorization/test_team_member_permissions.py
Normal file
107
tests/integration/authorization/test_team_member_permissions.py
Normal file
|
|
@ -0,0 +1,107 @@
|
|||
from dataclasses import dataclass
|
||||
from hashlib import sha256
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from integration._support.client import Gateway, Scenario, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from pydantic import JsonValue
|
||||
|
||||
PERMISSION_ERROR: Final = "team_member_permission_error"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Member:
|
||||
team_id: str
|
||||
team_key: str
|
||||
member_key: str
|
||||
|
||||
|
||||
def _member(scenario: Scenario, permissions: list[JsonValue] | None) -> Member:
|
||||
team_id: Final = scenario.team() if permissions is None else scenario.team(team_member_permissions=permissions)
|
||||
team_key: Final = scenario.key(team_id=team_id, metadata={"owner": "team"})
|
||||
member: Final = scenario.member(team_id)
|
||||
return Member(team_id, team_key, scenario.key(user_id=member))
|
||||
|
||||
|
||||
def _team_key_row(key: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT team_id, metadata FROM "LiteLLM_VerificationToken" WHERE token = %s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
)
|
||||
|
||||
|
||||
def _team_key_count(team_id: str) -> int:
|
||||
return len(read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE team_id = %s', (team_id,)))
|
||||
|
||||
|
||||
def _refused(response: httpx.Response, status: int, error_type: str | None = None) -> None:
|
||||
assert response.status_code == status, response.text
|
||||
if error_type is not None:
|
||||
assert object_value(response.json()["error"])["type"] == error_type, response.text
|
||||
|
||||
|
||||
def test_default_member_permissions_only_allow_reading_team_keys(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
member: Final = _member(scenario, None)
|
||||
generated: Final = gateway.request("POST", "/key/generate", {"team_id": member.team_id}, key=member.member_key)
|
||||
_refused(generated, 401, PERMISSION_ERROR)
|
||||
updated: Final = gateway.request(
|
||||
"POST",
|
||||
"/key/update",
|
||||
{"key": member.team_key, "team_id": "ATTACKER_TEAM_ID", "metadata": {"owner": "member"}},
|
||||
key=member.member_key,
|
||||
)
|
||||
_refused(updated, 401, PERMISSION_ERROR)
|
||||
_refused(gateway.request("POST", "/key/delete", {"keys": [member.team_key]}, key=member.member_key), 403)
|
||||
_refused(gateway.request("POST", "/key/regenerate", {"key": member.team_key}, key=member.member_key), 401)
|
||||
info: Final = gateway.request("GET", "/key/info", key=member.member_key, params={"key": member.team_key})
|
||||
assert info.status_code == 200, info.text
|
||||
assert object_value(info.json()["info"])["team_id"] == member.team_id
|
||||
assert _team_key_row(member.team_key) == [{"team_id": member.team_id, "metadata": {"owner": "team"}}]
|
||||
assert _team_key_count(member.team_id) == 1
|
||||
|
||||
|
||||
def test_update_and_delete_permissions_let_a_member_edit_but_not_create_delete_or_regenerate(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
member: Final = _member(scenario, ["/key/update", "/key/delete", "/key/info"])
|
||||
updated: Final = gateway.request(
|
||||
"POST",
|
||||
"/key/update",
|
||||
{"key": member.team_key, "team_id": member.team_id, "metadata": {"owner": "member"}},
|
||||
key=member.member_key,
|
||||
)
|
||||
assert updated.status_code == 200, updated.text
|
||||
assert _team_key_row(member.team_key) == [{"team_id": member.team_id, "metadata": {"owner": "member"}}]
|
||||
_refused(gateway.request("POST", "/key/delete", {"keys": [member.team_key]}, key=member.member_key), 403)
|
||||
generated: Final = gateway.request("POST", "/key/generate", {"team_id": member.team_id}, key=member.member_key)
|
||||
_refused(generated, 401, PERMISSION_ERROR)
|
||||
regenerated: Final = gateway.request(
|
||||
"POST", "/key/regenerate", {"key": member.team_key, "team_id": member.team_id}, key=member.member_key
|
||||
)
|
||||
_refused(regenerated, 401)
|
||||
assert _team_key_count(member.team_id) == 1
|
||||
|
||||
|
||||
def test_generate_permission_lets_a_member_create_team_keys_but_not_change_existing_ones(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
member: Final = _member(scenario, ["/key/generate"])
|
||||
generated: Final = gateway.request("POST", "/key/generate", {"team_id": member.team_id}, key=member.member_key)
|
||||
assert generated.status_code == 200, generated.text
|
||||
created: Final = string_value(generated.json()["key"])
|
||||
scenario.cleanups.callback(scenario.delete_key, created)
|
||||
assert _team_key_row(created) == [{"team_id": member.team_id, "metadata": {}}]
|
||||
updated: Final = gateway.request(
|
||||
"POST",
|
||||
"/key/update",
|
||||
{"key": member.team_key, "team_id": member.team_id, "metadata": {"owner": "member"}},
|
||||
key=member.member_key,
|
||||
)
|
||||
_refused(updated, 401, PERMISSION_ERROR)
|
||||
assert _team_key_row(member.team_key) == [{"team_id": member.team_id, "metadata": {"owner": "team"}}]
|
||||
_refused(gateway.request("POST", "/key/delete", {"keys": [member.team_key]}, key=member.member_key), 403)
|
||||
regenerated: Final = gateway.request(
|
||||
"POST", "/key/regenerate", {"key": member.team_key, "team_id": member.team_id}, key=member.member_key
|
||||
)
|
||||
_refused(regenerated, 401, PERMISSION_ERROR)
|
||||
assert _team_key_count(member.team_id) == 2
|
||||
105
tests/integration/authorization/test_team_scoped_models.py
Normal file
105
tests/integration/authorization/test_team_scoped_models.py
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, object_value, string_value
|
||||
from pydantic import JsonValue
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def upstream(gateway: Gateway) -> Iterator[httpx.Client]:
|
||||
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as client:
|
||||
client.get("/__observations").raise_for_status()
|
||||
yield client
|
||||
|
||||
|
||||
def _observed_models(upstream: httpx.Client) -> list[JsonValue]:
|
||||
observed: Final = upstream.get("/__observations")
|
||||
observed.raise_for_status()
|
||||
return [request["body"]["model"] for request in observed.json()["requests"]]
|
||||
|
||||
|
||||
def _chat(gateway: Gateway, model: str, key: str) -> httpx.Response:
|
||||
return gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": f"team model {uuid.uuid4().hex}"}]},
|
||||
key=key,
|
||||
)
|
||||
|
||||
|
||||
def _ids(listing: dict[str, JsonValue]) -> set[JsonValue]:
|
||||
data: Final = listing["data"]
|
||||
assert isinstance(data, list)
|
||||
return {object_value(entry)["id"] for entry in data}
|
||||
|
||||
|
||||
def test_a_model_created_for_a_team_is_listed_in_that_teams_models(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
team: Final = scenario.team(models=[])
|
||||
other: Final = scenario.team(models=[])
|
||||
model: Final = scenario.model(model_info={"team_id": team})
|
||||
own_models: Final = object_value(gateway.get("/team/info", {"team_id": team})["team_info"])["models"]
|
||||
other_models: Final = object_value(gateway.get("/team/info", {"team_id": other})["team_info"])["models"]
|
||||
assert isinstance(own_models, list) and isinstance(other_models, list)
|
||||
assert model in own_models
|
||||
assert model not in other_models
|
||||
|
||||
|
||||
def test_a_team_model_is_listed_and_served_only_for_keys_of_its_team(gateway: Gateway, upstream: httpx.Client) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
team: Final = scenario.team(models=[])
|
||||
other: Final = scenario.team(models=[])
|
||||
provider_model: Final = f"team-scoped-{uuid.uuid4().hex}"
|
||||
model: Final = scenario.model(model=f"openai/{provider_model}", model_info={"team_id": team})
|
||||
team_key: Final = scenario.key(team_id=team)
|
||||
other_key: Final = scenario.key(team_id=other)
|
||||
assert model in _ids(object_value(gateway.request("GET", "/models", key=team_key).json()))
|
||||
assert model not in _ids(object_value(gateway.request("GET", "/models", key=other_key).json()))
|
||||
served: Final = _chat(gateway, model, team_key)
|
||||
assert served.status_code == 200, served.text
|
||||
refused: Final = _chat(gateway, model, other_key)
|
||||
assert refused.status_code == 400, refused.text
|
||||
assert _observed_models(upstream) == [provider_model]
|
||||
|
||||
|
||||
def _v2_team_public_names(gateway: Gateway, key: str, model: str) -> list[JsonValue]:
|
||||
response: Final = gateway.request("GET", "/v2/model/info", key=key, params={"model_name": model})
|
||||
assert response.status_code == 200, response.text
|
||||
return [entry["model_info"].get("team_public_model_name") for entry in response.json()["data"]]
|
||||
|
||||
|
||||
def test_v2_model_info_reports_a_team_model_to_team_and_non_team_keys(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
team: Final = scenario.team(models=[])
|
||||
model: Final = scenario.model(model_info={"team_id": team})
|
||||
assert _v2_team_public_names(gateway, scenario.key(team_id=team), model) == [model]
|
||||
assert _v2_team_public_names(gateway, scenario.key(), model) == [model]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("set_on", ["team_new", "team_update"])
|
||||
def test_team_model_alias_routes_a_team_key_to_its_target(
|
||||
gateway: Gateway, upstream: httpx.Client, set_on: str
|
||||
) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
provider_model: Final = f"team-alias-{uuid.uuid4().hex}"
|
||||
model: Final = scenario.model(model=f"openai/{provider_model}")
|
||||
alias: Final = f"alias-{uuid.uuid4().hex}"
|
||||
team: Final = (
|
||||
scenario.team(models=[model], model_aliases={alias: model})
|
||||
if set_on == "team_new"
|
||||
else scenario.team(models=[model])
|
||||
)
|
||||
if set_on == "team_update":
|
||||
gateway.post("/team/update", {"team_id": team, "model_aliases": {alias: model}})
|
||||
key: Final = scenario.key(team_id=team, models=[model])
|
||||
response: Final = _chat(gateway, alias, key)
|
||||
assert response.status_code == 200, response.text
|
||||
assert string_value(response.json()["model"]) == alias
|
||||
assert _observed_models(upstream) == [provider_model]
|
||||
unaliased: Final = _chat(gateway, f"alias-{uuid.uuid4().hex}", key)
|
||||
assert unaliased.status_code == 403, unaliased.text
|
||||
assert unaliased.json()["error"]["type"] == "key_model_access_denied"
|
||||
assert _observed_models(upstream) == []
|
||||
|
|
@ -0,0 +1,90 @@
|
|||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, gateway_from_environment
|
||||
from integration._support.process import owned_proxy
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(180)
|
||||
|
||||
|
||||
def _deployment(model_name: str, model: str, upstream_url: str) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {"model": model, "api_base": f"{upstream_url}/v1", "api_key": "synthetic-wildcard-key"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def candidate(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("wildcard-access")
|
||||
with gateway_from_environment() as base:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
_deployment("*", "openai/*", base.upstream_url),
|
||||
_deployment("anthropic/*", "openai/*", base.upstream_url),
|
||||
_deployment("groq/*", "openai/*", base.upstream_url),
|
||||
_deployment("good-model", "openai/good-model-upstream", base.upstream_url),
|
||||
]
|
||||
path: Final = directory / "wildcard-access.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(base, directory, {}, config=path) as proxy:
|
||||
yield proxy
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def upstream(candidate: Gateway) -> Iterator[httpx.Client]:
|
||||
with httpx.Client(base_url=candidate.upstream_url, timeout=5, trust_env=False) as client:
|
||||
client.get("/__observations").raise_for_status()
|
||||
yield client
|
||||
|
||||
|
||||
def _chat(gateway: Gateway, model: str, key: str) -> httpx.Response:
|
||||
return gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": f"wildcard {uuid.uuid4().hex}"}]},
|
||||
key=key,
|
||||
)
|
||||
|
||||
|
||||
def _observed_models(upstream: httpx.Client) -> list[JsonValue]:
|
||||
return [request["body"]["model"] for request in upstream.get("/__observations").json()["requests"]]
|
||||
|
||||
|
||||
def test_an_all_models_key_reaches_a_model_served_only_by_the_catch_all_deployment(
|
||||
candidate: Gateway, upstream: httpx.Client
|
||||
) -> None:
|
||||
with candidate.scenario() as scenario:
|
||||
key: Final = scenario.key(models=["*"])
|
||||
unlisted: Final = f"unlisted-{uuid.uuid4().hex}"
|
||||
response: Final = _chat(candidate, unlisted, key)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _observed_models(upstream) == [unlisted]
|
||||
|
||||
|
||||
def test_a_key_without_models_inherits_the_users_exact_and_wildcard_grants(
|
||||
candidate: Gateway, upstream: httpx.Client
|
||||
) -> None:
|
||||
with candidate.scenario() as scenario:
|
||||
user_id: Final = scenario.user(models=["good-model", "anthropic/*"])
|
||||
key: Final = scenario.key(user_id=user_id, models=[])
|
||||
wildcard_model: Final = f"claude-{uuid.uuid4().hex}"
|
||||
assert _chat(candidate, f"anthropic/{wildcard_model}", key).status_code == 200
|
||||
assert _chat(candidate, "good-model", key).status_code == 200
|
||||
assert _observed_models(upstream) == [wildcard_model, "good-model-upstream"]
|
||||
denied: Final = tuple(
|
||||
_chat(candidate, outside, key)
|
||||
for outside in (f"groq/{wildcard_model}", f"bedrock/anthropic.{wildcard_model}")
|
||||
)
|
||||
assert [(response.status_code, response.json()["error"]["type"]) for response in denied] == [
|
||||
(403, "user_model_access_denied")
|
||||
] * 2, [response.text for response in denied]
|
||||
assert _observed_models(upstream) == []
|
||||
assert _chat(candidate, f"groq/{wildcard_model}", candidate.key).status_code == 200
|
||||
assert _observed_models(upstream) == [wildcard_model]
|
||||
34
tests/integration/management/test_model_health_check.py
Normal file
34
tests/integration/management/test_model_health_check.py
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
import uuid
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from integration._support.client import Gateway, object_value
|
||||
|
||||
|
||||
def test_health_check_of_a_model_added_through_the_api_calls_its_upstream_and_reports_it_healthy(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
with (
|
||||
gateway.scenario() as scenario,
|
||||
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
|
||||
):
|
||||
provider_model: Final = f"health-{uuid.uuid4().hex}"
|
||||
model: Final = scenario.model(model=f"openai/{provider_model}")
|
||||
key: Final = scenario.key(models=[model])
|
||||
listed: Final = gateway.request("GET", "/v2/model/info", key=key, params={"model": model})
|
||||
assert listed.status_code == 200, listed.text
|
||||
assert [entry["model_name"] for entry in listed.json()["data"]] == [model]
|
||||
assert (
|
||||
object_value(gateway.chat(model, key=key, text=f"health {uuid.uuid4().hex}")["usage"])["total_tokens"] == 40
|
||||
)
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
health: Final = gateway.request("GET", "/health", params={"model": model})
|
||||
assert health.status_code == 200, health.text
|
||||
report: Final = health.json()
|
||||
assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report
|
||||
assert [(endpoint["model"], endpoint["api_base"]) for endpoint in report["healthy_endpoints"]] == [
|
||||
(f"openai/{provider_model}", f"{gateway.upstream_url}/v1")
|
||||
]
|
||||
assert [request["body"]["model"] for request in upstream.get("/__observations").json()["requests"]] == [
|
||||
provider_model
|
||||
]
|
||||
89
tests/integration/management/test_organization_lifecycle.py
Normal file
89
tests/integration/management/test_organization_lifecycle.py
Normal file
|
|
@ -0,0 +1,89 @@
|
|||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from integration._support.client import Gateway, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from pydantic import JsonValue
|
||||
|
||||
CONCURRENT_CREATES: Final = 8
|
||||
|
||||
|
||||
def _membership_rows(organization_id: str) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT user_id, user_role FROM "LiteLLM_OrganizationMembership" WHERE organization_id = %s',
|
||||
(organization_id,),
|
||||
)
|
||||
|
||||
|
||||
def _listed(gateway: Gateway, organization_id: str) -> dict[str, JsonValue]:
|
||||
response: Final = gateway.request("GET", "/organization/list")
|
||||
assert response.status_code == 200, response.text
|
||||
entries: Final = response.json()
|
||||
assert isinstance(entries, list)
|
||||
matches: Final = [entry for entry in entries if entry["organization_id"] == organization_id]
|
||||
assert len(matches) == 1, f"{organization_id} listed {len(matches)} times"
|
||||
return object_value(matches[0])
|
||||
|
||||
|
||||
def test_concurrent_creates_with_one_alias_each_persist_a_distinct_organization(gateway: Gateway) -> None:
|
||||
alias: Final = f"integration-{uuid.uuid4().hex}"
|
||||
|
||||
def create(_: int) -> httpx.Response:
|
||||
return gateway.request("POST", "/organization/new", {"organization_alias": alias})
|
||||
|
||||
with ThreadPoolExecutor(max_workers=CONCURRENT_CREATES) as pool:
|
||||
responses: Final = tuple(pool.map(create, range(CONCURRENT_CREATES)))
|
||||
created: Final = tuple(response.json() for response in responses if response.status_code == 200)
|
||||
with gateway.scenario() as scenario:
|
||||
for body in created:
|
||||
scenario.cleanups.callback(
|
||||
scenario.delete_organization, string_value(body["organization_id"]), string_value(body["budget_id"])
|
||||
)
|
||||
assert [response.status_code for response in responses] == [200] * CONCURRENT_CREATES, [
|
||||
response.text for response in responses
|
||||
]
|
||||
identities: Final = {string_value(body["organization_id"]) for body in created}
|
||||
assert len(identities) == CONCURRENT_CREATES
|
||||
rows: Final = read_rows(
|
||||
'SELECT organization_id FROM "LiteLLM_OrganizationTable" WHERE organization_alias = %s', (alias,)
|
||||
)
|
||||
assert {string_value(row["organization_id"]) for row in rows} == identities
|
||||
|
||||
|
||||
def test_list_returns_each_organization_with_its_budget_and_members(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
organization_id: Final = scenario.organization(max_budget=3.5, tpm_limit=120)
|
||||
member: Final = scenario.org_member(organization_id, role="internal_user")
|
||||
listed: Final = _listed(gateway, organization_id)
|
||||
budget: Final = object_value(listed["litellm_budget_table"])
|
||||
assert (budget["max_budget"], budget["tpm_limit"]) == (3.5, 120)
|
||||
members: Final = listed["members"]
|
||||
assert isinstance(members, list)
|
||||
assert [(object_value(entry)["user_id"], object_value(entry)["user_role"]) for entry in members] == [
|
||||
(member, "internal_user")
|
||||
]
|
||||
|
||||
|
||||
def test_member_role_update_and_removal_persist_and_read_back(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
organization_id: Final = scenario.organization()
|
||||
member: Final = scenario.org_member(organization_id, role="internal_user")
|
||||
assert _membership_rows(organization_id) == [{"user_id": member, "user_role": "internal_user"}]
|
||||
updated: Final = gateway.request(
|
||||
"PATCH",
|
||||
"/organization/member_update",
|
||||
{"organization_id": organization_id, "user_id": member, "role": "org_admin"},
|
||||
)
|
||||
assert updated.status_code == 200, updated.text
|
||||
assert _membership_rows(organization_id) == [{"user_id": member, "user_role": "org_admin"}]
|
||||
info_members: Final = gateway.get("/organization/info", {"organization_id": organization_id})["members"]
|
||||
assert isinstance(info_members, list)
|
||||
assert [object_value(entry)["user_role"] for entry in info_members] == ["org_admin"]
|
||||
removed: Final = gateway.request(
|
||||
"DELETE", "/organization/member_delete", {"organization_id": organization_id, "user_id": member}
|
||||
)
|
||||
assert removed.status_code == 200, removed.text
|
||||
assert _membership_rows(organization_id) == []
|
||||
assert _listed(gateway, organization_id)["members"] == []
|
||||
|
|
@ -0,0 +1,53 @@
|
|||
import uuid
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually, gateway_from_environment
|
||||
from integration._support.otlp_sink import Span, SpanSinks, recorded_spans
|
||||
from integration._support.process import owned_proxy
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(180)
|
||||
|
||||
AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path]
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def gateway(audit_sinks: SpanSinks) -> Iterator[Gateway]:
|
||||
with gateway_from_environment() as base:
|
||||
yield base
|
||||
|
||||
|
||||
def _traces(spans: tuple[Span, ...]) -> dict[str, frozenset[str]]:
|
||||
trace_ids: Final = {span["trace_id"] for span in spans}
|
||||
return {trace: frozenset(span["name"] for span in spans if span["trace_id"] == trace) for trace in trace_ids}
|
||||
|
||||
|
||||
def test_default_otel_logger_puts_datastore_model_and_spend_writer_spans_in_the_request_trace(
|
||||
gateway: Gateway, audit_sinks: SpanSinks, otel_audit_config: AuditConfigWriter, tmp_path: Path
|
||||
) -> None:
|
||||
config: Final = otel_audit_config(tmp_path, {})
|
||||
overrides: Final = {"OTEL_EXPORTER": "http/json", "OTEL_ENDPOINT": audit_sinks.operator}
|
||||
with owned_proxy(gateway, tmp_path, overrides, config=config) as candidate, candidate.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
key: Final = scenario.key(models=[model])
|
||||
start, _ = recorded_spans(audit_sinks.operator)
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": f"otel v1 {uuid.uuid4().hex}"}]},
|
||||
key=key,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
expected: Final = frozenset({"postgres", "redis", "raw_gen_ai_request", "batch_write_to_db"})
|
||||
traces: Final = eventually(
|
||||
lambda: _traces(recorded_spans(audit_sinks.operator, start)[1]),
|
||||
lambda grouped: any(expected <= names for names in grouped.values()),
|
||||
seconds=60,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
assert any(expected <= names for names in traces.values()), {
|
||||
trace: sorted(names) for trace, names in traces.items()
|
||||
}
|
||||
141
tests/integration/observability/test_presidio_entity_masking.py
Normal file
141
tests/integration/observability/test_presidio_entity_masking.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
import json
|
||||
import re
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, string_value
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
CARD: Final = "4111-1111-1111-1111"
|
||||
EMAIL: Final = "jane.doe@example.com"
|
||||
PHONE: Final = "555-123-4567"
|
||||
SYSTEM_PROMPT: Final = "You are a helpful assistant."
|
||||
RECOGNIZERS: Final = {
|
||||
"CREDIT_CARD": re.escape(CARD),
|
||||
"EMAIL_ADDRESS": re.escape(EMAIL),
|
||||
"PHONE_NUMBER": re.escape(PHONE),
|
||||
}
|
||||
|
||||
|
||||
def _detect(entity: str, text: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(
|
||||
{"entity_type": entity, "start": match.start(), "end": match.end(), "score": 0.95}
|
||||
for match in re.finditer(RECOGNIZERS[entity], text)
|
||||
)
|
||||
|
||||
|
||||
def _analyze(request: Request) -> Reply:
|
||||
assert request.target == "/analyze", request.target
|
||||
body: Final = json.loads(request.body)
|
||||
requested: Final = body.get("entities") or list(RECOGNIZERS)
|
||||
findings: Final = list(chain.from_iterable(_detect(entity, body["text"]) for entity in requested))
|
||||
return Reply(body=json.dumps(findings).encode())
|
||||
|
||||
|
||||
def _anonymize(request: Request) -> Reply:
|
||||
assert request.target == "/anonymize", request.target
|
||||
body: Final = json.loads(request.body)
|
||||
spans: Final = sorted(body["analyzer_results"], key=lambda item: item["start"])
|
||||
pieces: Final = [
|
||||
body["text"][(spans[index - 1]["end"] if index else 0) : span["start"]] + f"<{span['entity_type']}>"
|
||||
for index, span in enumerate(spans)
|
||||
]
|
||||
tail: Final = body["text"][spans[-1]["end"] :] if spans else body["text"]
|
||||
return Reply(body=json.dumps({"text": "".join(pieces) + tail, "items": []}).encode())
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Presidio:
|
||||
name: str
|
||||
analyzer: Wire
|
||||
anonymizer: Wire
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _presidio(gateway: Gateway, mode: str, entities: Mapping[str, str] | None) -> Iterator[Presidio]:
|
||||
name: Final = f"presidio-{uuid.uuid4().hex}"
|
||||
with wire_server(_analyze) as analyzer, wire_server(_anonymize) as anonymizer:
|
||||
created: Final = gateway.request(
|
||||
"POST",
|
||||
"/guardrails",
|
||||
{
|
||||
"guardrail": {
|
||||
"guardrail_name": name,
|
||||
"litellm_params": {
|
||||
"guardrail": "presidio",
|
||||
"mode": mode,
|
||||
"default_on": False,
|
||||
"presidio_analyzer_api_base": analyzer.url,
|
||||
"presidio_anonymizer_api_base": anonymizer.url,
|
||||
**({} if entities is None else {"pii_entities_config": dict(entities)}),
|
||||
},
|
||||
}
|
||||
},
|
||||
)
|
||||
assert created.status_code == 200, created.text
|
||||
try:
|
||||
yield Presidio(name, analyzer, anonymizer)
|
||||
finally:
|
||||
deleted: Final = gateway.request("DELETE", f"/guardrails/{created.json()['guardrail_id']}")
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
|
||||
|
||||
def _requested_entities(analyzer: Wire) -> list[JsonValue]:
|
||||
return [json.loads(request.body).get("entities") for request in analyzer.drain()]
|
||||
|
||||
|
||||
def test_pre_call_masks_only_the_configured_entities_before_the_provider_sees_the_prompt(gateway: Gateway) -> None:
|
||||
with (
|
||||
_presidio(gateway, "pre_call", {"CREDIT_CARD": "MASK", "EMAIL_ADDRESS": "MASK"}) as presidio,
|
||||
gateway.scenario() as scenario,
|
||||
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
|
||||
):
|
||||
model: Final = scenario.model()
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
user_text: Final = f"{uuid.uuid4().hex} card {CARD}, email {EMAIL}, phone {PHONE}"
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"guardrails": [presidio.name],
|
||||
"messages": [{"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user_text}],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
observed: Final = upstream.get("/__observations").json()["requests"]
|
||||
assert len(observed) == 1
|
||||
messages: Final = observed[0]["body"]["messages"]
|
||||
assert messages[0] == {"role": "system", "content": SYSTEM_PROMPT}
|
||||
forwarded: Final = string_value(messages[1]["content"])
|
||||
assert CARD not in forwarded and EMAIL not in forwarded, forwarded
|
||||
assert "<CREDIT_CARD>" in forwarded and "<EMAIL_ADDRESS>" in forwarded, forwarded
|
||||
assert PHONE in forwarded, forwarded
|
||||
requested: Final = _requested_entities(presidio.analyzer)
|
||||
assert requested and all(sorted(entities) == ["CREDIT_CARD", "EMAIL_ADDRESS"] for entities in requested), (
|
||||
requested
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entities", [None, {}])
|
||||
def test_apply_guardrail_with_the_default_config_masks_every_detected_entity(
|
||||
gateway: Gateway, entities: Mapping[str, str] | None
|
||||
) -> None:
|
||||
with _presidio(gateway, "pre_call", entities) as presidio:
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/guardrails/apply_guardrail",
|
||||
{"guardrail_name": presidio.name, "text": f"card {CARD} and email {EMAIL}"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
masked: Final = string_value(response.json()["response_text"])
|
||||
assert masked == "card <CREDIT_CARD> and email <EMAIL_ADDRESS>", masked
|
||||
assert _requested_entities(presidio.analyzer) == [None]
|
||||
assert len(presidio.anonymizer.drain()) == 1
|
||||
|
|
@ -0,0 +1,209 @@
|
|||
import asyncio
|
||||
import itertools
|
||||
import json
|
||||
import ssl
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import websockets
|
||||
from integration._support.client import Gateway, gateway_from_environment
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.tls import server_context, write_self_signed_cert
|
||||
from pydantic import JsonValue
|
||||
from websockets.asyncio.server import ServerConnection, serve
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(180)
|
||||
|
||||
PROVIDER_MODEL: Final = "ws-peer-model"
|
||||
USAGE: Final = {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7}
|
||||
TERMINAL: Final = frozenset({"response.completed", "response.failed", "error"})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Peer:
|
||||
url: str
|
||||
paths: SimpleQueue[str]
|
||||
frames: SimpleQueue[dict[str, JsonValue]]
|
||||
|
||||
|
||||
def _events(response_id: str, text: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
message: Final = {
|
||||
"type": "message",
|
||||
"id": f"msg_{response_id}",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": text, "annotations": []}],
|
||||
}
|
||||
response: Final = {"id": response_id, "object": "response", "created_at": 1700000000, "model": PROVIDER_MODEL}
|
||||
return (
|
||||
{"type": "response.created", "response": {**response, "status": "in_progress", "output": []}},
|
||||
{
|
||||
"type": "response.output_text.delta",
|
||||
"item_id": f"msg_{response_id}",
|
||||
"output_index": 0,
|
||||
"content_index": 0,
|
||||
"delta": text,
|
||||
},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {**response, "status": "completed", "output": [message], "usage": USAGE},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _answer(
|
||||
connection: ServerConnection, paths: SimpleQueue[str], frames: SimpleQueue[dict[str, JsonValue]]
|
||||
) -> None:
|
||||
paths.put(connection.request.path if connection.request is not None else "")
|
||||
turns: Final = itertools.count(1)
|
||||
async for raw in connection:
|
||||
frame: Final = json.loads(raw)
|
||||
frames.put(frame)
|
||||
if frame.get("type") != "response.create":
|
||||
continue
|
||||
for event in _events(f"resp_peer_{next(turns)}", "seven"):
|
||||
await connection.send(json.dumps(event))
|
||||
|
||||
|
||||
async def _serve(
|
||||
tls: ssl.SSLContext,
|
||||
paths: SimpleQueue[str],
|
||||
frames: SimpleQueue[dict[str, JsonValue]],
|
||||
ports: SimpleQueue[int],
|
||||
stop: asyncio.Event,
|
||||
) -> None:
|
||||
async with serve(lambda connection: _answer(connection, paths, frames), "127.0.0.1", 0, ssl=tls) as server:
|
||||
ports.put(next(iter(server.sockets)).getsockname()[1])
|
||||
await stop.wait()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def responses_peer(cert: tuple[Path, Path]) -> Iterator[Peer]:
|
||||
loop: Final = asyncio.new_event_loop()
|
||||
stop: Final = asyncio.Event()
|
||||
paths: Final = SimpleQueue[str]()
|
||||
frames: Final = SimpleQueue[dict[str, JsonValue]]()
|
||||
ports: Final = SimpleQueue[int]()
|
||||
thread: Final = threading.Thread(
|
||||
target=loop.run_until_complete, args=(_serve(server_context(*cert), paths, frames, ports, stop),), daemon=True
|
||||
)
|
||||
thread.start()
|
||||
try:
|
||||
yield Peer(f"https://127.0.0.1:{ports.get(timeout=10)}/v1", paths, frames)
|
||||
finally:
|
||||
loop.call_soon_threadsafe(stop.set)
|
||||
thread.join(timeout=10)
|
||||
loop.close()
|
||||
|
||||
|
||||
def _create(model: str, text: str, previous_response_id: str | None = None) -> str:
|
||||
return json.dumps(
|
||||
{
|
||||
"type": "response.create",
|
||||
"model": model,
|
||||
"store": True,
|
||||
"input": [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": text}]}],
|
||||
**({} if previous_response_id is None else {"previous_response_id": previous_response_id}),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def _turn(connection: websockets.ClientConnection, frame: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
await connection.send(frame)
|
||||
return await _until_terminal(connection, ())
|
||||
|
||||
|
||||
async def _until_terminal(
|
||||
connection: websockets.ClientConnection, received: tuple[dict[str, JsonValue], ...]
|
||||
) -> tuple[dict[str, JsonValue], ...]:
|
||||
event: Final = json.loads(await asyncio.wait_for(connection.recv(), timeout=20))
|
||||
collected: Final = (*received, event)
|
||||
if event.get("type") in TERMINAL or len(collected) >= 50:
|
||||
return collected
|
||||
return await _until_terminal(connection, collected)
|
||||
|
||||
|
||||
async def _session(
|
||||
proxy_url: str, key: str, model: str, texts: tuple[str, ...]
|
||||
) -> tuple[tuple[dict[str, JsonValue], ...], ...]:
|
||||
proxy: Final = proxy_url.rstrip("/").replace("http://", "ws://")
|
||||
async with websockets.connect(
|
||||
f"{proxy}/v1/responses?model={model}",
|
||||
additional_headers={"Authorization": f"Bearer {key}"},
|
||||
open_timeout=10,
|
||||
) as connection:
|
||||
first: Final = await _turn(connection, _create(model, texts[0]))
|
||||
if len(texts) == 1:
|
||||
return (first,)
|
||||
previous: Final = str(first[-1]["response"]["id"])
|
||||
second: Final = await _turn(connection, _create(model, texts[1], previous))
|
||||
return (first, second)
|
||||
|
||||
|
||||
def _completed(events: tuple[dict[str, JsonValue], ...]) -> dict[str, JsonValue]:
|
||||
assert events[-1]["type"] == "response.completed", [event.get("type") for event in events]
|
||||
return events[-1]["response"]
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def cert(tmp_path_factory: pytest.TempPathFactory) -> tuple[Path, Path]:
|
||||
return write_self_signed_cert(tmp_path_factory.mktemp("responses-ws-cert"))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def candidate(tmp_path_factory: pytest.TempPathFactory, cert: tuple[Path, Path]) -> Iterator[Gateway]:
|
||||
with gateway_from_environment() as base:
|
||||
with owned_proxy(base, tmp_path_factory.mktemp("responses-ws"), {"SSL_CERT_FILE": str(cert[0])}) as proxy:
|
||||
yield proxy
|
||||
|
||||
|
||||
def _drain(queue: SimpleQueue[dict[str, JsonValue]]) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(queue.get_nowait() for _ in range(queue.qsize()))
|
||||
|
||||
|
||||
def test_a_response_create_frame_streams_from_the_provider_socket_back_to_the_client(
|
||||
candidate: Gateway, cert: tuple[Path, Path]
|
||||
) -> None:
|
||||
with responses_peer(cert) as peer, candidate.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url)
|
||||
key: Final = scenario.key(models=[model])
|
||||
text: Final = f"say seven {uuid.uuid4().hex}"
|
||||
(events,) = asyncio.run(_session(str(candidate.client.base_url), key, model, (text,)))
|
||||
assert [event["type"] for event in events] == [
|
||||
"response.created",
|
||||
"response.output_text.delta",
|
||||
"response.completed",
|
||||
]
|
||||
completed: Final = _completed(events)
|
||||
assert completed["status"] == "completed"
|
||||
assert completed["usage"] == USAGE
|
||||
assert peer.paths.get_nowait() == f"/v1/responses?model={PROVIDER_MODEL}"
|
||||
(forwarded,) = _drain(peer.frames)
|
||||
assert forwarded["type"] == "response.create"
|
||||
assert forwarded["model"] == PROVIDER_MODEL
|
||||
assert forwarded["input"][0]["content"][0]["text"] == text
|
||||
|
||||
|
||||
def test_previous_response_id_from_the_first_turn_reaches_the_provider_as_its_own_id(
|
||||
candidate: Gateway, cert: tuple[Path, Path]
|
||||
) -> None:
|
||||
with responses_peer(cert) as peer, candidate.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"openai/{PROVIDER_MODEL}", api_base=peer.url)
|
||||
key: Final = scenario.key(models=[model])
|
||||
texts: Final = (f"remember seven {uuid.uuid4().hex}", f"which number {uuid.uuid4().hex}")
|
||||
first, second = asyncio.run(_session(str(candidate.client.base_url), key, model, texts))
|
||||
assert _completed(first)["status"] == "completed"
|
||||
assert str(_completed(first)["id"]).startswith("resp_") and _completed(first)["id"] != "resp_peer_1"
|
||||
assert _completed(second)["status"] == "completed"
|
||||
assert peer.paths.qsize() == 1
|
||||
forwarded: Final = _drain(peer.frames)
|
||||
assert [frame["input"][0]["content"][0]["text"] for frame in forwarded] == list(texts)
|
||||
assert "previous_response_id" not in forwarded[0]
|
||||
assert forwarded[1]["previous_response_id"] == "resp_peer_1"
|
||||
72
tests/integration/routing/test_end_user_region_routing.py
Normal file
72
tests/integration/routing/test_end_user_region_routing.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, gateway_from_environment
|
||||
from integration._support.process import owned_proxy
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(180)
|
||||
|
||||
MODEL: Final = "regional-model"
|
||||
UPSTREAM_BY_REGION: Final = {"eu": "regional-eu-upstream", "us": "regional-us-upstream"}
|
||||
CALLS: Final = 5
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def candidate(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("region-routing")
|
||||
with gateway_from_environment() as base:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
{
|
||||
"model_name": MODEL,
|
||||
"litellm_params": {
|
||||
"model": f"openai/{upstream}",
|
||||
"api_base": f"{base.upstream_url}/v1",
|
||||
"api_key": "synthetic-region-key",
|
||||
"region_name": region,
|
||||
},
|
||||
}
|
||||
for region, upstream in UPSTREAM_BY_REGION.items()
|
||||
]
|
||||
config["router_settings"] = {**config["router_settings"], "enable_pre_call_checks": True}
|
||||
path: Final = directory / "region-routing.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(base, directory, {}, config=path) as proxy:
|
||||
yield proxy
|
||||
|
||||
|
||||
@pytest.mark.parametrize("region", ["eu", "us"])
|
||||
def test_an_end_users_allowed_region_pins_every_call_to_that_regions_deployment(
|
||||
candidate: Gateway, region: str
|
||||
) -> None:
|
||||
with (
|
||||
candidate.scenario() as scenario,
|
||||
httpx.Client(base_url=candidate.upstream_url, timeout=5, trust_env=False) as upstream,
|
||||
):
|
||||
end_user: Final = f"integration-end-user-{uuid.uuid4().hex}"
|
||||
candidate.post("/end_user/new", {"user_id": end_user, "allowed_model_region": region})
|
||||
scenario.cleanups.callback(candidate.post, "/end_user/delete", {"user_ids": [end_user]})
|
||||
key: Final = scenario.key(models=[MODEL])
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
responses: Final = tuple(
|
||||
candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": MODEL,
|
||||
"user": end_user,
|
||||
"messages": [{"role": "user", "content": f"region {uuid.uuid4().hex}"}],
|
||||
},
|
||||
key=key,
|
||||
)
|
||||
for _ in range(CALLS)
|
||||
)
|
||||
assert [response.status_code for response in responses] == [200] * CALLS, [r.text for r in responses]
|
||||
assert [response.headers.get("x-litellm-model-region") for response in responses] == [region] * CALLS
|
||||
observed: Final = upstream.get("/__observations").json()["requests"]
|
||||
assert [request["body"]["model"] for request in observed] == [UPSTREAM_BY_REGION[region]] * CALLS
|
||||
27
tests/integration/routing/test_key_max_parallel_requests.py
Normal file
27
tests/integration/routing/test_key_max_parallel_requests.py
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
import uuid
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from integration._support.client import Gateway, object_value
|
||||
|
||||
|
||||
def test_zero_parallel_slots_refuse_before_the_provider_and_one_slot_serves(gateway: Gateway) -> None:
|
||||
with (
|
||||
gateway.scenario() as scenario,
|
||||
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
|
||||
):
|
||||
model: Final = scenario.model()
|
||||
blocked: Final = scenario.key(models=[model], max_parallel_requests=0)
|
||||
allowed: Final = scenario.key(models=[model], max_parallel_requests=1)
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
refused: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": f"no slots {uuid.uuid4().hex}"}]},
|
||||
key=blocked,
|
||||
)
|
||||
assert refused.status_code == 429, refused.text
|
||||
assert upstream.get("/__observations").json()["requests"] == []
|
||||
served: Final = tuple(gateway.chat(model, key=allowed, text=f"one slot {uuid.uuid4().hex}") for _ in range(2))
|
||||
assert [object_value(response["usage"])["total_tokens"] for response in served] == [40, 40]
|
||||
assert len(upstream.get("/__observations").json()["requests"]) == 2
|
||||
67
tests/integration/routing/test_team_tag_routing.py
Normal file
67
tests/integration/routing/test_team_tag_routing.py
Normal file
|
|
@ -0,0 +1,67 @@
|
|||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, gateway_from_environment
|
||||
from integration._support.process import owned_proxy
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(180)
|
||||
|
||||
MODEL: Final = "tagged-model"
|
||||
DEPLOYMENT_BY_TAG: Final = {"teamA": "team-a-deployment", "teamB": "team-b-deployment"}
|
||||
CALLS: Final = 5
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def candidate(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
|
||||
directory: Final = tmp_path_factory.mktemp("team-tag-routing")
|
||||
with gateway_from_environment() as base:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
{
|
||||
"model_name": MODEL,
|
||||
"litellm_params": {
|
||||
"model": f"openai/{deployment}",
|
||||
"api_base": f"{base.upstream_url}/v1",
|
||||
"api_key": "synthetic-tag-key",
|
||||
"tags": [tag],
|
||||
},
|
||||
"model_info": {"id": deployment},
|
||||
}
|
||||
for tag, deployment in DEPLOYMENT_BY_TAG.items()
|
||||
]
|
||||
config["router_settings"] = {**config["router_settings"], "enable_tag_filtering": True}
|
||||
path: Final = directory / "team-tag-routing.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(base, directory, {}, config=path) as proxy:
|
||||
yield proxy
|
||||
|
||||
|
||||
@pytest.mark.parametrize("tag", ["teamA", "teamB"])
|
||||
def test_a_teams_tags_route_every_call_of_its_keys_to_the_matching_deployment(candidate: Gateway, tag: str) -> None:
|
||||
with (
|
||||
candidate.scenario() as scenario,
|
||||
httpx.Client(base_url=candidate.upstream_url, timeout=5, trust_env=False) as upstream,
|
||||
):
|
||||
team_id: Final = scenario.team(tags=[tag])
|
||||
key: Final = scenario.key(team_id=team_id)
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
responses: Final = tuple(
|
||||
candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": MODEL, "messages": [{"role": "user", "content": f"tagged {uuid.uuid4().hex}"}]},
|
||||
key=key,
|
||||
)
|
||||
for _ in range(CALLS)
|
||||
)
|
||||
assert [response.status_code for response in responses] == [200] * CALLS, [r.text for r in responses]
|
||||
assert [response.headers.get("x-litellm-model-id") for response in responses] == [
|
||||
DEPLOYMENT_BY_TAG[tag]
|
||||
] * CALLS
|
||||
observed: Final = upstream.get("/__observations").json()["requests"]
|
||||
assert [request["body"]["model"] for request in observed] == [DEPLOYMENT_BY_TAG[tag]] * CALLS
|
||||
74
tests/integration/sdk/test_provider_budget_redis.py
Normal file
74
tests/integration/sdk/test_provider_budget_redis.py
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
import asyncio
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
|
||||
from litellm.types.utils import BudgetConfig
|
||||
from redis import Redis
|
||||
|
||||
WINDOWS: Final = {"openai": ("1d", 86400), "vertex_ai": ("1h", 3600)}
|
||||
SPEND_KEYS: Final = {provider: f"provider_spend:{provider}:{window}" for provider, (window, _) in WINDOWS.items()}
|
||||
START_KEYS: Final = tuple(f"provider_budget_start_time:{provider}" for provider in WINDOWS)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def redis_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[Redis]:
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), decode_responses=True) as client:
|
||||
client.delete(*SPEND_KEYS.values(), *START_KEYS)
|
||||
yield client
|
||||
client.delete(*SPEND_KEYS.values(), *START_KEYS)
|
||||
|
||||
|
||||
def _limiter() -> RouterBudgetLimiting:
|
||||
return RouterBudgetLimiting(
|
||||
dual_cache=DualCache(redis_cache=RedisCache(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]))),
|
||||
provider_budget_config={
|
||||
provider: BudgetConfig(budget_duration=window, max_budget=100) for provider, (window, _) in WINDOWS.items()
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _windows_opened(redis_client: Redis) -> bool:
|
||||
for _ in range(100):
|
||||
if all(int(redis_client.ttl(key)) > 0 for key in SPEND_KEYS.values()):
|
||||
return True
|
||||
await asyncio.sleep(0.1)
|
||||
return False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_written_to_redis_by_another_instance_is_pulled_into_memory(redis_client: Redis) -> None:
|
||||
limiter: Final = _limiter()
|
||||
assert await _windows_opened(redis_client)
|
||||
elsewhere: Final = {SPEND_KEYS["openai"]: 50.0, SPEND_KEYS["vertex_ai"]: 75.0}
|
||||
for key, value in elsewhere.items():
|
||||
redis_client.set(key, str(value), keepttl=True)
|
||||
await limiter._sync_in_memory_spend_with_redis()
|
||||
in_memory: Final = {key: await limiter.dual_cache.in_memory_cache.async_get_cache(key) for key in elsewhere}
|
||||
assert in_memory == elsewhere
|
||||
assert await limiter._get_current_provider_spend("openai") == 50.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_budget_reset_time_follows_the_redis_window_expiry(redis_client: Redis) -> None:
|
||||
limiter: Final = _limiter()
|
||||
assert await _windows_opened(redis_client)
|
||||
assert await limiter._get_current_provider_budget_reset_at("anthropic") is None
|
||||
reset_times: Final = {
|
||||
provider: await limiter._get_current_provider_budget_reset_at(provider) for provider in WINDOWS
|
||||
}
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
drift: Final = {
|
||||
provider: abs(
|
||||
(datetime.fromisoformat(str(reset_times[provider])) - (now + timedelta(seconds=seconds))).total_seconds()
|
||||
)
|
||||
for provider, (_, seconds) in WINDOWS.items()
|
||||
}
|
||||
assert all(seconds < 5 for seconds in drift.values()), (reset_times, drift)
|
||||
81
tests/integration/sdk/test_redis_service_metrics.py
Normal file
81
tests/integration/sdk/test_redis_service_metrics.py
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
import json
|
||||
import os
|
||||
import uuid
|
||||
from itertools import chain
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from litellm import Router
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from prometheus_client import REGISTRY
|
||||
|
||||
CHAT_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"id": "chatcmpl_redis_service_metrics",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "metrics"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
|
||||
}
|
||||
).encode()
|
||||
LABELS: Final = {"redis": "redis"}
|
||||
|
||||
|
||||
def _reply(request: Request) -> Reply:
|
||||
return Reply(body=CHAT_RESPONSE)
|
||||
|
||||
|
||||
def _redis_metrics() -> tuple[float, float, float]:
|
||||
failed_metrics: Final = tuple(
|
||||
metric for metric in REGISTRY.collect() if metric.name == "litellm_redis_failed_requests"
|
||||
)
|
||||
samples: Final = chain.from_iterable(metric.samples for metric in failed_metrics)
|
||||
failed: Final = sum(sample.value for sample in samples if sample.name.endswith("_total"))
|
||||
return (
|
||||
REGISTRY.get_sample_value("litellm_redis_total_requests_total", LABELS) or 0.0,
|
||||
REGISTRY.get_sample_value("litellm_redis_latency_count", LABELS) or 0.0,
|
||||
failed,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_redis_traffic_is_counted_in_the_prometheus_service_metrics(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "service_callback", ["prometheus_system"])
|
||||
with wire_server(_reply) as wire:
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "redis-metrics",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": f"{wire.url}/v1",
|
||||
"api_key": "synthetic-redis-metrics-key",
|
||||
"tpm": tpm,
|
||||
},
|
||||
}
|
||||
for tpm in (100, 1000)
|
||||
],
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
redis_host=os.environ["REDIS_HOST"],
|
||||
redis_port=int(os.environ["REDIS_PORT"]),
|
||||
)
|
||||
before: Final = _redis_metrics()
|
||||
responses: Final = [
|
||||
await router.acompletion(
|
||||
model="redis-metrics", messages=[{"role": "user", "content": f"metrics {uuid.uuid4().hex}"}]
|
||||
)
|
||||
for _ in range(2)
|
||||
]
|
||||
await GLOBAL_LOGGING_WORKER.flush()
|
||||
after: Final = _redis_metrics()
|
||||
assert [response.usage.total_tokens for response in responses] == [7, 7]
|
||||
assert len(wire.drain()) == 2
|
||||
total_delta, latency_delta, failed_delta = (now - then for now, then in zip(after, before, strict=True))
|
||||
assert total_delta > 0, (before, after)
|
||||
assert latency_delta > 0, (before, after)
|
||||
assert failed_delta == 0, (before, after)
|
||||
107
tests/integration/sdk/test_router_redis_tls_url.py
Normal file
107
tests/integration/sdk/test_router_redis_tls_url.py
Normal file
|
|
@ -0,0 +1,107 @@
|
|||
import os
|
||||
import socket
|
||||
import ssl
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.tls import server_context, write_self_signed_cert
|
||||
from litellm import Router
|
||||
from redis import Redis
|
||||
|
||||
PAYLOAD: Final = {"transport": "tls"}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TlsRelay:
|
||||
url: str
|
||||
handshakes: SimpleQueue[str]
|
||||
|
||||
|
||||
def _pipe(source: socket.socket, sink: socket.socket) -> None:
|
||||
try:
|
||||
while chunk := source.recv(65536):
|
||||
sink.sendall(chunk)
|
||||
except OSError:
|
||||
pass
|
||||
finally:
|
||||
sink.close()
|
||||
|
||||
|
||||
def _serve(listener: socket.socket, context: ssl.SSLContext, handshakes: SimpleQueue[str]) -> None:
|
||||
while True:
|
||||
try:
|
||||
raw, _ = listener.accept()
|
||||
except OSError:
|
||||
return
|
||||
try:
|
||||
secured = context.wrap_socket(raw, server_side=True)
|
||||
except (ssl.SSLError, OSError):
|
||||
raw.close()
|
||||
continue
|
||||
handshakes.put(str(secured.version()))
|
||||
backend = socket.create_connection((os.environ["REDIS_HOST"], int(os.environ["REDIS_PORT"])))
|
||||
threading.Thread(target=_pipe, args=(secured, backend), daemon=True).start()
|
||||
threading.Thread(target=_pipe, args=(backend, secured), daemon=True).start()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def tls_relay(directory: Path) -> Iterator[TlsRelay]:
|
||||
cert: Final = write_self_signed_cert(directory)
|
||||
handshakes: Final = SimpleQueue[str]()
|
||||
with socket.create_server(("127.0.0.1", 0)) as listener:
|
||||
thread: Final = threading.Thread(target=_serve, args=(listener, server_context(*cert), handshakes), daemon=True)
|
||||
thread.start()
|
||||
port: Final = listener.getsockname()[1]
|
||||
yield TlsRelay(f"rediss://127.0.0.1:{port}/0?ssl_ca_certs={cert[0]}", handshakes)
|
||||
listener.close()
|
||||
thread.join(timeout=5)
|
||||
|
||||
|
||||
def _router(redis_url: str) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "tls-cache",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "synthetic-tls-key"},
|
||||
}
|
||||
],
|
||||
redis_url=redis_url,
|
||||
)
|
||||
|
||||
|
||||
def _plain_redis() -> Redis:
|
||||
return Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), decode_responses=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_router_cache_built_from_a_rediss_url_talks_tls_to_redis(tmp_path: Path) -> None:
|
||||
with tls_relay(tmp_path) as relay, _plain_redis() as plain:
|
||||
cache: Final = _router(relay.url).cache.redis_cache
|
||||
assert cache is not None
|
||||
assert await cache.ping() is True
|
||||
key: Final = f"tls-async-{uuid.uuid4().hex}"
|
||||
await cache.async_set_cache(key, PAYLOAD, ttl=60)
|
||||
assert plain.exists(key) == 1
|
||||
assert await cache.async_get_cache(key) == PAYLOAD
|
||||
assert relay.handshakes.qsize() >= 1
|
||||
assert relay.handshakes.get_nowait().startswith("TLS")
|
||||
|
||||
|
||||
def test_sync_router_cache_built_from_a_rediss_url_talks_tls_to_redis(tmp_path: Path) -> None:
|
||||
with tls_relay(tmp_path) as relay, _plain_redis() as plain:
|
||||
cache: Final = _router(relay.url).cache.redis_cache
|
||||
assert cache is not None
|
||||
assert cache.sync_ping() is True
|
||||
key: Final = f"tls-sync-{uuid.uuid4().hex}"
|
||||
cache.set_cache(key, PAYLOAD, ttl=60)
|
||||
assert plain.exists(key) == 1
|
||||
assert cache.get_cache(key) == PAYLOAD
|
||||
assert relay.handshakes.qsize() >= 1
|
||||
assert relay.handshakes.get_nowait().startswith("TLS")
|
||||
89
tests/integration/sdk/test_slack_daily_report_redis.py
Normal file
89
tests/integration/sdk/test_slack_daily_report_redis.py
Normal file
|
|
@ -0,0 +1,89 @@
|
|||
import json
|
||||
import os
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from litellm import Router
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
from litellm.proxy._types import AlertType
|
||||
from litellm.types.integrations.slack_alerting import SlackAlertingCacheKeys
|
||||
from redis import Redis
|
||||
|
||||
REPORT_SENT_KEY: Final = SlackAlertingCacheKeys.report_sent_key.value
|
||||
FAILED_REQUESTS: Final = 3
|
||||
API_BASE: Final = "http://daily-report-upstream.invalid/v1"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def redis_client() -> Iterator[Redis]:
|
||||
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), decode_responses=True) as client:
|
||||
client.delete(REPORT_SENT_KEY)
|
||||
yield client
|
||||
client.delete(REPORT_SENT_KEY)
|
||||
|
||||
|
||||
def _accept(request: Request) -> Reply:
|
||||
return Reply(body=b"ok", content_type="text/plain")
|
||||
|
||||
|
||||
def _pod(webhook: Wire) -> SlackAlerting:
|
||||
redis_cache: Final = RedisCache(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]))
|
||||
return SlackAlerting(
|
||||
internal_usage_cache=DualCache(redis_cache=redis_cache),
|
||||
alerting=["slack"],
|
||||
alert_types=[AlertType.daily_reports],
|
||||
alerting_args={"daily_report_frequency": 0},
|
||||
default_webhook_url=webhook.url,
|
||||
)
|
||||
|
||||
|
||||
def _router(deployment_id: str) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "daily-report",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": API_BASE,
|
||||
"api_key": "synthetic-daily-report-key",
|
||||
},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_report_timestamp_one_pod_stores_in_redis_drives_the_next_pods_daily_report(
|
||||
redis_client: Redis,
|
||||
) -> None:
|
||||
deployment_id: Final = f"daily-report-{uuid.uuid4().hex}"
|
||||
failed_key: Final = f"{deployment_id}:{SlackAlertingCacheKeys.failed_requests_key.value}"
|
||||
redis_client.set(failed_key, json.dumps(FAILED_REQUESTS), ex=300)
|
||||
router: Final = _router(deployment_id)
|
||||
with wire_server(_accept) as webhook:
|
||||
first_pod: Final = _pod(webhook)
|
||||
assert await first_pod._run_scheduler_helper(llm_router=router) is False
|
||||
stored: Final = redis_client.get(REPORT_SENT_KEY)
|
||||
assert stored is not None
|
||||
first_sent: Final = json.loads(stored)
|
||||
assert isinstance(first_sent, float), stored
|
||||
await first_pod.flush_queue()
|
||||
assert webhook.drain() == ()
|
||||
|
||||
second_pod: Final = _pod(webhook)
|
||||
assert await second_pod._run_scheduler_helper(llm_router=router) is True
|
||||
await second_pod.flush_queue()
|
||||
delivered: Final = webhook.drain()
|
||||
assert len(delivered) == 1
|
||||
text: Final = json.loads(delivered[0].body)["text"]
|
||||
assert f"Failed Requests: `{FAILED_REQUESTS}`" in text, text
|
||||
assert API_BASE in text, text
|
||||
assert json.loads(redis_client.get(failed_key) or "null") == 0
|
||||
assert float(json.loads(redis_client.get(REPORT_SENT_KEY) or "null")) >= first_sent
|
||||
redis_client.delete(failed_key)
|
||||
97
tests/integration/sdk/test_usage_routing_counter_ttl.py
Normal file
97
tests/integration/sdk/test_usage_routing_counter_ttl.py
Normal file
|
|
@ -0,0 +1,97 @@
|
|||
import json
|
||||
import os
|
||||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import eventually
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from litellm import Router
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from redis import Redis
|
||||
|
||||
COUNTER_TTL_SECONDS: Final = 60
|
||||
CHAT_RESPONSE: Final = json.dumps(
|
||||
{
|
||||
"id": "chatcmpl_usage_counter_ttl",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ttl"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 7, "completion_tokens": 4, "total_tokens": 11},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _reply(request: Request) -> Reply:
|
||||
assert request.target == "/v1/chat/completions", request.target
|
||||
return Reply(body=CHAT_RESPONSE)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def redis_client() -> Iterator[Redis]:
|
||||
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"]), decode_responses=True) as client:
|
||||
yield client
|
||||
|
||||
|
||||
def _router(wire: Wire, deployment_id: str) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "usage-ttl",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": f"{wire.url}/v1",
|
||||
"api_key": "synthetic-usage-ttl-key",
|
||||
"tpm": 1440,
|
||||
},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
],
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
redis_host=os.environ["REDIS_HOST"],
|
||||
redis_port=int(os.environ["REDIS_PORT"]),
|
||||
)
|
||||
|
||||
|
||||
def _counter_ttls(redis_client: Redis, deployment_id: str) -> dict[str, int]:
|
||||
keys: Final = tuple(redis_client.scan_iter(match=f"{deployment_id}:*"))
|
||||
return {key: int(redis_client.ttl(key)) for key in keys}
|
||||
|
||||
|
||||
def _expiring(ttls: dict[str, int]) -> bool:
|
||||
kinds: Final = {key.split(":")[-2] for key in ttls}
|
||||
return "tpm" in kinds and all(0 < ttl <= COUNTER_TTL_SECONDS for ttl in ttls.values())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_usage_counters_land_in_redis_with_a_one_minute_expiry(redis_client: Redis) -> None:
|
||||
deployment_id: Final = f"usage-ttl-{uuid.uuid4().hex}"
|
||||
with wire_server(_reply) as wire:
|
||||
router: Final = _router(wire, deployment_id)
|
||||
response: Final = await router.acompletion(
|
||||
model="usage-ttl", messages=[{"role": "user", "content": f"async {uuid.uuid4().hex}"}]
|
||||
)
|
||||
assert response.usage.total_tokens == 11
|
||||
await GLOBAL_LOGGING_WORKER.flush()
|
||||
ttls: Final = eventually(
|
||||
lambda: _counter_ttls(redis_client, deployment_id), _expiring, seconds=15, return_last_on_timeout=True
|
||||
)
|
||||
assert _expiring(ttls), ttls
|
||||
assert len(wire.drain()) == 1
|
||||
|
||||
|
||||
def test_sync_usage_counters_land_in_redis_with_a_one_minute_expiry(redis_client: Redis) -> None:
|
||||
deployment_id: Final = f"usage-ttl-{uuid.uuid4().hex}"
|
||||
with wire_server(_reply) as wire:
|
||||
router: Final = _router(wire, deployment_id)
|
||||
response: Final = router.completion(
|
||||
model="usage-ttl", messages=[{"role": "user", "content": f"sync {uuid.uuid4().hex}"}]
|
||||
)
|
||||
assert response.usage.total_tokens == 11
|
||||
ttls: Final = eventually(
|
||||
lambda: _counter_ttls(redis_client, deployment_id), _expiring, seconds=15, return_last_on_timeout=True
|
||||
)
|
||||
assert _expiring(ttls), ttls
|
||||
assert len(wire.drain()) == 1
|
||||
73
tests/integration/spend/test_global_spend_report.py
Normal file
73
tests/integration/spend/test_global_spend_report.py
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
import uuid
|
||||
from hashlib import sha256
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from pydantic import JsonValue
|
||||
|
||||
COST_PER_REQUEST: Final = 20 * 0.001 + 20 * 0.002
|
||||
|
||||
|
||||
def _logged(key: str, requests: int) -> list[dict[str, JsonValue]]:
|
||||
return eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT model, to_char("startTime", \'YYYY-MM-DD\') AS day FROM "LiteLLM_SpendLogs" WHERE api_key=%s',
|
||||
(sha256(key.encode()).hexdigest(),),
|
||||
),
|
||||
lambda rows: len(rows) == requests,
|
||||
seconds=70,
|
||||
)
|
||||
|
||||
|
||||
def _team_entries(report: JsonValue, day: str, team_names: frozenset[str]) -> dict[str, dict[str, JsonValue]]:
|
||||
assert isinstance(report, list), report
|
||||
days: Final = [
|
||||
object_value(row) for row in report if string_value(object_value(row)["group_by_day"]).startswith(day)
|
||||
]
|
||||
assert len(days) == 1, report
|
||||
teams: Final = days[0]["teams"]
|
||||
assert isinstance(teams, list)
|
||||
return {
|
||||
string_value(object_value(team)["team_name"]): object_value(team)
|
||||
for team in teams
|
||||
if object_value(team)["team_name"] in team_names
|
||||
}
|
||||
|
||||
|
||||
def test_default_report_groups_each_days_spend_by_team_with_per_key_breakdown(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
busy_alias: Final = f"integration-{uuid.uuid4().hex}"
|
||||
quiet_alias: Final = f"integration-{uuid.uuid4().hex}"
|
||||
busy: Final = scenario.team(team_alias=busy_alias, models=[model])
|
||||
quiet: Final = scenario.team(team_alias=quiet_alias, models=[model])
|
||||
busy_key: Final = scenario.key(team_id=busy, models=[model])
|
||||
quiet_key: Final = scenario.key(team_id=quiet, models=[model])
|
||||
traffic: Final = tuple(
|
||||
gateway.chat(model, key=key, text=f"report {uuid.uuid4().hex}") for key in (busy_key, busy_key, quiet_key)
|
||||
)
|
||||
assert len({response["id"] for response in traffic}) == 3
|
||||
busy_rows: Final = _logged(busy_key, 2)
|
||||
_logged(quiet_key, 1)
|
||||
day: Final = string_value(busy_rows[0]["day"])
|
||||
stored_model: Final = busy_rows[0]["model"]
|
||||
response: Final = gateway.request("GET", "/global/spend/report", params={"start_date": day, "end_date": day})
|
||||
assert response.status_code == 200, response.text
|
||||
entries: Final = _team_entries(response.json(), day, frozenset({busy_alias, quiet_alias}))
|
||||
assert sorted(entries) == sorted((busy_alias, quiet_alias))
|
||||
assert float(str(entries[busy_alias]["total_spend"])) == pytest.approx(2 * COST_PER_REQUEST)
|
||||
assert float(str(entries[quiet_alias]["total_spend"])) == pytest.approx(COST_PER_REQUEST)
|
||||
breakdown: Final = entries[busy_alias]["metadata"]
|
||||
assert isinstance(breakdown, list)
|
||||
assert [
|
||||
(entry["model"], entry["api_key"], float(str(entry["spend"])), entry["total_tokens"])
|
||||
for entry in map(object_value, breakdown)
|
||||
] == [(stored_model, sha256(busy_key.encode()).hexdigest(), pytest.approx(2 * COST_PER_REQUEST), 80)]
|
||||
filtered: Final = gateway.request(
|
||||
"GET", "/global/spend/report", params={"start_date": day, "end_date": day, "team_id": quiet}
|
||||
)
|
||||
assert filtered.status_code == 200, filtered.text
|
||||
only: Final = filtered.json()
|
||||
assert len(only) == 1 and [object_value(team)["team_name"] for team in only[0]["teams"]] == [quiet_alias], only
|
||||
57
tests/integration/spend/test_image_generation_key_spend.py
Normal file
57
tests/integration/spend/test_image_generation_key_spend.py
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
import json
|
||||
from hashlib import sha256
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
|
||||
PROMPT: Final = "a scripted sea otter"
|
||||
PRICE_PER_IMAGE: Final = 0.25
|
||||
|
||||
|
||||
def _image(request: Request) -> Reply:
|
||||
assert (request.method, request.target) == ("POST", "/images/generations")
|
||||
return Reply(body=json.dumps({"created": 1700000000, "data": [{"b64_json": "aW1n"}]}).encode())
|
||||
|
||||
|
||||
def test_identical_image_generations_each_charge_the_key(gateway: Gateway) -> None:
|
||||
with wire_server(_image) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/dall-e-3",
|
||||
api_base=wire.url,
|
||||
api_key="synthetic-image-key",
|
||||
output_cost_per_image=PRICE_PER_IMAGE,
|
||||
)
|
||||
key: Final = scenario.key(models=[model])
|
||||
digest: Final = sha256(key.encode()).hexdigest()
|
||||
body: Final = {"model": model, "prompt": PROMPT, "size": "1024x1024", "n": 1}
|
||||
first: Final = gateway.request("POST", "/v1/images/generations", body, key=key)
|
||||
assert first.status_code == 200, first.text
|
||||
logged: Final = eventually(
|
||||
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,)),
|
||||
lambda rows: len(rows) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
charge: Final = float(str(logged[0]["spend"]))
|
||||
assert charge == pytest.approx(PRICE_PER_IMAGE)
|
||||
eventually(
|
||||
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)),
|
||||
lambda rows: float(str(rows[0]["spend"])) == pytest.approx(charge),
|
||||
seconds=70,
|
||||
)
|
||||
repeat: Final = gateway.request("POST", "/v1/images/generations", body, key=key)
|
||||
assert repeat.status_code == 200, repeat.text
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,)),
|
||||
lambda values: len(values) == 2,
|
||||
seconds=70,
|
||||
)
|
||||
assert [float(str(row["spend"])) for row in rows] == pytest.approx([charge, charge])
|
||||
eventually(
|
||||
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)),
|
||||
lambda values: float(str(values[0]["spend"])) == pytest.approx(2 * charge),
|
||||
seconds=70,
|
||||
)
|
||||
assert len(wire.drain()) == 2
|
||||
84
tests/integration/spend/test_key_budget_lockout.py
Normal file
84
tests/integration/spend/test_key_budget_lockout.py
Normal file
|
|
@ -0,0 +1,84 @@
|
|||
import uuid
|
||||
from hashlib import sha256
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
|
||||
|
||||
def test_an_exhausted_key_is_refused_inference_but_can_still_read_its_own_info(gateway: Gateway) -> None:
|
||||
with (
|
||||
gateway.scenario() as scenario,
|
||||
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
|
||||
):
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
key: Final = scenario.key(models=[model], max_budget=0.06)
|
||||
assert (
|
||||
object_value(gateway.chat(model, key=key, text=f"spend {uuid.uuid4().hex}")["usage"])["total_tokens"] == 40
|
||||
)
|
||||
digest: Final = sha256(key.encode()).hexdigest()
|
||||
eventually(
|
||||
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)),
|
||||
lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) >= 0.06,
|
||||
seconds=70,
|
||||
)
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
denied: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": f"over budget {uuid.uuid4().hex}"}]},
|
||||
key=key,
|
||||
)
|
||||
assert denied.status_code == 422, denied.text
|
||||
error: Final = denied.json()["error"]
|
||||
assert error["type"] == "budget_exceeded"
|
||||
assert "Budget has been exceeded!" in error["message"]
|
||||
assert upstream.get("/__observations").json()["requests"] == []
|
||||
info: Final = gateway.request("GET", "/key/info", key=key, params={"key": key})
|
||||
assert info.status_code == 200, info.text
|
||||
own: Final = object_value(info.json()["info"])
|
||||
assert float(str(own["spend"])) == pytest.approx(0.06)
|
||||
assert own["max_budget"] == 0.06
|
||||
|
||||
|
||||
def _bounded_chat(gateway: Gateway, model: str, key: str) -> httpx.Response:
|
||||
return gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"max_tokens": 20,
|
||||
"messages": [{"role": "user", "content": f"key recovery {uuid.uuid4().hex}"}],
|
||||
},
|
||||
key=key,
|
||||
)
|
||||
|
||||
|
||||
def test_raising_a_spent_keys_budget_restores_serving(gateway: Gateway) -> None:
|
||||
with (
|
||||
gateway.scenario() as scenario,
|
||||
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
|
||||
):
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
key: Final = scenario.key(models=[model], max_budget=0.06)
|
||||
first: Final = _bounded_chat(gateway, model, key)
|
||||
assert first.status_code == 200, first.text
|
||||
eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (sha256(key.encode()).hexdigest(),)
|
||||
),
|
||||
lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) >= 0.06,
|
||||
seconds=70,
|
||||
)
|
||||
eventually(lambda: _bounded_chat(gateway, model, key), lambda response: response.status_code != 200, seconds=30)
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
denied: Final = _bounded_chat(gateway, model, key)
|
||||
assert denied.status_code == 422, denied.text
|
||||
assert object_value(denied.json()["error"])["type"] == "budget_exceeded"
|
||||
assert upstream.get("/__observations").json()["requests"] == []
|
||||
gateway.post("/key/update", {"key": key, "max_budget": 1.0})
|
||||
served: Final = tuple(_bounded_chat(gateway, model, key) for _ in range(3))
|
||||
assert [response.status_code for response in served] == [200, 200, 200], [response.text for response in served]
|
||||
assert len(upstream.get("/__observations").json()["requests"]) == 3
|
||||
69
tests/integration/spend/test_spend_rollup_accuracy.py
Normal file
69
tests/integration/spend/test_spend_rollup_accuracy.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually, object_value
|
||||
|
||||
COST_PER_REQUEST: Final = 20 * 0.001 + 20 * 0.002
|
||||
FIRST_BURST: Final = 6
|
||||
SECOND_BURST: Final = 4
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Owners:
|
||||
key: str
|
||||
team_id: str
|
||||
user_id: str
|
||||
organization_id: str
|
||||
|
||||
|
||||
def _reported(gateway: Gateway, owners: Owners) -> tuple[float, float, float, float]:
|
||||
key_info: Final = object_value(gateway.get("/key/info", {"key": owners.key})["info"])
|
||||
team_info: Final = object_value(gateway.get("/team/info", {"team_id": owners.team_id})["team_info"])
|
||||
user_info: Final = object_value(gateway.get("/user/info", {"user_id": owners.user_id})["user_info"])
|
||||
organization: Final = gateway.get("/organization/info", {"organization_id": owners.organization_id})
|
||||
return (
|
||||
float(str(key_info["spend"])),
|
||||
float(str(team_info["spend"])),
|
||||
float(str(user_info["spend"])),
|
||||
float(str(organization["spend"])),
|
||||
)
|
||||
|
||||
|
||||
def _matches(observed: tuple[float, float, float, float], expected: float) -> bool:
|
||||
return all(value == pytest.approx(expected, rel=1e-9) for value in observed)
|
||||
|
||||
|
||||
def _burst(
|
||||
gateway: Gateway, model: str, owners: Owners, requests: int, total_requests: int
|
||||
) -> tuple[float, float, float, float]:
|
||||
usage: Final = tuple(
|
||||
object_value(gateway.chat(model, key=owners.key, text=f"burst {uuid.uuid4().hex}")["usage"])
|
||||
for _ in range(requests)
|
||||
)
|
||||
assert [(entry["prompt_tokens"], entry["completion_tokens"]) for entry in usage] == [(20, 20)] * requests
|
||||
return eventually(
|
||||
lambda: _reported(gateway, owners),
|
||||
lambda observed: _matches(observed, total_requests * COST_PER_REQUEST),
|
||||
seconds=70,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
|
||||
|
||||
def test_every_burst_rolls_up_exactly_to_key_team_user_and_organization(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
organization_id: Final = scenario.organization()
|
||||
team_id: Final = scenario.team(organization_id=organization_id, models=[model])
|
||||
user_id: Final = scenario.user(user_role="internal_user")
|
||||
owners: Final = Owners(
|
||||
key=scenario.key(user_id=user_id, team_id=team_id, models=[model]),
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
organization_id=organization_id,
|
||||
)
|
||||
first: Final = _burst(gateway, model, owners, FIRST_BURST, FIRST_BURST)
|
||||
assert first == pytest.approx((FIRST_BURST * COST_PER_REQUEST,) * 4, rel=1e-9), first
|
||||
both: Final = _burst(gateway, model, owners, SECOND_BURST, FIRST_BURST + SECOND_BURST)
|
||||
assert both == pytest.approx(((FIRST_BURST + SECOND_BURST) * COST_PER_REQUEST,) * 4, rel=1e-9), both
|
||||
72
tests/integration/spend/test_team_budget_enforcement.py
Normal file
72
tests/integration/spend/test_team_budget_enforcement.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
import uuid
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, Scenario, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
|
||||
TEAM_BUDGET: Final = 0.06
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ExhaustedTeam:
|
||||
scenario: Scenario
|
||||
upstream: httpx.Client
|
||||
model: str
|
||||
team_id: str
|
||||
key: str
|
||||
|
||||
|
||||
def _chat(gateway: Gateway, model: str, key: str) -> httpx.Response:
|
||||
return gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"max_tokens": 20,
|
||||
"messages": [{"role": "user", "content": f"team budget {uuid.uuid4().hex}"}],
|
||||
},
|
||||
key=key,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def exhausted(gateway: Gateway) -> Iterator[ExhaustedTeam]:
|
||||
with (
|
||||
gateway.scenario() as scenario,
|
||||
httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream,
|
||||
):
|
||||
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
|
||||
team_id: Final = scenario.team(models=[model], max_budget=TEAM_BUDGET)
|
||||
key: Final = scenario.key(team_id=team_id, models=[model], max_budget=1.0)
|
||||
first: Final = _chat(gateway, model, key)
|
||||
assert first.status_code == 200, first.text
|
||||
eventually(
|
||||
lambda: read_rows('SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id=%s', (team_id,)),
|
||||
lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) >= TEAM_BUDGET,
|
||||
seconds=70,
|
||||
)
|
||||
eventually(lambda: _chat(gateway, model, key), lambda response: response.status_code != 200, seconds=30)
|
||||
upstream.get("/__observations").raise_for_status()
|
||||
yield ExhaustedTeam(scenario, upstream, model, team_id, key)
|
||||
|
||||
|
||||
def test_the_team_budget_blocks_a_key_whose_own_budget_has_room(gateway: Gateway, exhausted: ExhaustedTeam) -> None:
|
||||
denied: Final = _chat(gateway, exhausted.model, exhausted.key)
|
||||
assert denied.status_code == 422, denied.text
|
||||
error: Final = object_value(denied.json()["error"])
|
||||
assert error["type"] == "budget_exceeded"
|
||||
assert f"Budget has been exceeded! Team={exhausted.team_id}" in str(error["message"])
|
||||
assert exhausted.upstream.get("/__observations").json()["requests"] == []
|
||||
|
||||
|
||||
def test_raising_an_exhausted_team_budget_restores_serving(gateway: Gateway, exhausted: ExhaustedTeam) -> None:
|
||||
denied: Final = _chat(gateway, exhausted.model, exhausted.key)
|
||||
assert denied.status_code == 422, denied.text
|
||||
gateway.post("/team/update", {"team_id": exhausted.team_id, "max_budget": 1.0})
|
||||
served: Final = tuple(_chat(gateway, exhausted.model, exhausted.key) for _ in range(3))
|
||||
assert [response.status_code for response in served] == [200, 200, 200], [response.text for response in served]
|
||||
assert len(exhausted.upstream.get("/__observations").json()["requests"]) == 3
|
||||
|
|
@ -83,63 +83,6 @@ async def test_completion_with_caching_bad_call():
|
|||
assert sl.mock_testing_sync_success_hook == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_with_caching():
|
||||
"""
|
||||
- Run router with usage-based-routing-v2
|
||||
- Assert success callback gets called
|
||||
"""
|
||||
try:
|
||||
|
||||
def get_openai_params():
|
||||
params = {
|
||||
"model": "gpt-4.1-nano",
|
||||
"api_key": os.environ["OPENAI_API_KEY"],
|
||||
}
|
||||
return params
|
||||
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "azure/gpt-4",
|
||||
"litellm_params": get_openai_params(),
|
||||
"tpm": 100,
|
||||
},
|
||||
{
|
||||
"model_name": "azure/gpt-4",
|
||||
"litellm_params": get_openai_params(),
|
||||
"tpm": 1000,
|
||||
},
|
||||
]
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=model_list,
|
||||
set_verbose=True,
|
||||
debug_level="DEBUG",
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
redis_host=os.environ["REDIS_HOST"],
|
||||
redis_port=os.environ["REDIS_PORT"],
|
||||
redis_password=os.environ["REDIS_PASSWORD"],
|
||||
)
|
||||
|
||||
litellm.service_callback = ["prometheus_system"]
|
||||
|
||||
sl = ServiceLogging(mock_testing=True)
|
||||
sl.prometheusServicesLogger.mock_testing = True
|
||||
router.cache.redis_cache.service_logger_obj = sl
|
||||
|
||||
messages = [{"role": "user", "content": "Hey, how's it going?"}]
|
||||
response1 = await router.acompletion(model="azure/gpt-4", messages=messages)
|
||||
response1 = await router.acompletion(model="azure/gpt-4", messages=messages)
|
||||
|
||||
assert sl.mock_testing_async_success_hook > 0
|
||||
assert sl.mock_testing_sync_failure_hook == 0
|
||||
assert sl.mock_testing_async_failure_hook == 0
|
||||
assert sl.prometheusServicesLogger.mock_testing_success_calls > 0
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"An exception occured - {str(e)}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_service_logger_db_monitoring():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -356,62 +356,6 @@ async def test_increment_spend_in_current_window():
|
|||
assert queued_op["ttl"] == ttl
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_in_memory_spend_with_redis():
|
||||
"""
|
||||
Test _sync_in_memory_spend_with_redis helper method
|
||||
|
||||
Expected behavior:
|
||||
- Push all provider spend increments to Redis
|
||||
- Fetch all current provider spend from Redis to update in-memory cache
|
||||
"""
|
||||
cleanup_redis()
|
||||
provider_budget_config = {
|
||||
"openai": BudgetConfig(time_period="1d", budget_limit=100),
|
||||
"anthropic": BudgetConfig(time_period="1d", budget_limit=200),
|
||||
}
|
||||
|
||||
provider_budget = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(
|
||||
redis_cache=RedisCache(
|
||||
host=os.getenv("REDIS_HOST"),
|
||||
port=int(os.getenv("REDIS_PORT")),
|
||||
password=os.getenv("REDIS_PASSWORD"),
|
||||
)
|
||||
),
|
||||
provider_budget_config=provider_budget_config,
|
||||
)
|
||||
|
||||
# Allow background _init_provider_budget_in_cache tasks to complete
|
||||
# before overwriting Redis values (avoids race where init overwrites with 0.0)
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
# Set some values in Redis
|
||||
spend_key_openai = "provider_spend:openai:1d"
|
||||
spend_key_anthropic = "provider_spend:anthropic:1d"
|
||||
|
||||
await provider_budget.dual_cache.redis_cache.async_set_cache(
|
||||
key=spend_key_openai, value=50.0
|
||||
)
|
||||
await provider_budget.dual_cache.redis_cache.async_set_cache(
|
||||
key=spend_key_anthropic, value=75.0
|
||||
)
|
||||
|
||||
# Test syncing with Redis
|
||||
await provider_budget._sync_in_memory_spend_with_redis()
|
||||
|
||||
# Verify in-memory cache was updated
|
||||
openai_spend = await provider_budget.dual_cache.in_memory_cache.async_get_cache(
|
||||
spend_key_openai
|
||||
)
|
||||
anthropic_spend = await provider_budget.dual_cache.in_memory_cache.async_get_cache(
|
||||
spend_key_anthropic
|
||||
)
|
||||
|
||||
assert float(openai_spend) == 50.0
|
||||
assert float(anthropic_spend) == 75.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_current_provider_spend():
|
||||
"""
|
||||
|
|
@ -446,59 +390,6 @@ async def test_get_current_provider_spend():
|
|||
assert spend == 50.5
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=6, delay=2)
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_current_provider_budget_reset_at():
|
||||
"""
|
||||
Test _get_current_provider_budget_reset_at helper method
|
||||
|
||||
Scenarios:
|
||||
1. Provider with no budget config returns None
|
||||
2. Provider with budget config but no TTL returns None
|
||||
3. Provider with budget config and TTL returns correct ISO timestamp
|
||||
"""
|
||||
cleanup_redis()
|
||||
provider_budget = RouterBudgetLimiting(
|
||||
dual_cache=DualCache(
|
||||
redis_cache=RedisCache(
|
||||
host=os.getenv("REDIS_HOST"),
|
||||
port=int(os.getenv("REDIS_PORT")),
|
||||
password=os.getenv("REDIS_PASSWORD"),
|
||||
)
|
||||
),
|
||||
provider_budget_config={
|
||||
"openai": BudgetConfig(budget_duration="1d", max_budget=100),
|
||||
"vertex_ai": BudgetConfig(budget_duration="1h", max_budget=100),
|
||||
},
|
||||
)
|
||||
|
||||
await asyncio.sleep(2)
|
||||
|
||||
# Test provider with no budget config
|
||||
reset_at = await provider_budget._get_current_provider_budget_reset_at("anthropic")
|
||||
assert reset_at is None
|
||||
|
||||
# Test provider with budget config but no TTL
|
||||
reset_at = await provider_budget._get_current_provider_budget_reset_at("openai")
|
||||
assert reset_at is not None
|
||||
reset_time = datetime.fromisoformat(reset_at.replace("Z", "+00:00"))
|
||||
expected_time = datetime.now(timezone.utc) + timedelta(seconds=(24 * 60 * 60))
|
||||
time_difference = abs((reset_time - expected_time).total_seconds())
|
||||
assert time_difference < 5
|
||||
|
||||
# Test provider with budget config and TTL
|
||||
reset_at = await provider_budget._get_current_provider_budget_reset_at("vertex_ai")
|
||||
assert reset_at is not None
|
||||
|
||||
# Verify the timestamp format and approximate time
|
||||
reset_time = datetime.fromisoformat(reset_at.replace("Z", "+00:00"))
|
||||
expected_time = datetime.now(timezone.utc) + timedelta(seconds=3600)
|
||||
|
||||
# Allow for small time differences (within 5 seconds)
|
||||
time_difference = abs((reset_time - expected_time).total_seconds())
|
||||
assert time_difference < 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deployment_budget_limits_e2e_test():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -18,61 +18,6 @@ from litellm.caching import RedisCache, RedisClusterCache
|
|||
## 2. 2 models - openai, azure - 2 diff model groups, 1 caching group
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_async_caching_with_ssl_url():
|
||||
"""
|
||||
Tests when a redis url is passed to the router, if caching is correctly setup
|
||||
"""
|
||||
try:
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
"tpm": 100000,
|
||||
"rpm": 10000,
|
||||
},
|
||||
],
|
||||
redis_url=os.getenv("REDIS_SSL_URL"),
|
||||
)
|
||||
|
||||
response = await router.cache.redis_cache.ping()
|
||||
print(f"response: {response}")
|
||||
assert response == True
|
||||
except Exception as e:
|
||||
pytest.fail(f"An exception occurred - {str(e)}")
|
||||
|
||||
|
||||
def test_router_sync_caching_with_ssl_url():
|
||||
"""
|
||||
Tests when a redis url is passed to the router, if caching is correctly setup
|
||||
"""
|
||||
try:
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-3.5-turbo",
|
||||
"litellm_params": {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
"tpm": 100000,
|
||||
"rpm": 10000,
|
||||
},
|
||||
],
|
||||
redis_url=os.getenv("REDIS_SSL_URL"),
|
||||
)
|
||||
|
||||
response = router.cache.redis_cache.sync_ping()
|
||||
print(f"response: {response}")
|
||||
assert response == True
|
||||
except Exception as e:
|
||||
pytest.fail(f"An exception occurred - {str(e)}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.flaky(retries=3, delay=1)
|
||||
async def test_acompletion_caching_on_router():
|
||||
|
|
|
|||
|
|
@ -505,159 +505,6 @@ async def test_router_completion_streaming():
|
|||
"""
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_caching_ttl():
|
||||
"""
|
||||
Confirm caching ttl's work as expected.
|
||||
|
||||
Relevant issue: https://github.com/BerriAI/litellm/issues/5609
|
||||
"""
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello, can you generate a 500 words poem?"}
|
||||
]
|
||||
model = "azure-model"
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "azure-model",
|
||||
"litellm_params": {
|
||||
"model": "azure/gpt-turbo",
|
||||
"api_key": "os.environ/AZURE_FRANCE_API_KEY",
|
||||
"api_base": "https://openai-france-1234.openai.azure.com",
|
||||
"tpm": 1440,
|
||||
"mock_response": "Hello world",
|
||||
},
|
||||
"model_info": {"id": 1},
|
||||
}
|
||||
]
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
set_verbose=False,
|
||||
redis_host=os.getenv("REDIS_HOST"),
|
||||
redis_password=os.getenv("REDIS_PASSWORD"),
|
||||
redis_port=os.getenv("REDIS_PORT"),
|
||||
)
|
||||
|
||||
assert router.cache.redis_cache is not None
|
||||
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
increment_cache_kwargs = {}
|
||||
with patch.object(
|
||||
router.cache,
|
||||
"async_increment_cache_pipeline",
|
||||
new=AsyncMock(),
|
||||
) as mock_client:
|
||||
await router.acompletion(model=model, messages=messages)
|
||||
|
||||
# Async success callbacks are dispatched to GLOBAL_LOGGING_WORKER's
|
||||
# background queue; drain it before asserting the mock was invoked.
|
||||
await GLOBAL_LOGGING_WORKER.flush()
|
||||
|
||||
# mock_client.assert_called_once()
|
||||
print(f"mock_client.call_args.kwargs: {mock_client.call_args.kwargs}")
|
||||
print(f"mock_client.call_args.args: {mock_client.call_args.args}")
|
||||
|
||||
# Get the increment_list from the first positional argument or the keyword argument
|
||||
increment_list = mock_client.call_args.kwargs.get(
|
||||
"increment_list",
|
||||
mock_client.call_args.args[0] if mock_client.call_args.args else None,
|
||||
)
|
||||
assert increment_list is not None
|
||||
assert len(increment_list) > 0
|
||||
|
||||
# Check that TTL is set to 60 for all operations
|
||||
for operation in increment_list:
|
||||
assert operation["ttl"] == 60
|
||||
|
||||
# Get the first operation for testing the redis increment
|
||||
first_operation = increment_list[0]
|
||||
increment_cache_kwargs = {
|
||||
"key": first_operation["key"],
|
||||
"value": first_operation["increment_value"],
|
||||
"ttl": first_operation["ttl"],
|
||||
}
|
||||
|
||||
## call redis async increment and check if ttl correctly set
|
||||
await router.cache.redis_cache.async_increment(**increment_cache_kwargs)
|
||||
|
||||
_redis_client = router.cache.redis_cache.init_async_client()
|
||||
|
||||
async with _redis_client as redis_client:
|
||||
current_ttl = await redis_client.ttl(increment_cache_kwargs["key"])
|
||||
|
||||
assert current_ttl >= 0
|
||||
|
||||
print(f"current_ttl: {current_ttl}")
|
||||
|
||||
|
||||
def test_router_caching_ttl_sync():
|
||||
"""
|
||||
Confirm caching ttl's work as expected.
|
||||
|
||||
Relevant issue: https://github.com/BerriAI/litellm/issues/5609
|
||||
"""
|
||||
messages = [
|
||||
{"role": "user", "content": "Hello, can you generate a 500 words poem?"}
|
||||
]
|
||||
model = "azure-model"
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "azure-model",
|
||||
"litellm_params": {
|
||||
"model": "azure/gpt-turbo",
|
||||
"api_key": "os.environ/AZURE_FRANCE_API_KEY",
|
||||
"api_base": "https://openai-france-1234.openai.azure.com",
|
||||
"tpm": 1440,
|
||||
"mock_response": "Hello world",
|
||||
},
|
||||
"model_info": {"id": 1},
|
||||
}
|
||||
]
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
routing_strategy="usage-based-routing-v2",
|
||||
set_verbose=False,
|
||||
redis_host=os.getenv("REDIS_HOST"),
|
||||
redis_password=os.getenv("REDIS_PASSWORD"),
|
||||
redis_port=os.getenv("REDIS_PORT"),
|
||||
)
|
||||
|
||||
assert router.cache.redis_cache is not None
|
||||
|
||||
increment_cache_kwargs = {}
|
||||
with patch.object(
|
||||
router.cache.redis_cache,
|
||||
"increment_cache",
|
||||
new=MagicMock(),
|
||||
) as mock_client:
|
||||
router.completion(model=model, messages=messages)
|
||||
|
||||
print(mock_client.call_args_list)
|
||||
mock_client.assert_called()
|
||||
print(f"mock_client.call_args.kwargs: {mock_client.call_args.kwargs}")
|
||||
print(f"mock_client.call_args.args: {mock_client.call_args.args}")
|
||||
|
||||
increment_cache_kwargs = {
|
||||
"key": mock_client.call_args.args[0],
|
||||
"value": mock_client.call_args.args[1],
|
||||
"ttl": mock_client.call_args.kwargs["ttl"],
|
||||
}
|
||||
|
||||
assert mock_client.call_args.kwargs["ttl"] == 60
|
||||
|
||||
## call redis async increment and check if ttl correctly set
|
||||
router.cache.redis_cache.increment_cache(**increment_cache_kwargs)
|
||||
|
||||
_redis_client = router.cache.redis_cache.redis_client
|
||||
|
||||
current_ttl = _redis_client.ttl(increment_cache_kwargs["key"])
|
||||
|
||||
assert current_ttl >= 0
|
||||
|
||||
print(f"current_ttl: {current_ttl}")
|
||||
|
||||
|
||||
def test_return_potential_deployments():
|
||||
"""
|
||||
Assert deployment at limit is filtered out
|
||||
|
|
|
|||
|
|
@ -128,8 +128,6 @@ def test_init():
|
|||
print("passed testing slack alerting init")
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def slack_alerting():
|
||||
return SlackAlerting(
|
||||
|
|
@ -326,52 +324,6 @@ async def test_daily_reports_completion(slack_alerting):
|
|||
mock_send_alert.assert_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_daily_reports_redis_cache_scheduler():
|
||||
redis_cache = RedisCache()
|
||||
slack_alerting = SlackAlerting(
|
||||
internal_usage_cache=DualCache(redis_cache=redis_cache)
|
||||
)
|
||||
|
||||
# we need this to be 0 so it actualy sends the report
|
||||
slack_alerting.alerting_args.daily_report_frequency = 0
|
||||
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-5.5",
|
||||
"litellm_params": {
|
||||
"model": "gpt-5-mini",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(slack_alerting, "send_alert", new=AsyncMock()) as mock_send_alert,
|
||||
patch.object(
|
||||
redis_cache, "async_set_cache", new=AsyncMock()
|
||||
) as mock_redis_set_cache,
|
||||
):
|
||||
# initial call - expect empty
|
||||
await slack_alerting._run_scheduler_helper(llm_router=router)
|
||||
|
||||
try:
|
||||
json.dumps(mock_redis_set_cache.call_args[0][1])
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
"Cache value can't be json dumped - {}".format(
|
||||
mock_redis_set_cache.call_args[0][1]
|
||||
)
|
||||
)
|
||||
|
||||
mock_redis_set_cache.assert_awaited_once()
|
||||
|
||||
# second call - expect empty
|
||||
await slack_alerting._run_scheduler_helper(llm_router=router)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skip(reason="Local test. Test if slack alerts are sent.")
|
||||
async def test_send_llm_exception_to_slack():
|
||||
|
|
|
|||
|
|
@ -1,241 +0,0 @@
|
|||
"""
|
||||
E2E tests for OpenAI Responses API WebSocket mode through the LiteLLM proxy.
|
||||
|
||||
Connects to ws://0.0.0.0:4000/v1/responses, sends response.create events,
|
||||
and validates the streamed response events.
|
||||
|
||||
Requires:
|
||||
- Proxy running: python -m litellm.proxy.proxy_cli --config <config> --port 4000
|
||||
- Model configured in proxy (e.g. gpt-5-mini)
|
||||
|
||||
See: https://developers.openai.com/api/docs/guides/websocket-mode/
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
# ── Configuration ─────────────────────────────────────────────────────────────
|
||||
PROXY_BASE_URL = os.environ.get("LITELLM_PROXY_BASE_URL", "ws://0.0.0.0:4000")
|
||||
PROXY_MASTER_KEY = os.environ.get("LITELLM_PROXY_KEY", "sk-1234")
|
||||
PROXY_MODEL = os.environ.get("LITELLM_PROXY_RESPONSES_MODEL", "gpt-5-mini")
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _generate_key() -> str:
|
||||
"""Generate a key for testing via proxy key/generate endpoint."""
|
||||
url = "http://0.0.0.0:4000/key/generate"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {PROXY_MASTER_KEY}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
response = httpx.post(url, headers=headers, json={}, timeout=10)
|
||||
if response.status_code != 200:
|
||||
raise Exception(
|
||||
f"Key generation failed with status: {response.status_code}. "
|
||||
"Is the proxy running?"
|
||||
)
|
||||
return response.json()["key"]
|
||||
|
||||
|
||||
def _assert_basic_response(events: list[dict], label: str = "") -> None:
|
||||
"""Assert that events contain response.created, response.completed, and usage."""
|
||||
prefix = f"[{label}] " if label else ""
|
||||
types = [e.get("type") for e in events]
|
||||
assert len(events) > 0, f"{prefix}no events received"
|
||||
assert (
|
||||
"response.created" in types
|
||||
), f"{prefix}missing response.created, got: {types}"
|
||||
assert (
|
||||
"response.completed" in types
|
||||
), f"{prefix}missing response.completed, got: {types}"
|
||||
completed = next(e for e in events if e.get("type") == "response.completed")
|
||||
resp = completed.get("response", {})
|
||||
assert (
|
||||
resp.get("status") == "completed"
|
||||
), f"{prefix}status != completed: {resp.get('status')}"
|
||||
usage = resp.get("usage", {})
|
||||
assert usage.get("input_tokens", 0) > 0, f"{prefix}input_tokens=0"
|
||||
assert usage.get("output_tokens", 0) > 0, f"{prefix}output_tokens=0"
|
||||
streaming_types = {
|
||||
"response.output_item.added",
|
||||
"response.content_part.added",
|
||||
"response.output_text.delta",
|
||||
"response.output_item.done",
|
||||
}
|
||||
found = streaming_types & set(types)
|
||||
assert found, f"{prefix}no streaming delta events found, got: {types}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_websocket_proxy_basic():
|
||||
"""
|
||||
Sends a simple response.create event to the proxy WebSocket endpoint
|
||||
and validates response.created, response.completed, and streaming events.
|
||||
"""
|
||||
try:
|
||||
import websockets
|
||||
except ImportError:
|
||||
pytest.skip("websockets not installed")
|
||||
|
||||
try:
|
||||
key = _generate_key()
|
||||
except Exception as e:
|
||||
pytest.skip(
|
||||
f"Proxy not available or key generation failed: {e}. "
|
||||
"Start proxy: python -m litellm.proxy.proxy_cli --config <config> --port 4000"
|
||||
)
|
||||
|
||||
url = f"{PROXY_BASE_URL}/v1/responses?model={PROXY_MODEL}"
|
||||
headers = {"Authorization": f"Bearer {key}"}
|
||||
events: list[dict] = []
|
||||
|
||||
try:
|
||||
async with websockets.connect(
|
||||
url, additional_headers=headers, open_timeout=5
|
||||
) as ws:
|
||||
payload = {
|
||||
"type": "response.create",
|
||||
"model": PROXY_MODEL,
|
||||
"store": False,
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "Say hello in one word."}
|
||||
],
|
||||
}
|
||||
],
|
||||
"tools": [],
|
||||
}
|
||||
await ws.send(json.dumps(payload))
|
||||
for _ in range(50):
|
||||
msg = await asyncio.wait_for(ws.recv(), timeout=15)
|
||||
event = json.loads(msg)
|
||||
events.append(event)
|
||||
if event.get("type") in (
|
||||
"response.completed",
|
||||
"response.failed",
|
||||
"error",
|
||||
):
|
||||
break
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
f"WebSocket connection failed: {e}. "
|
||||
"Ensure proxy is running and model is configured."
|
||||
)
|
||||
|
||||
_assert_basic_response(events, "proxy-basic")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_websocket_proxy_multi_turn():
|
||||
"""
|
||||
Sends two sequential response.create events with previous_response_id
|
||||
to validate multi-turn conversation over a single WebSocket.
|
||||
"""
|
||||
try:
|
||||
import websockets
|
||||
except ImportError:
|
||||
pytest.skip("websockets not installed")
|
||||
|
||||
try:
|
||||
key = _generate_key()
|
||||
except Exception as e:
|
||||
pytest.skip(
|
||||
f"Proxy not available or key generation failed: {e}. "
|
||||
"Start proxy: python -m litellm.proxy.proxy_cli --config <config> --port 4000"
|
||||
)
|
||||
|
||||
url = f"{PROXY_BASE_URL}/v1/responses?model={PROXY_MODEL}"
|
||||
headers = {"Authorization": f"Bearer {key}"}
|
||||
all_events: list[dict] = []
|
||||
completed: list[dict] = []
|
||||
first_id = None
|
||||
|
||||
try:
|
||||
async with websockets.connect(
|
||||
url, additional_headers=headers, open_timeout=5
|
||||
) as ws:
|
||||
# Turn 1
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "response.create",
|
||||
"model": PROXY_MODEL,
|
||||
"store": True,
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "input_text",
|
||||
"text": "Remember the number 7. Just say OK.",
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
for _ in range(50):
|
||||
msg = await asyncio.wait_for(ws.recv(), timeout=15)
|
||||
event = json.loads(msg)
|
||||
all_events.append(event)
|
||||
if event.get("type") == "response.completed":
|
||||
completed.append(event)
|
||||
first_id = event.get("response", {}).get("id")
|
||||
break
|
||||
if event.get("type") in ("response.failed", "error"):
|
||||
break
|
||||
|
||||
assert first_id, "Turn 1 never completed"
|
||||
|
||||
# Turn 2
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "response.create",
|
||||
"model": PROXY_MODEL,
|
||||
"store": True,
|
||||
"previous_response_id": first_id,
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "input_text",
|
||||
"text": "What number did I tell you to remember?",
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
for _ in range(50):
|
||||
msg = await asyncio.wait_for(ws.recv(), timeout=15)
|
||||
event = json.loads(msg)
|
||||
all_events.append(event)
|
||||
if event.get("type") == "response.completed":
|
||||
completed.append(event)
|
||||
break
|
||||
if event.get("type") in ("response.failed", "error"):
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
f"WebSocket multi-turn failed: {e}. "
|
||||
"Ensure proxy is running and model is configured."
|
||||
)
|
||||
|
||||
assert (
|
||||
len(completed) >= 2
|
||||
), f"Expected 2 response.completed events, got {len(completed)}"
|
||||
assert completed[1].get("response", {}).get("status") == "completed"
|
||||
|
|
@ -83,18 +83,6 @@ async def chat_completion(session, key: str, model: str):
|
|||
return response
|
||||
|
||||
|
||||
async def update_key_budget(session, key: str, max_budget: float):
|
||||
"""Helper function to update a key's max budget"""
|
||||
url = "http://0.0.0.0:4000/key/update"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {
|
||||
"key": key,
|
||||
"max_budget": max_budget,
|
||||
}
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
return await response.json()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_low_budget():
|
||||
"""
|
||||
|
|
@ -174,51 +162,6 @@ async def test_chat_completion_high_budget():
|
|||
), "Should make at least one successful call before budget exceeded"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_budget_update():
|
||||
"""
|
||||
Test that requests continue working after updating a key's budget:
|
||||
1. Create key with low budget
|
||||
2. Make calls until budget exceeded
|
||||
3. Update key with higher budget
|
||||
4. Verify calls work again
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Create key with very low budget
|
||||
key_gen = await generate_key(session=session, max_budget=0.0000000005)
|
||||
key = key_gen["key"]
|
||||
|
||||
# Make calls until budget exceeded
|
||||
calls_made = await make_calls_until_budget_exceeded(
|
||||
session=session,
|
||||
key=key,
|
||||
call_function=chat_completion,
|
||||
model="fake-openai-endpoint",
|
||||
)
|
||||
|
||||
assert (
|
||||
calls_made > 0
|
||||
), "Should make at least one successful call before budget exceeded"
|
||||
|
||||
# Update key with higher budget
|
||||
await update_key_budget(session, key, max_budget=0.001)
|
||||
|
||||
# Verify calls work again
|
||||
for _ in range(3):
|
||||
try:
|
||||
response = await chat_completion(
|
||||
session=session, key=key, model="fake-openai-endpoint"
|
||||
)
|
||||
print("response: ", response)
|
||||
assert (
|
||||
response is not None
|
||||
), "Should get valid response after budget update"
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
f"Request should succeed after budget update but got error: {e}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"field",
|
||||
[
|
||||
|
|
@ -610,112 +553,4 @@ async def test_team_budget_enforcement_cli_sso_token():
|
|||
), "Should make at least one successful call before team budget exceeded"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_and_key_budget_enforcement():
|
||||
"""
|
||||
Test budget enforcement when both team and key have budgets:
|
||||
1. Create team with low budget
|
||||
2. Create key with higher budget
|
||||
3. Verify team budget is enforced first
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Create team with very low budget
|
||||
team_response = await create_team(session=session, max_budget=0.0000000005)
|
||||
team_id = team_response["team_id"]
|
||||
|
||||
# Create key with higher budget
|
||||
key_gen = await generate_team_key(
|
||||
session=session,
|
||||
team_id=team_id,
|
||||
max_budget=0.001, # Higher than team budget
|
||||
)
|
||||
key = key_gen["key"]
|
||||
|
||||
# Make calls until budget exceeded
|
||||
calls_made = await make_calls_until_budget_exceeded(
|
||||
session=session,
|
||||
key=key,
|
||||
call_function=chat_completion,
|
||||
model="fake-openai-endpoint",
|
||||
)
|
||||
|
||||
assert (
|
||||
calls_made > 0
|
||||
), "Should make at least one successful call before team budget exceeded"
|
||||
|
||||
# Verify it was the team budget that was exceeded
|
||||
try:
|
||||
await chat_completion(
|
||||
session=session, key=key, model="fake-openai-endpoint"
|
||||
)
|
||||
except Exception as e:
|
||||
error_dict = e.body
|
||||
assert (
|
||||
"Budget has been exceeded! Team=" in error_dict["message"]
|
||||
), "Error should mention team budget being exceeded"
|
||||
|
||||
assert team_id in error_dict["message"], "Error should mention team id"
|
||||
|
||||
|
||||
async def update_team_budget(session, team_id: str, max_budget: float):
|
||||
"""Helper function to update a team's max budget"""
|
||||
url = "http://0.0.0.0:4000/team/update"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {
|
||||
"team_id": team_id,
|
||||
"max_budget": max_budget,
|
||||
}
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
return await response.json()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_budget_update():
|
||||
"""
|
||||
Test that requests continue working after updating a team's budget:
|
||||
1. Create team with low budget
|
||||
2. Create key for that team
|
||||
3. Make calls until team budget exceeded
|
||||
4. Update team with higher budget
|
||||
5. Verify calls work again
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Create team with very low budget
|
||||
team_response = await create_team(session=session, max_budget=0.0000000005)
|
||||
team_id = team_response["team_id"]
|
||||
|
||||
# Create key for team (no specific budget)
|
||||
key_gen = await generate_team_key(session=session, team_id=team_id)
|
||||
key = key_gen["key"]
|
||||
|
||||
# Make calls until budget exceeded
|
||||
calls_made = await make_calls_until_budget_exceeded(
|
||||
session=session,
|
||||
key=key,
|
||||
call_function=chat_completion,
|
||||
model="fake-openai-endpoint",
|
||||
)
|
||||
|
||||
assert (
|
||||
calls_made > 0
|
||||
), "Should make at least one successful call before team budget exceeded"
|
||||
|
||||
# Update team with higher budget
|
||||
await update_team_budget(session, team_id, max_budget=0.001)
|
||||
|
||||
# Verify calls work again
|
||||
for _ in range(3):
|
||||
try:
|
||||
response = await chat_completion(
|
||||
session=session, key=key, model="fake-openai-endpoint"
|
||||
)
|
||||
print("response: ", response)
|
||||
assert (
|
||||
response is not None
|
||||
), "Should get valid response after budget update"
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
f"Request should succeed after team budget update but got error: {e}"
|
||||
)
|
||||
|
||||
# Verify it was the team budget that was exceeded
|
||||
|
|
|
|||
|
|
@ -1,135 +0,0 @@
|
|||
# What this tests ?
|
||||
## Tests /chat/completions by generating a key and then making a chat completions request
|
||||
import pytest
|
||||
import asyncio
|
||||
import aiohttp, openai
|
||||
from openai import OpenAI, AsyncOpenAI
|
||||
from typing import Optional, List, Union
|
||||
from litellm._uuid import uuid
|
||||
|
||||
|
||||
async def generate_key(
|
||||
session,
|
||||
models=[
|
||||
"gpt-5.5",
|
||||
"text-embedding-3-small",
|
||||
"gpt-image-1",
|
||||
"fake-openai-endpoint",
|
||||
"mistral-embed",
|
||||
],
|
||||
):
|
||||
url = "http://0.0.0.0:4000/key/generate"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {
|
||||
"models": models,
|
||||
"duration": None,
|
||||
}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def chat_completion(session, key, model: Union[str, List] = "gpt-5.5"):
|
||||
url = "http://0.0.0.0:4000/chat/completions"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "user", "content": f"Hello! {str(uuid.uuid4())}"},
|
||||
],
|
||||
}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def get_otel_spans(session, key):
|
||||
url = "http://0.0.0.0:4000/otel-spans"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
async with session.get(url, headers=headers) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chat_completion_check_otel_spans():
|
||||
"""
|
||||
- Create key
|
||||
Make chat completion call
|
||||
- Create user
|
||||
make chat completion call
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key_gen = await generate_key(session=session)
|
||||
key = key_gen["key"]
|
||||
await chat_completion(session=session, key=key, model="fake-openai-endpoint")
|
||||
|
||||
await asyncio.sleep(3)
|
||||
|
||||
# /otel-spans requires proxy admin; use the master key.
|
||||
otel_spans = await get_otel_spans(session=session, key="sk-1234")
|
||||
print("otel_spans: ", otel_spans)
|
||||
|
||||
all_otel_spans = otel_spans["otel_spans"]
|
||||
spans_grouped_by_parent = otel_spans["spans_grouped_by_parent"]
|
||||
print("\n spans grouped by parent: ", spans_grouped_by_parent)
|
||||
|
||||
# The GET /otel-spans request itself produces auth spans that beat
|
||||
# the chat-completion spans on start_time, so `most_recent_parent`
|
||||
# points at the wrong trace. Pick the chat-completion trace by
|
||||
# content: it's the one carrying the full set of expected markers.
|
||||
chat_completion_markers = {
|
||||
"postgres",
|
||||
"redis",
|
||||
"raw_gen_ai_request",
|
||||
"batch_write_to_db",
|
||||
}
|
||||
parent_trace_spans = next(
|
||||
spans
|
||||
for spans in spans_grouped_by_parent.values()
|
||||
if chat_completion_markers.issubset(spans)
|
||||
)
|
||||
|
||||
print("Parent trace spans: ", parent_trace_spans)
|
||||
|
||||
# either 5 or 6 traces depending on how many redis calls were made
|
||||
assert len(parent_trace_spans) >= 5
|
||||
|
||||
# 'postgres', 'redis', 'raw_gen_ai_request', 'litellm_request', 'Received Proxy Server Request' in the span
|
||||
assert "postgres" in parent_trace_spans
|
||||
assert "redis" in parent_trace_spans
|
||||
assert "raw_gen_ai_request" in parent_trace_spans
|
||||
assert "batch_write_to_db" in parent_trace_spans
|
||||
|
|
@ -1,490 +0,0 @@
|
|||
"""
|
||||
1. Default permissions for members in a team - allowed to call /key/info and /key/health
|
||||
- Create a team, create a member in a team (role = "user")
|
||||
|
||||
|
||||
Invalid Permissions:
|
||||
- User tries creating a key with team_id = team_id -> expect to fail. Invalid Permissions
|
||||
- User tries editing a key with team_id = team_id -> expect to fail. Invalid Permissions
|
||||
- User tries deleting a key with team_id = team_id -> expect to fail. Invalid Permissions
|
||||
- User tries regenerating a key with team_id = team_id -> expect to fail. Invalid Permissions
|
||||
|
||||
Valid Permissions:
|
||||
- User tries calling /key/info with team_id, expect to get valid response
|
||||
|
||||
|
||||
|
||||
2. Permissions - members allowd to edit, delete keys but not allowed to create keys
|
||||
- Create a team with member_permissions = ["/key/update", "/key/delete", "/key/info"]
|
||||
- Create a member in the team with role = "user"
|
||||
|
||||
Valid Permissions:
|
||||
- User tries editing a key with team_id = team_id -> expect to pass. Valid Permissions
|
||||
- Note: Delete/regenerate require key ownership or team admin status, not just team member permissions
|
||||
- User tries deleting a key with team_id = team_id -> expect to fail (403) unless user owns the key or is team admin
|
||||
- User tries regenerating a key with team_id = team_id -> expect to fail (403) unless user owns the key or is team admin
|
||||
|
||||
Invalid Permissions:
|
||||
- User tries creating a key with team_id = team_id -> expect to fail. Invalid Permissions
|
||||
- User tries calling /key/info with team_id, expect to get valid response
|
||||
|
||||
|
||||
|
||||
3. Permissions - members allowed to create keys but not allowed to edit, delete keys
|
||||
- Create a team with member_permissions = ["/key/generate"]
|
||||
- Create a member in the team with role = "user"
|
||||
|
||||
Valid Permissions:
|
||||
- User tries creating a key with team_id = team_id -> expect to pass. Valid Permissions
|
||||
|
||||
Invalid Permissions:
|
||||
- User tries editing a key with team_id = team_id -> expect to fail. Invalid Permissions
|
||||
- User tries deleting a key with team_id = team_id -> expect to fail. Invalid Permissions
|
||||
- User tries regenerating a key with team_id = team_id -> expect to fail. Invalid Permissions
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import asyncio
|
||||
import aiohttp, openai
|
||||
from litellm._uuid import uuid
|
||||
import json
|
||||
from litellm.proxy._types import ProxyErrorTypes
|
||||
from typing import Optional
|
||||
|
||||
LITELLM_MASTER_KEY = "sk-1234"
|
||||
|
||||
|
||||
async def create_team(session, key, member_permissions=None):
|
||||
url = "http://0.0.0.0:4000/team/new"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {"team_member_permissions": member_permissions}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(response_text)
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def create_user(session, key, user_id, team_id=None):
|
||||
url = "http://0.0.0.0:4000/user/new"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {"user_id": user_id}
|
||||
if team_id:
|
||||
data["team_id"] = team_id
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(response_text)
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def add_team_member(session, key, team_id, user_id, role="user"):
|
||||
url = "http://0.0.0.0:4000/team/member_add"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {"team_id": team_id, "member": {"role": role, "user_id": user_id}}
|
||||
print("Adding team member with data: ", data)
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(response_text)
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def generate_key(session, key, team_id=None, user_id=None):
|
||||
url = "http://0.0.0.0:4000/key/generate"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {}
|
||||
if team_id:
|
||||
data["team_id"] = team_id
|
||||
if user_id:
|
||||
data["user_id"] = user_id
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
return {"status": status, "error": response_text}
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def key_info(session, key, key_id):
|
||||
url = f"http://0.0.0.0:4000/key/info?key={key_id}"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
async with session.get(url, headers=headers) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
return {"status": status, "error": response_text}
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def update_key(
|
||||
session: aiohttp.ClientSession,
|
||||
key: str,
|
||||
key_id: str,
|
||||
team_id: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
Update a key
|
||||
|
||||
Args:
|
||||
key: key to use for authentication
|
||||
key_id: key to update
|
||||
"""
|
||||
url = "http://0.0.0.0:4000/key/update"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {"key": key_id, "metadata": {"updated": True}}
|
||||
if team_id:
|
||||
data["team_id"] = team_id
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
return {"status": status, "error": response_text}
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def delete_key(session, key, key_id):
|
||||
url = "http://0.0.0.0:4000/key/delete"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {"keys": [key_id]}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
return {"status": status, "error": response_text}
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def regenerate_key(session, key, key_id, team_id=None):
|
||||
url = "http://0.0.0.0:4000/key/regenerate"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {"key": key_id}
|
||||
if team_id:
|
||||
data["team_id"] = team_id
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
return {"status": status, "error": response_text}
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_default_member_permissions():
|
||||
"""
|
||||
Test default permissions for members in a team - allowed to call /key/info and /key/health
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
master_key = LITELLM_MASTER_KEY
|
||||
|
||||
# Create a team
|
||||
team_data = await create_team(session=session, key=master_key)
|
||||
team_id = team_data["team_id"]
|
||||
|
||||
# create a team key
|
||||
team_key_data = await generate_key(
|
||||
session=session, key=master_key, team_id=team_id
|
||||
)
|
||||
team_key = team_key_data["key"]
|
||||
|
||||
# create a user
|
||||
user_data = await create_user(
|
||||
session=session,
|
||||
key=master_key,
|
||||
user_id=f"user_{uuid.uuid4().hex[:8]}",
|
||||
team_id=team_id,
|
||||
)
|
||||
user_id = user_data["user_id"]
|
||||
|
||||
# Create a user key
|
||||
print("New user data: ", user_data)
|
||||
|
||||
# Create a user key
|
||||
user_key_data = await generate_key(
|
||||
session=session, key=master_key, user_id=user_id
|
||||
)
|
||||
print("new user key: ", user_key_data)
|
||||
user_key = user_key_data["key"]
|
||||
|
||||
# Test invalid permissions
|
||||
# User tries creating a key with team_id
|
||||
print(
|
||||
"Regular team member trying to create a key with team_id. Expecting error."
|
||||
)
|
||||
create_result = await generate_key(
|
||||
session=session, key=user_key, team_id=team_id
|
||||
)
|
||||
print("result: ", create_result)
|
||||
assert (
|
||||
"status" in create_result and create_result["status"] == 401
|
||||
), "User should not be able to create keys for team"
|
||||
error_data = json.loads(create_result["error"])
|
||||
print("error response =", json.dumps(error_data, indent=4))
|
||||
assert (
|
||||
error_data["error"]["type"]
|
||||
== ProxyErrorTypes.team_member_permission_error.value
|
||||
), "Error should be a team member permission error"
|
||||
|
||||
# User tries editing a key with team_id
|
||||
print("Regular team member trying to edit a key with team_id. Expecting error.")
|
||||
update_result = await update_key(
|
||||
session=session, key=user_key, key_id=team_key, team_id="ATTACKER_TEAM_ID"
|
||||
)
|
||||
assert (
|
||||
"status" in update_result and update_result["status"] == 401
|
||||
), "User should not be able to update keys for team"
|
||||
error_data = json.loads(update_result["error"])
|
||||
print("error response =", json.dumps(error_data, indent=4))
|
||||
assert (
|
||||
error_data["error"]["type"]
|
||||
== ProxyErrorTypes.team_member_permission_error.value
|
||||
), "Error should be a team member permission error"
|
||||
|
||||
# User tries deleting a key with team_id
|
||||
print(
|
||||
"Regular team member trying to delete a key with team_id. Expecting error."
|
||||
)
|
||||
delete_result = await delete_key(
|
||||
session=session,
|
||||
key=user_key,
|
||||
key_id=team_key,
|
||||
)
|
||||
assert (
|
||||
"status" in delete_result and delete_result["status"] == 403
|
||||
), "User should not be able to delete keys for team"
|
||||
error_data = json.loads(delete_result["error"])
|
||||
print("error response =", json.dumps(error_data, indent=4))
|
||||
# Delete endpoint now returns 403 with authorization error, not team_member_permission_error
|
||||
assert "error" in error_data, "Error should contain error field"
|
||||
|
||||
# User tries regenerating a key with team_id
|
||||
print(
|
||||
"Regular team member trying to regenerate a key with team_id. Expecting error."
|
||||
)
|
||||
regenerate_result = await regenerate_key(
|
||||
session=session,
|
||||
key=user_key,
|
||||
key_id=team_key,
|
||||
)
|
||||
assert (
|
||||
"status" in regenerate_result and regenerate_result["status"] == 401
|
||||
), "User should not be able to regenerate keys for team"
|
||||
error_data = json.loads(regenerate_result["error"])
|
||||
print("error response =", json.dumps(error_data, indent=4))
|
||||
# Regenerate endpoint now returns 403 with authorization error, not team_member_permission_error
|
||||
assert "error" in error_data, "Error should contain error field"
|
||||
|
||||
# Test valid permissions
|
||||
# User tries calling /key/info with team_id
|
||||
print(
|
||||
"Regular team member trying to get key info with team_id. Expecting success."
|
||||
)
|
||||
info_result = await key_info(
|
||||
session=session,
|
||||
key=user_key,
|
||||
key_id=team_key,
|
||||
)
|
||||
print("info result =", info_result)
|
||||
assert "status" not in info_result, "Admin should be able to get key info"
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_edit_delete_permissions():
|
||||
"""
|
||||
Test permissions - members allowed to edit, delete keys but not allowed to create keys
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
master_key = LITELLM_MASTER_KEY
|
||||
|
||||
# Create a team with specific member permissions
|
||||
team_data = await create_team(
|
||||
session=session,
|
||||
key=master_key,
|
||||
member_permissions=["/key/update", "/key/delete", "/key/info"],
|
||||
)
|
||||
team_id = team_data["team_id"]
|
||||
|
||||
# create a user in team=team_id
|
||||
user_data = await create_user(
|
||||
session=session,
|
||||
key=master_key,
|
||||
user_id=f"user_{uuid.uuid4().hex[:8]}",
|
||||
team_id=team_id,
|
||||
)
|
||||
user_id = user_data["user_id"]
|
||||
|
||||
# Generate an admin key for the team
|
||||
admin_key_data = await generate_key(session, master_key, team_id)
|
||||
key_id = admin_key_data["key"]
|
||||
|
||||
# Create a user key
|
||||
user_key_data = await generate_key(
|
||||
session=session, key=master_key, user_id=user_id
|
||||
)
|
||||
user_key = user_key_data["key"]
|
||||
|
||||
# Test valid permissions
|
||||
# User tries editing a key with team_id
|
||||
update_result = await update_key(
|
||||
session=session, key=user_key, key_id=key_id, team_id=team_id
|
||||
)
|
||||
assert (
|
||||
"status" not in update_result
|
||||
), "User should be able to update keys for team"
|
||||
|
||||
# User tries deleting a key with team_id
|
||||
# Note: Even with /key/delete permission, users can only delete keys they own or if they're team admin
|
||||
# The delete endpoint checks ownership/team admin status, not just team member permissions
|
||||
delete_result = await delete_key(session=session, key=user_key, key_id=key_id)
|
||||
assert (
|
||||
"status" in delete_result and delete_result["status"] == 403
|
||||
), "User should not be able to delete keys they don't own (even with /key/delete permission, ownership is required)"
|
||||
|
||||
# Test invalid permissions
|
||||
# User tries creating a key with team_id
|
||||
create_result = await generate_key(
|
||||
session=session, key=user_key, team_id=team_id
|
||||
)
|
||||
assert (
|
||||
"status" in create_result and create_result["status"] != 200
|
||||
), "User should not be able to create keys for team"
|
||||
|
||||
# User tries regenerating a key with team_id
|
||||
# Note: Even with /key/regenerate permission, users can only regenerate keys they own or if they're team admin
|
||||
regenerate_result = await regenerate_key(
|
||||
session=session, key=user_key, key_id=key_id, team_id=team_id
|
||||
)
|
||||
assert (
|
||||
"status" in regenerate_result and regenerate_result["status"] == 401
|
||||
), "User should not be able to regenerate keys they don't own (even with /key/regenerate permission, ownership is required)"
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_create_permissions():
|
||||
"""
|
||||
Test permissions - members allowed to create keys but not allowed to edit, delete keys
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
master_key = LITELLM_MASTER_KEY
|
||||
|
||||
# Create a team with specific member permissions
|
||||
team_data = await create_team(
|
||||
session=session, key=master_key, member_permissions=["/key/generate"]
|
||||
)
|
||||
team_id = team_data["team_id"]
|
||||
|
||||
# Create a user in the team
|
||||
user_id = f"user_{uuid.uuid4().hex[:8]}"
|
||||
await add_team_member(
|
||||
session=session,
|
||||
key=master_key,
|
||||
team_id=team_id,
|
||||
user_id=user_id,
|
||||
role="user",
|
||||
)
|
||||
|
||||
# Generate an admin key for the team
|
||||
admin_key_data = await generate_key(
|
||||
session=session, key=master_key, team_id=team_id
|
||||
)
|
||||
admin_key = admin_key_data["key"]
|
||||
key_id = admin_key_data["key"]
|
||||
|
||||
# Create a user key
|
||||
user_key_data = await generate_key(
|
||||
session=session, key=master_key, user_id=user_id
|
||||
)
|
||||
user_key = user_key_data["key"]
|
||||
|
||||
# Test valid permissions
|
||||
# User tries creating a key with team_id
|
||||
create_result = await generate_key(
|
||||
session=session, key=user_key, team_id=team_id
|
||||
)
|
||||
print("success, user created key for team=", create_result)
|
||||
assert "key" in create_result, "User should be able to create keys for team"
|
||||
assert (
|
||||
create_result["team_id"] == team_id
|
||||
), "User should be able to create keys for team"
|
||||
assert (
|
||||
"status" not in create_result
|
||||
), "User should be able to create keys for team"
|
||||
|
||||
# Test invalid permissions
|
||||
# User tries editing a key with team_id
|
||||
update_result = await update_key(
|
||||
session=session, key=user_key, key_id=key_id, team_id=team_id
|
||||
)
|
||||
assert (
|
||||
"status" in update_result and update_result["status"] != 200
|
||||
), "User should not be able to update keys for team"
|
||||
|
||||
# User tries deleting a key with team_id
|
||||
delete_result = await delete_key(session=session, key=user_key, key_id=key_id)
|
||||
assert (
|
||||
"status" in delete_result and delete_result["status"] == 403
|
||||
), "User should not be able to delete keys for team"
|
||||
|
||||
# User tries regenerating a key with team_id
|
||||
# User doesn't have /key/regenerate permission, so should get 401 (team member permission error)
|
||||
regenerate_result = await regenerate_key(
|
||||
session=session, key=user_key, key_id=key_id, team_id=team_id
|
||||
)
|
||||
assert (
|
||||
"status" in regenerate_result and regenerate_result["status"] == 401
|
||||
), "User should not be able to regenerate keys for team (no /key/regenerate permission)"
|
||||
error_data = json.loads(regenerate_result["error"])
|
||||
assert (
|
||||
error_data["error"]["type"]
|
||||
== ProxyErrorTypes.team_member_permission_error.value
|
||||
), "Error should be a team member permission error"
|
||||
|
|
@ -36,45 +36,6 @@ async def chat_completion(
|
|||
return await response.json(), response.headers
|
||||
|
||||
|
||||
async def create_team_with_tags(session, key, tags: List[str]):
|
||||
url = "http://0.0.0.0:4000/team/new"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {
|
||||
"tags": tags,
|
||||
}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(response_text)
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def create_key_with_team(session, key, team_id: str):
|
||||
url = f"http://0.0.0.0:4000/key/generate"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {
|
||||
"team_id": team_id,
|
||||
}
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(response_text)
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def model_info_get_call(session, key, model_id: str):
|
||||
# make get call pass "litellm_model_id" in query params
|
||||
url = f"http://0.0.0.0:4000/model/info?litellm_model_id={model_id}"
|
||||
|
|
@ -92,45 +53,6 @@ async def model_info_get_call(session, key, model_id: str):
|
|||
return await response.json()
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_team_tag_routing():
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key = LITELLM_MASTER_KEY
|
||||
team_a_data = await create_team_with_tags(session, key, ["teamA"])
|
||||
print("team_a_data=", team_a_data)
|
||||
team_a_id = team_a_data["team_id"]
|
||||
|
||||
team_b_data = await create_team_with_tags(session, key, ["teamB"])
|
||||
print("team_b_data=", team_b_data)
|
||||
team_b_id = team_b_data["team_id"]
|
||||
|
||||
key_with_team_a = await create_key_with_team(session, key, team_a_id)
|
||||
print("key_with_team_a=", key_with_team_a)
|
||||
_key_with_team_a = key_with_team_a["key"]
|
||||
for _ in range(5):
|
||||
response_a, headers = await chat_completion(
|
||||
session=session, key=_key_with_team_a
|
||||
)
|
||||
|
||||
headers = dict(headers)
|
||||
print(response_a)
|
||||
print(headers)
|
||||
assert (
|
||||
headers["x-litellm-model-id"] == "team-a-model"
|
||||
), "Model ID should be teamA"
|
||||
|
||||
key_with_team_b = await create_key_with_team(session, key, team_b_id)
|
||||
_key_with_team_b = key_with_team_b["key"]
|
||||
for _ in range(5):
|
||||
response_b, headers = await chat_completion(session, _key_with_team_b)
|
||||
headers = dict(headers)
|
||||
print(response_b)
|
||||
print(headers)
|
||||
assert (
|
||||
headers["x-litellm-model-id"] == "team-b-model"
|
||||
), "Model ID should be teamB"
|
||||
|
||||
|
||||
@pytest.mark.asyncio()
|
||||
async def test_chat_completion_with_no_tags():
|
||||
async with aiohttp.ClientSession() as session:
|
||||
|
|
|
|||
|
|
@ -1,395 +0,0 @@
|
|||
import pytest
|
||||
import asyncio
|
||||
import aiohttp
|
||||
import time
|
||||
|
||||
import litellm
|
||||
from litellm._uuid import uuid
|
||||
|
||||
"""
|
||||
Tests to run
|
||||
|
||||
Basic Tests:
|
||||
1. Basic Spend Accuracy Test:
|
||||
- Make N requests, compute expected total spend locally from each response's usage
|
||||
- Poll until batch writer has flushed spend to the DB
|
||||
- Expect spend for Key, Team, User, Org (/info endpoints) to equal the computed total
|
||||
|
||||
2. Long term spend accuracy test (with 2 bursts of requests)
|
||||
- Burst 1: compute expected from responses, verify
|
||||
- Burst 2: compute expected from responses, verify total = burst1 + burst2
|
||||
|
||||
Additional Test Scenarios:
|
||||
|
||||
3. Concurrent Request Accuracy Test:
|
||||
- Make 20 concurrent requests
|
||||
- Check for race conditions in spend tracking
|
||||
|
||||
4. Error Case Test:
|
||||
- Make 10 successful requests
|
||||
- Make 5 failed requests
|
||||
- Verify spend is only counted for successful requests
|
||||
|
||||
5. Mixed Request Type Test:
|
||||
- Make different types of requests with varying costs
|
||||
- Verify accurate total spend calculation
|
||||
"""
|
||||
|
||||
# Upstream model the proxy is configured with (spend_tracking_config.yaml).
|
||||
# The proxy computes spend using this model's pricing; the local ground-truth
|
||||
# calculation uses the same pricing table via litellm.cost_per_token.
|
||||
UPSTREAM_MODEL = "gpt-5-mini"
|
||||
|
||||
# Batch writer flush cadence in CI is ~2-7s (PROXY_BATCH_WRITE_AT=2 + up to 5s jitter).
|
||||
# Poll every 2s for 60s — plenty of headroom for multiple ticks to land.
|
||||
POLL_INTERVAL_SECONDS = 2
|
||||
POLL_TIMEOUT_SECONDS = 60
|
||||
|
||||
TOLERANCE = 1e-10
|
||||
|
||||
|
||||
def _make_test_session() -> aiohttp.ClientSession:
|
||||
"""
|
||||
Session tuned for CI reliability:
|
||||
- force_close: avoid aiohttp reusing a TCP connection that the proxy/kernel
|
||||
silently closed during the long idle window between setup POSTs and the
|
||||
later poll loop (observed failure mode: ConnectionTimeoutError on the
|
||||
first /key/info call after 20 chat completions).
|
||||
- explicit connect timeout: surface a blocked proxy event loop quickly
|
||||
instead of hanging on aiohttp's 5-minute default total timeout.
|
||||
"""
|
||||
return aiohttp.ClientSession(
|
||||
connector=aiohttp.TCPConnector(force_close=True),
|
||||
timeout=aiohttp.ClientTimeout(total=30, connect=10),
|
||||
)
|
||||
|
||||
|
||||
async def create_organization(session, organization_alias: str):
|
||||
"""Helper function to create a new organization"""
|
||||
url = "http://0.0.0.0:4000/organization/new"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {"organization_alias": organization_alias}
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def create_team(session, org_id: str):
|
||||
"""Helper function to create a new team under an organization"""
|
||||
url = "http://0.0.0.0:4000/team/new"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {"organization_id": org_id, "team_alias": f"test-team-{uuid.uuid4()}"}
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def create_user(session, org_id: str):
|
||||
"""Helper function to create a new user"""
|
||||
url = "http://0.0.0.0:4000/user/new"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {"user_name": f"test-user-{uuid.uuid4()}"}
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def generate_key(session, user_id: str, team_id: str):
|
||||
"""Helper function to generate a key for a specific user and team"""
|
||||
url = "http://0.0.0.0:4000/key/generate"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {"user_id": user_id, "team_id": team_id}
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def chat_completion(session, key: str):
|
||||
"""Make a chat completion request"""
|
||||
from openai import AsyncOpenAI
|
||||
from litellm._uuid import uuid
|
||||
|
||||
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000/v1")
|
||||
|
||||
response = await client.chat.completions.create(
|
||||
model="fake-openai-endpoint",
|
||||
messages=[{"role": "user", "content": f"Test message {uuid.uuid4()}"}],
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
async def get_spend_info(session, entity_type: str, entity_id: str):
|
||||
"""Helper function to get spend information for an entity"""
|
||||
url = f"http://0.0.0.0:4000/{entity_type}/info"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
if entity_type == "key":
|
||||
data = {"key": entity_id}
|
||||
else:
|
||||
data = {f"{entity_type}_id": entity_id}
|
||||
|
||||
async with session.get(url, headers=headers, params=data) as response:
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def get_proxy_readiness(session):
|
||||
"""Fetch authenticated readiness details. Used both as a fail-fast gate and as a diagnostic on poll timeout."""
|
||||
url = "http://0.0.0.0:4000/health/readiness/details"
|
||||
headers = {"Authorization": "Bearer sk-1234"}
|
||||
async with session.get(url, headers=headers) as response:
|
||||
return response.status, await response.json()
|
||||
|
||||
|
||||
async def assert_proxy_healthy(session):
|
||||
"""Fail fast if the proxy's DB or cache is not reachable — no point running the test."""
|
||||
status, body = await get_proxy_readiness(session)
|
||||
if status != 200 or body.get("db") != "connected":
|
||||
pytest.fail(
|
||||
f"Proxy /health/readiness/details unhealthy (status={status}). "
|
||||
f"Cannot run spend accuracy test. Response: {body}"
|
||||
)
|
||||
print(f"Proxy readiness OK: {body}")
|
||||
|
||||
|
||||
def compute_expected_spend(responses) -> float:
|
||||
"""
|
||||
Compute the expected total spend locally from each response's usage tokens,
|
||||
using the same pricing table the proxy uses. This is the independent ground
|
||||
truth we compare the proxy's reported spend against.
|
||||
"""
|
||||
total = 0.0
|
||||
for r in responses:
|
||||
usage = r.usage
|
||||
prompt_cost, completion_cost = litellm.cost_per_token(
|
||||
model=UPSTREAM_MODEL,
|
||||
prompt_tokens=usage.prompt_tokens,
|
||||
completion_tokens=usage.completion_tokens,
|
||||
)
|
||||
total += prompt_cost + completion_cost
|
||||
return total
|
||||
|
||||
|
||||
async def poll_key_spend_until(session, key: str, expected: float) -> float:
|
||||
"""
|
||||
Poll key spend until it matches `expected` within TOLERANCE, or timeout.
|
||||
Returns the last observed spend either way; caller decides how to report.
|
||||
"""
|
||||
start = time.time()
|
||||
last_spend = 0.0
|
||||
while time.time() - start < POLL_TIMEOUT_SECONDS:
|
||||
try:
|
||||
key_info = await get_spend_info(session, "key", key)
|
||||
except (aiohttp.ClientError, asyncio.TimeoutError) as exc:
|
||||
print(
|
||||
f"Transient transport error during spend poll: "
|
||||
f"{type(exc).__name__}: {exc}. Retrying... "
|
||||
f"({time.time() - start:.1f}s elapsed)"
|
||||
)
|
||||
await asyncio.sleep(POLL_INTERVAL_SECONDS)
|
||||
continue
|
||||
last_spend = key_info["info"]["spend"]
|
||||
if abs(last_spend - expected) < TOLERANCE:
|
||||
print(
|
||||
f"Key spend reached expected {expected} after {time.time() - start:.1f}s"
|
||||
)
|
||||
return last_spend
|
||||
print(
|
||||
f"Key spend {last_spend}, expected {expected}, waiting... "
|
||||
f"({time.time() - start:.1f}s elapsed)"
|
||||
)
|
||||
await asyncio.sleep(POLL_INTERVAL_SECONDS)
|
||||
return last_spend
|
||||
|
||||
|
||||
async def fail_with_diagnostics(session, stage: str, expected: float, observed: float):
|
||||
"""Emit a failure with readiness state so CI output points at the real cause."""
|
||||
_, readiness = await get_proxy_readiness(session)
|
||||
pytest.fail(
|
||||
f"{stage}: key spend did not match expected after {POLL_TIMEOUT_SECONDS}s poll. "
|
||||
f"expected={expected}, observed={observed}, diff={expected - observed}. "
|
||||
f"Proxy readiness: {readiness}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_basic_spend_accuracy():
|
||||
"""
|
||||
Test basic spend accuracy across different entities:
|
||||
1. Create org, team, user, and key
|
||||
2. Make N requests, keeping each response
|
||||
3. Compute expected spend locally from response usage (independent ground truth)
|
||||
4. Poll until proxy-reported spend matches expected
|
||||
5. Verify spend is consistent across key, team, user, and org entities
|
||||
"""
|
||||
NUM_LLM_REQUESTS = 20
|
||||
|
||||
async with _make_test_session() as session:
|
||||
await assert_proxy_healthy(session)
|
||||
|
||||
org_response = await create_organization(
|
||||
session=session, organization_alias=f"test-org-{uuid.uuid4()}"
|
||||
)
|
||||
print("org_response: ", org_response)
|
||||
org_id = org_response["organization_id"]
|
||||
|
||||
team_response = await create_team(session, org_id)
|
||||
print("team_response: ", team_response)
|
||||
team_id = team_response["team_id"]
|
||||
|
||||
user_response = await create_user(session, org_id)
|
||||
print("user_response: ", user_response)
|
||||
user_id = user_response["user_id"]
|
||||
|
||||
key_response = await generate_key(session, user_id, team_id)
|
||||
print("key_response: ", key_response)
|
||||
key = key_response["key"]
|
||||
|
||||
responses = []
|
||||
for i in range(NUM_LLM_REQUESTS):
|
||||
response = await chat_completion(session, key)
|
||||
responses.append(response)
|
||||
print(f"Request {i + 1}/{NUM_LLM_REQUESTS} completed")
|
||||
|
||||
expected_spend = compute_expected_spend(responses)
|
||||
assert expected_spend > 0, (
|
||||
f"Locally computed expected spend is {expected_spend}. Either cost calc "
|
||||
f"is broken or upstream returned zero tokens. "
|
||||
f"Usage: {[r.usage.model_dump() for r in responses]}"
|
||||
)
|
||||
print(f"Expected total spend (local ground truth): {expected_spend}")
|
||||
|
||||
final_spend = await poll_key_spend_until(session, key, expected_spend)
|
||||
if abs(final_spend - expected_spend) >= TOLERANCE:
|
||||
await fail_with_diagnostics(
|
||||
session,
|
||||
stage="test_basic_spend_accuracy",
|
||||
expected=expected_spend,
|
||||
observed=final_spend,
|
||||
)
|
||||
|
||||
# Allow a final scheduler tick for team/user/org aggregations to settle
|
||||
await asyncio.sleep(5)
|
||||
|
||||
key_info = await get_spend_info(session, "key", key)
|
||||
print("key_info: ", key_info)
|
||||
team_info = await get_spend_info(session, "team", team_id)
|
||||
print("team_info: ", team_info)
|
||||
user_info = await get_spend_info(session, "user", user_id)
|
||||
print("user_info: ", user_info)
|
||||
org_info = await get_spend_info(session, "organization", org_id)
|
||||
print("org_info: ", org_info)
|
||||
|
||||
assert (
|
||||
abs(key_info["info"]["spend"] - expected_spend) < TOLERANCE
|
||||
), f"Key spend {key_info['info']['spend']} does not match expected {expected_spend}"
|
||||
|
||||
assert (
|
||||
abs(user_info["user_info"]["spend"] - expected_spend) < TOLERANCE
|
||||
), f"User spend {user_info['user_info']['spend']} does not match expected {expected_spend}"
|
||||
|
||||
assert (
|
||||
abs(team_info["team_info"]["spend"] - expected_spend) < TOLERANCE
|
||||
), f"Team spend {team_info['team_info']['spend']} does not match expected {expected_spend}"
|
||||
|
||||
assert (
|
||||
abs(org_info["spend"] - expected_spend) < TOLERANCE
|
||||
), f"Organization spend {org_info['spend']} does not match expected {expected_spend}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_long_term_spend_accuracy_with_bursts():
|
||||
"""
|
||||
Test long-term spend accuracy with multiple bursts of requests:
|
||||
1. Create org, team, user, and key
|
||||
2. Burst 1: make requests, compute expected locally, verify proxy matches
|
||||
3. Burst 2: make more requests, verify proxy total == burst1 + burst2
|
||||
4. Verify total spend is consistent across all entities
|
||||
"""
|
||||
BURST_1_REQUESTS = 22
|
||||
BURST_2_REQUESTS = 12
|
||||
|
||||
async with _make_test_session() as session:
|
||||
await assert_proxy_healthy(session)
|
||||
|
||||
org_response = await create_organization(
|
||||
session=session, organization_alias=f"test-org-{uuid.uuid4()}"
|
||||
)
|
||||
print("org_response: ", org_response)
|
||||
org_id = org_response["organization_id"]
|
||||
|
||||
team_response = await create_team(session, org_id)
|
||||
print("team_response: ", team_response)
|
||||
team_id = team_response["team_id"]
|
||||
|
||||
user_response = await create_user(session, org_id)
|
||||
print("user_response: ", user_response)
|
||||
user_id = user_response["user_id"]
|
||||
|
||||
key_response = await generate_key(session, user_id, team_id)
|
||||
print("key_response: ", key_response)
|
||||
key = key_response["key"]
|
||||
|
||||
print(f"Starting first burst of {BURST_1_REQUESTS} requests...")
|
||||
burst_1_responses = []
|
||||
for i in range(BURST_1_REQUESTS):
|
||||
response = await chat_completion(session, key)
|
||||
burst_1_responses.append(response)
|
||||
print(f"Burst 1 - Request {i + 1}/{BURST_1_REQUESTS} completed")
|
||||
|
||||
burst_1_expected = compute_expected_spend(burst_1_responses)
|
||||
assert burst_1_expected > 0, (
|
||||
f"Burst 1 expected spend is {burst_1_expected}. "
|
||||
f"Usage: {[r.usage.model_dump() for r in burst_1_responses]}"
|
||||
)
|
||||
print(f"Burst 1 expected spend: {burst_1_expected}")
|
||||
|
||||
final_burst_1 = await poll_key_spend_until(session, key, burst_1_expected)
|
||||
if abs(final_burst_1 - burst_1_expected) >= TOLERANCE:
|
||||
await fail_with_diagnostics(
|
||||
session,
|
||||
stage="test_long_term_spend_accuracy burst 1",
|
||||
expected=burst_1_expected,
|
||||
observed=final_burst_1,
|
||||
)
|
||||
|
||||
print(f"Starting second burst of {BURST_2_REQUESTS} requests...")
|
||||
burst_2_responses = []
|
||||
for i in range(BURST_2_REQUESTS):
|
||||
response = await chat_completion(session, key)
|
||||
burst_2_responses.append(response)
|
||||
print(f"Burst 2 - Request {i + 1}/{BURST_2_REQUESTS} completed")
|
||||
|
||||
total_expected = burst_1_expected + compute_expected_spend(burst_2_responses)
|
||||
print(f"Total expected spend (burst 1 + burst 2): {total_expected}")
|
||||
|
||||
final_total = await poll_key_spend_until(session, key, total_expected)
|
||||
if abs(final_total - total_expected) >= TOLERANCE:
|
||||
await fail_with_diagnostics(
|
||||
session,
|
||||
stage="test_long_term_spend_accuracy total",
|
||||
expected=total_expected,
|
||||
observed=final_total,
|
||||
)
|
||||
|
||||
await asyncio.sleep(5)
|
||||
|
||||
key_info = await get_spend_info(session, "key", key)
|
||||
team_info = await get_spend_info(session, "team", team_id)
|
||||
user_info = await get_spend_info(session, "user", user_id)
|
||||
org_info = await get_spend_info(session, "organization", org_id)
|
||||
|
||||
print(f"Final key spend: {key_info['info']['spend']}")
|
||||
print(f"Final team spend: {team_info['team_info']['spend']}")
|
||||
print(f"Final user spend: {user_info['user_info']['spend']}")
|
||||
print(f"Final org spend: {org_info['spend']}")
|
||||
|
||||
assert (
|
||||
abs(key_info["info"]["spend"] - total_expected) < TOLERANCE
|
||||
), f"Key spend {key_info['info']['spend']} does not match expected {total_expected}"
|
||||
|
||||
assert (
|
||||
abs(user_info["user_info"]["spend"] - total_expected) < TOLERANCE
|
||||
), f"User spend {user_info['user_info']['spend']} does not match expected {total_expected}"
|
||||
|
||||
assert (
|
||||
abs(team_info["team_info"]["spend"] - total_expected) < TOLERANCE
|
||||
), f"Team spend {team_info['team_info']['spend']} does not match expected {total_expected}"
|
||||
|
||||
assert (
|
||||
abs(org_info["spend"] - total_expected) < TOLERANCE
|
||||
), f"Organization spend {org_info['spend']} does not match expected {total_expected}"
|
||||
|
|
@ -1,311 +0,0 @@
|
|||
import pytest
|
||||
import asyncio
|
||||
import aiohttp
|
||||
import json
|
||||
from openai import AsyncOpenAI
|
||||
from litellm._uuid import uuid
|
||||
from httpx import AsyncClient
|
||||
import os
|
||||
|
||||
TEST_MASTER_KEY = "sk-1234"
|
||||
PROXY_BASE_URL = "http://0.0.0.0:4000"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_model_alias():
|
||||
"""
|
||||
Test model alias functionality with teams:
|
||||
1. Add a new model with model_name="gpt-4-team1" and litellm_params.model="gpt-4o"
|
||||
2. Create a new team
|
||||
3. Update team with model_alias mapping
|
||||
4. Generate key for team
|
||||
5. Make request with aliased model name
|
||||
"""
|
||||
client = AsyncClient(base_url=PROXY_BASE_URL)
|
||||
headers = {"Authorization": f"Bearer {TEST_MASTER_KEY}"}
|
||||
|
||||
# Add new model
|
||||
model_response = await client.post(
|
||||
"/model/new",
|
||||
json={
|
||||
"model_name": "gpt-4o-team1",
|
||||
"litellm_params": {
|
||||
"model": "gpt-4o",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
assert model_response.status_code == 200
|
||||
|
||||
# Create new team
|
||||
team_response = await client.post(
|
||||
"/team/new",
|
||||
json={
|
||||
"models": ["gpt-4o-team1"],
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
assert team_response.status_code == 200
|
||||
team_data = team_response.json()
|
||||
team_id = team_data["team_id"]
|
||||
|
||||
# Update team with model alias
|
||||
update_response = await client.post(
|
||||
"/team/update",
|
||||
json={"team_id": team_id, "model_aliases": {"gpt-4o": "gpt-4o-team1"}},
|
||||
headers=headers,
|
||||
)
|
||||
assert update_response.status_code == 200
|
||||
|
||||
# Generate key for team
|
||||
key_response = await client.post(
|
||||
"/key/generate", json={"team_id": team_id}, headers=headers
|
||||
)
|
||||
assert key_response.status_code == 200
|
||||
key = key_response.json()["key"]
|
||||
|
||||
# Make request with model alias
|
||||
openai_client = AsyncOpenAI(api_key=key, base_url=f"{PROXY_BASE_URL}/v1")
|
||||
|
||||
response = await openai_client.chat.completions.create(
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": f"Test message {uuid.uuid4()}"}],
|
||||
)
|
||||
|
||||
assert response is not None, "Should get valid response when using model alias"
|
||||
|
||||
# Cleanup - delete the model
|
||||
model_id = model_response.json()["model_info"]["id"]
|
||||
delete_response = await client.post(
|
||||
"/model/delete",
|
||||
json={"id": model_id},
|
||||
headers={"Authorization": f"Bearer {TEST_MASTER_KEY}"},
|
||||
)
|
||||
assert delete_response.status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_model_association():
|
||||
"""
|
||||
Test that models created with a team_id are properly associated with the team:
|
||||
1. Create a new team
|
||||
2. Add a model with team_id in model_info
|
||||
3. Verify the model appears in team info
|
||||
"""
|
||||
client = AsyncClient(base_url=PROXY_BASE_URL)
|
||||
headers = {"Authorization": f"Bearer {TEST_MASTER_KEY}"}
|
||||
|
||||
# Create new team
|
||||
team_response = await client.post(
|
||||
"/team/new",
|
||||
json={
|
||||
"models": [], # Start with empty model list
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
assert team_response.status_code == 200
|
||||
team_data = team_response.json()
|
||||
team_id = team_data["team_id"]
|
||||
|
||||
# Add new model with team_id
|
||||
model_response = await client.post(
|
||||
"/model/new",
|
||||
json={
|
||||
"model_name": "gpt-4-team-test",
|
||||
"litellm_params": {
|
||||
"model": "gpt-4",
|
||||
"custom_llm_provider": "openai",
|
||||
"api_key": "fake_key",
|
||||
},
|
||||
"model_info": {"team_id": team_id},
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
assert model_response.status_code == 200
|
||||
|
||||
# Get team info and verify model association
|
||||
team_info_response = await client.get(
|
||||
f"/team/info",
|
||||
headers=headers,
|
||||
params={"team_id": team_id},
|
||||
)
|
||||
assert team_info_response.status_code == 200
|
||||
team_info = team_info_response.json()["team_info"]
|
||||
|
||||
print("team_info", json.dumps(team_info, indent=4))
|
||||
|
||||
# Verify the model is in team_models
|
||||
assert (
|
||||
"gpt-4-team-test" in team_info["models"]
|
||||
), "Model should be associated with team"
|
||||
|
||||
# Cleanup - delete the model
|
||||
model_id = model_response.json()["model_info"]["id"]
|
||||
delete_response = await client.post(
|
||||
"/model/delete",
|
||||
json={"id": model_id},
|
||||
headers=headers,
|
||||
)
|
||||
assert delete_response.status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_model_visibility_in_models_endpoint():
|
||||
"""
|
||||
Test that team-specific models are only visible to the correct team in /models endpoint:
|
||||
1. Create two teams
|
||||
2. Add a model associated with team1
|
||||
3. Generate keys for both teams
|
||||
4. Verify team1's key can see the model in /models
|
||||
5. Verify team2's key cannot see the model in /models
|
||||
"""
|
||||
client = AsyncClient(base_url=PROXY_BASE_URL)
|
||||
headers = {"Authorization": f"Bearer {TEST_MASTER_KEY}"}
|
||||
|
||||
# Create team1
|
||||
team1_response = await client.post(
|
||||
"/team/new",
|
||||
json={"models": []},
|
||||
headers=headers,
|
||||
)
|
||||
assert team1_response.status_code == 200
|
||||
team1_id = team1_response.json()["team_id"]
|
||||
|
||||
# Create team2
|
||||
team2_response = await client.post(
|
||||
"/team/new",
|
||||
json={"models": []},
|
||||
headers=headers,
|
||||
)
|
||||
assert team2_response.status_code == 200
|
||||
team2_id = team2_response.json()["team_id"]
|
||||
|
||||
# Add model associated with team1
|
||||
model_response = await client.post(
|
||||
"/model/new",
|
||||
json={
|
||||
"model_name": "gpt-4-team-test",
|
||||
"litellm_params": {
|
||||
"model": "gpt-4",
|
||||
"custom_llm_provider": "openai",
|
||||
"api_key": "fake_key",
|
||||
},
|
||||
"model_info": {"team_id": team1_id},
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
assert model_response.status_code == 200
|
||||
|
||||
# Generate keys for both teams
|
||||
team1_key = (
|
||||
await client.post("/key/generate", json={"team_id": team1_id}, headers=headers)
|
||||
).json()["key"]
|
||||
team2_key = (
|
||||
await client.post("/key/generate", json={"team_id": team2_id}, headers=headers)
|
||||
).json()["key"]
|
||||
|
||||
# Check models visibility for team1's key
|
||||
team1_models = await client.get(
|
||||
"/models", headers={"Authorization": f"Bearer {team1_key}"}
|
||||
)
|
||||
assert team1_models.status_code == 200
|
||||
print("team1_models", json.dumps(team1_models.json(), indent=4))
|
||||
assert any(
|
||||
model["id"] == "gpt-4-team-test" for model in team1_models.json()["data"]
|
||||
), "Team1 should see their model"
|
||||
|
||||
# Check models visibility for team2's key
|
||||
team2_models = await client.get(
|
||||
"/models", headers={"Authorization": f"Bearer {team2_key}"}
|
||||
)
|
||||
assert team2_models.status_code == 200
|
||||
print("team2_models", json.dumps(team2_models.json(), indent=4))
|
||||
assert not any(
|
||||
model["id"] == "gpt-4-team-test" for model in team2_models.json()["data"]
|
||||
), "Team2 should not see team1's model"
|
||||
|
||||
# Cleanup
|
||||
model_id = model_response.json()["model_info"]["id"]
|
||||
await client.post("/model/delete", json={"id": model_id}, headers=headers)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_model_visibility_in_model_info_endpoint():
|
||||
"""
|
||||
Test that team-specific models are visible to all users in /v2/model/info endpoint:
|
||||
Note: /v2/model/info is used by the Admin UI to display model info
|
||||
1. Create a team
|
||||
2. Add a model associated with the team
|
||||
3. Generate a team key
|
||||
4. Verify both team key and non-team key can see the model in /v2/model/info
|
||||
"""
|
||||
client = AsyncClient(base_url=PROXY_BASE_URL)
|
||||
headers = {"Authorization": f"Bearer {TEST_MASTER_KEY}"}
|
||||
|
||||
# Create team
|
||||
team_response = await client.post(
|
||||
"/team/new",
|
||||
json={"models": []},
|
||||
headers=headers,
|
||||
)
|
||||
assert team_response.status_code == 200
|
||||
team_id = team_response.json()["team_id"]
|
||||
|
||||
# Add model associated with team
|
||||
model_response = await client.post(
|
||||
"/model/new",
|
||||
json={
|
||||
"model_name": "gpt-4-team-test",
|
||||
"litellm_params": {
|
||||
"model": "gpt-4",
|
||||
"custom_llm_provider": "openai",
|
||||
"api_key": "fake_key",
|
||||
},
|
||||
"model_info": {"team_id": team_id},
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
assert model_response.status_code == 200
|
||||
|
||||
# Generate team key
|
||||
team_key = (
|
||||
await client.post("/key/generate", json={"team_id": team_id}, headers=headers)
|
||||
).json()["key"]
|
||||
|
||||
# Generate non-team key
|
||||
non_team_key = (
|
||||
await client.post("/key/generate", json={}, headers=headers)
|
||||
).json()["key"]
|
||||
|
||||
# Check model info visibility with team key
|
||||
team_model_info = await client.get(
|
||||
"/v2/model/info",
|
||||
headers={"Authorization": f"Bearer {team_key}"},
|
||||
params={"model_name": "gpt-4-team-test"},
|
||||
)
|
||||
assert team_model_info.status_code == 200
|
||||
team_model_info = team_model_info.json()
|
||||
print("Team 1 model info", json.dumps(team_model_info, indent=4))
|
||||
assert any(
|
||||
model["model_info"].get("team_public_model_name") == "gpt-4-team-test"
|
||||
for model in team_model_info["data"]
|
||||
), "Team1 should see their model"
|
||||
|
||||
# Check model info visibility with non-team key
|
||||
non_team_model_info = await client.get(
|
||||
"/v2/model/info",
|
||||
headers={"Authorization": f"Bearer {non_team_key}"},
|
||||
params={"model_name": "gpt-4-team-test"},
|
||||
)
|
||||
assert non_team_model_info.status_code == 200
|
||||
non_team_model_info = non_team_model_info.json()
|
||||
print("Non-team model info", json.dumps(non_team_model_info, indent=4))
|
||||
assert any(
|
||||
model["model_info"].get("team_public_model_name") == "gpt-4-team-test"
|
||||
for model in non_team_model_info["data"]
|
||||
), "Non-team should see the model"
|
||||
|
||||
# Cleanup
|
||||
model_id = model_response.json()["model_info"]["id"]
|
||||
await client.post("/model/delete", json={"id": model_id}, headers=headers)
|
||||
|
|
@ -118,45 +118,6 @@ async def test_end_user_new():
|
|||
await asyncio.gather(*tasks)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aaaend_user_specific_region():
|
||||
"""
|
||||
- Specify region user can make calls in
|
||||
- Make a generic call
|
||||
- assert returned api base is for model in region
|
||||
|
||||
Repeat 3 times
|
||||
"""
|
||||
key: str = ""
|
||||
## CREATE USER ##
|
||||
async with aiohttp.ClientSession() as session:
|
||||
end_user_obj = await new_end_user(
|
||||
session=session,
|
||||
i=0,
|
||||
user_id=str(uuid.uuid4()),
|
||||
model_region="eu",
|
||||
)
|
||||
|
||||
## MAKE CALL ##
|
||||
key_gen = await generate_key(
|
||||
session=session, i=0, models=["gpt-5-mini-end-user-test"]
|
||||
)
|
||||
|
||||
key = key_gen["key"]
|
||||
|
||||
for _ in range(3):
|
||||
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000", max_retries=0)
|
||||
|
||||
print("SENDING USER PARAM - {}".format(end_user_obj["user_id"]))
|
||||
result = await client.chat.completions.with_raw_response.create(
|
||||
model="gpt-5-mini-end-user-test",
|
||||
messages=[{"role": "user", "content": "Hey!"}],
|
||||
user=end_user_obj["user_id"],
|
||||
)
|
||||
|
||||
assert result.headers.get("x-litellm-model-region") == "eu"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enduser_tpm_limits_non_master_key():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -147,55 +147,6 @@ async def test_key_gen_bad_key():
|
|||
pass
|
||||
|
||||
|
||||
async def update_key(session, get_key, metadata: Optional[dict] = None):
|
||||
"""
|
||||
Make sure only models user has access to are returned
|
||||
"""
|
||||
url = "http://0.0.0.0:4000/key/update"
|
||||
headers = {
|
||||
"Authorization": "Bearer sk-1234",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {"key": get_key}
|
||||
|
||||
if metadata is not None:
|
||||
data["metadata"] = metadata
|
||||
else:
|
||||
data.update({"models": ["gpt-4"], "duration": "120s"})
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def update_proxy_budget(session):
|
||||
"""
|
||||
Make sure only models user has access to are returned
|
||||
"""
|
||||
url = "http://0.0.0.0:4000/user/update"
|
||||
headers = {
|
||||
"Authorization": f"Bearer sk-1234",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {"user_id": "litellm-proxy-budget", "spend": 0}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def chat_completion(session, key, model="gpt-4"):
|
||||
url = "http://0.0.0.0:4000/chat/completions"
|
||||
headers = {
|
||||
|
|
@ -232,39 +183,6 @@ async def chat_completion(session, key, model="gpt-4"):
|
|||
pass
|
||||
|
||||
|
||||
async def image_generation(session, key, model="gpt-image-1"):
|
||||
url = "http://0.0.0.0:4000/v1/images/generations"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {
|
||||
"model": model,
|
||||
"prompt": "A cute baby sea otter",
|
||||
}
|
||||
|
||||
for i in range(3):
|
||||
try:
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
print("/images/generations response", response_text)
|
||||
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(
|
||||
f"Request did not return a 200 status code: {status}. Response: {response_text}"
|
||||
)
|
||||
|
||||
return await response.json()
|
||||
except Exception as e:
|
||||
if "Request did not return a 200 status code" in str(e):
|
||||
raise e
|
||||
else:
|
||||
pass
|
||||
|
||||
|
||||
async def chat_completion_streaming(session, key, model="gpt-4"):
|
||||
client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4000")
|
||||
messages = [
|
||||
|
|
@ -292,29 +210,6 @@ async def chat_completion_streaming(session, key, model="gpt-4"):
|
|||
return prompt_tokens, completion_tokens
|
||||
|
||||
|
||||
@pytest.mark.parametrize("metadata", [{"test": "new"}, {}])
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_update(metadata):
|
||||
"""
|
||||
Create key
|
||||
Update key with new model
|
||||
Test key w/ model
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key_gen = await generate_key(session=session, i=0, metadata={"test": "test"})
|
||||
key = key_gen["key"]
|
||||
assert key_gen["metadata"]["test"] == "test"
|
||||
updated_key = await update_key(
|
||||
session=session,
|
||||
get_key=key,
|
||||
metadata=metadata,
|
||||
)
|
||||
print(f"updated_key['metadata']: {updated_key['metadata']}")
|
||||
assert updated_key["metadata"] == metadata
|
||||
await update_proxy_budget(session=session) # resets proxy spend
|
||||
await chat_completion(session=session, key=key)
|
||||
|
||||
|
||||
async def delete_key(session, get_key, auth_key="sk-1234"):
|
||||
"""
|
||||
Delete key
|
||||
|
|
@ -583,61 +478,6 @@ async def test_aaaaakey_info_spend_values_streaming():
|
|||
), f"Expected={rounded_response_cost}, Got={rounded_key_info_spend}"
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=3, delay=1)
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_info_spend_values_image_generation():
|
||||
"""
|
||||
Test to ensure spend is correctly calculated
|
||||
- create key
|
||||
- make image gen call
|
||||
- assert cost is expected value
|
||||
"""
|
||||
|
||||
async def retry_request(func, *args, _max_attempts=5, **kwargs):
|
||||
for attempt in range(_max_attempts):
|
||||
try:
|
||||
return await func(*args, **kwargs)
|
||||
except aiohttp.client_exceptions.ClientOSError as e:
|
||||
if attempt + 1 == _max_attempts:
|
||||
raise # re-raise the last ClientOSError if all attempts failed
|
||||
print(f"Attempt {attempt+1} failed, retrying...")
|
||||
|
||||
async with aiohttp.ClientSession(
|
||||
timeout=aiohttp.ClientTimeout(total=600)
|
||||
) as session:
|
||||
## Test Spend Update ##
|
||||
# completion
|
||||
key_gen = await generate_key(session=session, i=0)
|
||||
key = key_gen["key"]
|
||||
response = await image_generation(session=session, key=key)
|
||||
await asyncio.sleep(5)
|
||||
key_info = await retry_request(
|
||||
get_key_info, session=session, get_key=key, call_key=key
|
||||
)
|
||||
spend = key_info["info"]["spend"]
|
||||
assert spend > 0
|
||||
|
||||
# The record/replay proxy serves this identical second call from its
|
||||
# cassette (free), but the proxy must still bill it. Spend logging is
|
||||
# async/batched, so poll for the increase rather than reading once after a
|
||||
# fixed sleep; a spend that never grows means the repeat was not billed
|
||||
# (e.g. the proxy response cache is on), which this still catches.
|
||||
await image_generation(session=session, key=key)
|
||||
spend_after = spend
|
||||
for _ in range(12):
|
||||
await asyncio.sleep(5)
|
||||
key_info = await retry_request(
|
||||
get_key_info, session=session, get_key=key, call_key=key
|
||||
)
|
||||
spend_after = key_info["info"]["spend"]
|
||||
if spend_after > spend:
|
||||
break
|
||||
assert spend_after > spend, (
|
||||
"spend did not increase on an identical repeat image call; the repeat "
|
||||
"was not billed (the proxy response cache may be on)"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="Frequent check on ci/cd leads to read timeout issue.")
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_with_budgets():
|
||||
|
|
@ -684,33 +524,6 @@ async def test_key_with_budgets():
|
|||
assert reset_at_init_value != reset_at_new_value
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_crossing_budget():
|
||||
"""
|
||||
- Create key with budget with budget=0.00000001
|
||||
- make a /chat/completions call
|
||||
- wait 5s
|
||||
- make a /chat/completions call - should fail with key crossed it's budget
|
||||
|
||||
- Check if value updated
|
||||
"""
|
||||
from litellm.proxy.utils import hash_token
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key_gen = await generate_key(session=session, i=0, budget=0.0000001)
|
||||
key = key_gen["key"]
|
||||
hashed_token = hash_token(token=key)
|
||||
print(f"hashed_token: {hashed_token}")
|
||||
|
||||
response = await chat_completion(session=session, key=key)
|
||||
print("response 1: ", response)
|
||||
await asyncio.sleep(10)
|
||||
with pytest.raises(Exception, match="Budget has been exceeded!") as exc_info:
|
||||
response = await chat_completion(session=session, key=key)
|
||||
e = exc_info.value
|
||||
assert "Budget has been exceeded!" in str(e)
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="AWS Suspended Account")
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_info_spend_values_sagemaker():
|
||||
|
|
@ -736,32 +549,6 @@ async def test_key_info_spend_values_sagemaker():
|
|||
# assert rounded_response_cost == rounded_key_info_spend
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_rate_limit():
|
||||
"""
|
||||
Tests backoff/retry logic on parallel request error.
|
||||
- Create key with max parallel requests 0
|
||||
- run 2 requests -> both fail
|
||||
- Create key with max parallel request 1
|
||||
- run 2 requests
|
||||
- both should succeed
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key_gen = await generate_key(session=session, i=0, max_parallel_requests=0)
|
||||
new_key = key_gen["key"]
|
||||
try:
|
||||
await chat_completion(session=session, key=new_key)
|
||||
pytest.fail(f"Expected this call to fail")
|
||||
except Exception as e:
|
||||
pass
|
||||
key_gen = await generate_key(session=session, i=0, max_parallel_requests=1)
|
||||
new_key = key_gen["key"]
|
||||
try:
|
||||
await chat_completion(session=session, key=new_key)
|
||||
except Exception as e:
|
||||
pytest.fail(f"Expected this call to work - {str(e)}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_delete_ui():
|
||||
"""
|
||||
|
|
@ -845,43 +632,3 @@ async def test_key_model_list(model_access, model_access_level, model_endpoint):
|
|||
assert len(model_list["data"]) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_user_not_in_db():
|
||||
"""
|
||||
- Create a key with unique user-id (not in db)
|
||||
- Check if key can make `/chat/completion` call
|
||||
"""
|
||||
my_unique_user = str(uuid.uuid4())
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key_gen = await generate_key(
|
||||
session=session,
|
||||
i=0,
|
||||
user_id=my_unique_user,
|
||||
)
|
||||
key = key_gen["key"]
|
||||
try:
|
||||
await chat_completion(session=session, key=key)
|
||||
except Exception as e:
|
||||
pytest.fail(f"Expected this call to work - {str(e)}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_over_budget():
|
||||
"""
|
||||
Test if key over budget is handled as expected.
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key_gen = await generate_key(session=session, i=0, budget=0.0000001)
|
||||
key = key_gen["key"]
|
||||
try:
|
||||
await chat_completion(session=session, key=key)
|
||||
except Exception as e:
|
||||
pytest.fail(f"Expected this call to work - {str(e)}")
|
||||
|
||||
## CALL `/models` - expect to work
|
||||
model_list = await get_key_info(session=session, get_key=key, call_key=key)
|
||||
## CALL `/chat/completions` - expect to fail
|
||||
with pytest.raises(Exception, match="Budget has been exceeded!") as exc_info:
|
||||
await chat_completion(session=session, key=key)
|
||||
e = exc_info.value
|
||||
assert "Budget has been exceeded!" in str(e)
|
||||
|
|
|
|||
|
|
@ -106,37 +106,6 @@ async def add_models(
|
|||
return response_json
|
||||
|
||||
|
||||
async def update_model(
|
||||
session, model_id="123", model_name="azure-gpt-3.5", key="sk-1234"
|
||||
):
|
||||
url = "http://0.0.0.0:4000/model/update"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
data = {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4.1-nano",
|
||||
"api_key": "os.environ/OPENAI_API_KEY",
|
||||
},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
print(f"Add models {response_text}")
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
|
||||
response_json = await response.json()
|
||||
return response_json
|
||||
|
||||
|
||||
async def get_model_info(session, key, litellm_model_id=None):
|
||||
"""
|
||||
Make sure only models user has access to are returned
|
||||
|
|
@ -301,169 +270,6 @@ async def test_add_and_delete_models():
|
|||
pass
|
||||
|
||||
|
||||
async def add_model_for_health_checking(session, model_id="123"):
|
||||
url = "http://0.0.0.0:4000/model/new"
|
||||
headers = {
|
||||
"Authorization": f"Bearer sk-1234",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
data = {
|
||||
"model_name": f"azure-model-health-check-{model_id}",
|
||||
"litellm_params": {
|
||||
"model": "gpt-4.1-nano",
|
||||
"api_key": os.getenv("OPENAI_API_KEY"),
|
||||
},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(f"Add models {response_text}")
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
|
||||
|
||||
async def get_model_info_v2(session, key):
|
||||
url = "http://0.0.0.0:4000/v2/model/info"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
async with session.get(url, headers=headers) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
print("response from v2/model/info")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
|
||||
|
||||
async def get_specific_model_info_v2(session, key, model_name):
|
||||
url = "http://0.0.0.0:4000/v2/model/info?debug=True&model=" + model_name
|
||||
print("running /model/info check for model=", model_name)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
async with session.get(url, headers=headers) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
print("response from v2/model/info")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
_json_response = await response.json()
|
||||
print("JSON response from /v2/model/info?model=", model_name, _json_response)
|
||||
|
||||
_model_info = _json_response["data"]
|
||||
assert len(_model_info) == 1, f"Expected 1 model, got {len(_model_info)}"
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
return _model_info[0]
|
||||
|
||||
|
||||
async def get_model_health(session, key, model_name):
|
||||
url = "http://0.0.0.0:4000/health?model=" + model_name
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
async with session.get(url, headers=headers) as response:
|
||||
status = response.status
|
||||
response_text = await response.json()
|
||||
print("response from /health?model=", model_name)
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
return response_text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_model_run_health():
|
||||
"""
|
||||
Add model
|
||||
Call /model/info and v2/model/info
|
||||
-> Admin UI calls v2/model/info
|
||||
Call /chat/completions
|
||||
Call /health
|
||||
-> Ensure the health check for the endpoint is working as expected
|
||||
"""
|
||||
from litellm._uuid import uuid
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key_gen = await generate_key(session=session)
|
||||
key = key_gen["key"]
|
||||
master_key = "sk-1234"
|
||||
model_id = str(uuid.uuid4())
|
||||
model_name = f"azure-model-health-check-{model_id}"
|
||||
print("adding model", model_name)
|
||||
await add_model_for_health_checking(session=session, model_id=model_id)
|
||||
_old_model_info = await get_specific_model_info_v2(
|
||||
session=session, key=key, model_name=model_name
|
||||
)
|
||||
print("model info before test", _old_model_info)
|
||||
|
||||
await asyncio.sleep(30)
|
||||
print("calling /model/info")
|
||||
await get_model_info(session=session, key=key)
|
||||
print("calling v2/model/info")
|
||||
await get_model_info_v2(session=session, key=key)
|
||||
|
||||
print("calling /chat/completions -> expect to work")
|
||||
await chat_completion(session=session, key=key, model=model_name)
|
||||
|
||||
print("calling /health?model=", model_name)
|
||||
_health_info = await get_model_health(
|
||||
session=session, key=master_key, model_name=model_name
|
||||
)
|
||||
_healthy_endpooint = _health_info["healthy_endpoints"][0]
|
||||
|
||||
assert _health_info["healthy_count"] == 1
|
||||
assert (
|
||||
_healthy_endpooint["model"] == "gpt-4.1-nano"
|
||||
) # this is the model that got added
|
||||
|
||||
# assert httpx client is is unchanges
|
||||
|
||||
await asyncio.sleep(10)
|
||||
|
||||
_model_info_after_test = await get_specific_model_info_v2(
|
||||
session=session, key=key, model_name=model_name
|
||||
)
|
||||
|
||||
print("model info after test", _model_info_after_test)
|
||||
old_openai_client = _old_model_info["openai_client"]
|
||||
new_openai_client = _model_info_after_test["openai_client"]
|
||||
print("old openai client", old_openai_client)
|
||||
print("new openai client", new_openai_client)
|
||||
|
||||
"""
|
||||
PROD TEST - This is extremly important
|
||||
The OpenAI client used should be the same after 30 seconds
|
||||
It is a serious bug if the openai client does not match here
|
||||
"""
|
||||
assert (
|
||||
old_openai_client == new_openai_client
|
||||
), "OpenAI client does not match for the same model after 30 seconds"
|
||||
|
||||
# cleanup
|
||||
await delete_model(session=session, model_id=model_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_personal_models_for_user():
|
||||
"""
|
||||
|
|
@ -506,52 +312,3 @@ async def test_model_group_info_e2e():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_model_e2e():
|
||||
"""
|
||||
Test team model e2e
|
||||
|
||||
- create team
|
||||
- create user
|
||||
- add user to team as admin
|
||||
- add model to team
|
||||
- update model
|
||||
- delete model
|
||||
"""
|
||||
from tests.test_users import new_user
|
||||
from tests.test_team import new_team
|
||||
from litellm._uuid import uuid
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Creat a user
|
||||
user_data = await new_user(session=session, i=0)
|
||||
user_id = user_data["user_id"]
|
||||
user_api_key = user_data["key"]
|
||||
|
||||
# Create a team
|
||||
member_list = [
|
||||
{"role": "admin", "user_id": user_id},
|
||||
]
|
||||
team_data = await new_team(session=session, member_list=member_list, i=0)
|
||||
team_id = team_data["team_id"]
|
||||
|
||||
model_id = str(uuid.uuid4())
|
||||
model_name = "my-test-model"
|
||||
# Add model to team
|
||||
model_data = await add_models(
|
||||
session=session,
|
||||
model_id=model_id,
|
||||
model_name=model_name,
|
||||
key=user_api_key,
|
||||
team_id=team_id,
|
||||
)
|
||||
model_id = model_data["model_id"]
|
||||
|
||||
# Update model
|
||||
model_data = await update_model(
|
||||
session=session, model_id=model_id, model_name=model_name, key=user_api_key
|
||||
)
|
||||
model_id = model_data["model_id"]
|
||||
|
||||
# Delete model
|
||||
await delete_model(session=session, model_id=model_id, key=user_api_key)
|
||||
|
|
|
|||
|
|
@ -523,23 +523,6 @@ async def test_image_generation():
|
|||
await image_generation(session=session, key=key_2)
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=5, delay=1)
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_wildcard_chat_completion():
|
||||
"""
|
||||
- Create key for model = "*" -> this has access to all models
|
||||
- proxy_server_config.yaml has model = *
|
||||
- Make chat completion call
|
||||
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
key_gen = await generate_key(session=session, models=["*"])
|
||||
key = key_gen["key"]
|
||||
|
||||
# call chat/completions with a model that the key was not created for + the model is not on the config.yaml
|
||||
await chat_completion(session=session, key=key, model="gpt-3.5-turbo-0125")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_all_models():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,319 +0,0 @@
|
|||
# What this tests ?
|
||||
## Tests /organization endpoints.
|
||||
import pytest
|
||||
import asyncio
|
||||
import aiohttp
|
||||
import time, uuid
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
|
||||
async def new_user(
|
||||
session,
|
||||
i,
|
||||
user_id=None,
|
||||
budget=None,
|
||||
budget_duration=None,
|
||||
models=["azure-models"],
|
||||
team_id=None,
|
||||
user_email=None,
|
||||
):
|
||||
url = "http://0.0.0.0:4000/user/new"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {
|
||||
"models": models,
|
||||
"aliases": {"mistral-7b": "gpt-3.5-turbo"},
|
||||
"duration": None,
|
||||
"max_budget": budget,
|
||||
"budget_duration": budget_duration,
|
||||
"user_email": user_email,
|
||||
}
|
||||
|
||||
if user_id is not None:
|
||||
data["user_id"] = user_id
|
||||
|
||||
if team_id is not None:
|
||||
data["team_id"] = team_id
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(f"Response {i} (Status code: {status}):")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(
|
||||
f"Request {i} did not return a 200 status code: {status}, response: {response_text}"
|
||||
)
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def new_organization(session, i, organization_alias, max_budget=None):
|
||||
url = "http://0.0.0.0:4000/organization/new"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {
|
||||
"organization_alias": organization_alias,
|
||||
"models": ["azure-models"],
|
||||
"max_budget": max_budget,
|
||||
}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(f"Response {i} (Status code: {status}):")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request {i} did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def add_member_to_org(
|
||||
session, i, organization_id, user_id, user_role="internal_user"
|
||||
):
|
||||
url = "http://0.0.0.0:4000/organization/member_add"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {
|
||||
"organization_id": organization_id,
|
||||
"member": {
|
||||
"user_id": user_id,
|
||||
"role": user_role,
|
||||
},
|
||||
}
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(f"Response {i} (Status code: {status}):")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request {i} did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def update_member_role(
|
||||
session, i, organization_id, user_id, user_role="internal_user"
|
||||
):
|
||||
url = "http://0.0.0.0:4000/organization/member_update"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {
|
||||
"organization_id": organization_id,
|
||||
"user_id": user_id,
|
||||
"role": user_role,
|
||||
}
|
||||
|
||||
async with session.patch(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(f"Response {i} (Status code: {status}):")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request {i} did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def delete_member_from_org(session, i, organization_id, user_id):
|
||||
url = "http://0.0.0.0:4000/organization/member_delete"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {
|
||||
"organization_id": organization_id,
|
||||
"user_id": user_id,
|
||||
}
|
||||
|
||||
async with session.delete(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(f"Response {i} (Status code: {status}):")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request {i} did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def delete_organization(session, i, organization_id):
|
||||
url = "http://0.0.0.0:4000/organization/delete"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
data = {"organization_ids": [organization_id]}
|
||||
|
||||
async with session.delete(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(f"Response {i} (Status code: {status}):")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request {i} did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
async def list_organization(session, i):
|
||||
url = "http://0.0.0.0:4000/organization/list"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
|
||||
async with session.get(url, headers=headers) as response:
|
||||
status = response.status
|
||||
response_json = await response.json()
|
||||
|
||||
print(f"Response {i} (Status code: {status}):")
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request {i} did not return a 200 status code: {status}")
|
||||
|
||||
# Assert that budget info is returned for each organization
|
||||
for org in response_json:
|
||||
assert (
|
||||
"litellm_budget_table" in org
|
||||
), "Missing budget info in organization response"
|
||||
# Optionally also check that it's not null
|
||||
assert org["litellm_budget_table"] is not None, "Budget info is None"
|
||||
|
||||
return response_json
|
||||
|
||||
|
||||
@pytest.mark.flaky(retries=5, delay=1)
|
||||
@pytest.mark.asyncio
|
||||
async def test_organization_new():
|
||||
"""
|
||||
Make 20 parallel calls to /organization/new. Assert all worked.
|
||||
"""
|
||||
organization_alias = f"Organization: {uuid.uuid4()}"
|
||||
async with aiohttp.ClientSession() as session:
|
||||
tasks = [
|
||||
new_organization(
|
||||
session=session, i=0, organization_alias=organization_alias
|
||||
)
|
||||
for i in range(1, 20)
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_organization_list():
|
||||
"""
|
||||
create 2 new Organizations
|
||||
check if the Organization list is not empty
|
||||
"""
|
||||
organization_alias = f"Organization: {uuid.uuid4()}"
|
||||
async with aiohttp.ClientSession() as session:
|
||||
tasks = [
|
||||
new_organization(
|
||||
session=session, i=0, organization_alias=organization_alias
|
||||
)
|
||||
for i in range(1, 2)
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
response_json = await list_organization(session, i=0)
|
||||
print(len(response_json))
|
||||
|
||||
if len(response_json) == 0:
|
||||
raise Exception("Return empty list of organization")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_organization_delete():
|
||||
"""
|
||||
create a new organization
|
||||
delete the organization
|
||||
check if the Organization list is set
|
||||
"""
|
||||
organization_alias = f"Organization: {uuid.uuid4()}"
|
||||
async with aiohttp.ClientSession() as session:
|
||||
tasks = [
|
||||
new_organization(
|
||||
session=session, i=0, organization_alias=organization_alias
|
||||
)
|
||||
]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
response_json = await list_organization(session, i=0)
|
||||
print(len(response_json))
|
||||
|
||||
organization_id = response_json[0]["organization_id"]
|
||||
await delete_organization(session, i=0, organization_id=organization_id)
|
||||
|
||||
response_json = await list_organization(session, i=0)
|
||||
print(len(response_json))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_organization_member_flow():
|
||||
"""
|
||||
create a new organization
|
||||
add a new member to the organization
|
||||
check if the member is added to the organization
|
||||
update the member's role in the organization
|
||||
delete the member from the organization
|
||||
check if the member is deleted from the organization
|
||||
"""
|
||||
organization_alias = f"Organization: {uuid.uuid4()}"
|
||||
async with aiohttp.ClientSession() as session:
|
||||
response_json = await new_organization(
|
||||
session=session, i=0, organization_alias=organization_alias
|
||||
)
|
||||
organization_id = response_json["organization_id"]
|
||||
|
||||
response_json = await list_organization(session, i=0)
|
||||
print(len(response_json))
|
||||
|
||||
new_user_response_json = await new_user(
|
||||
session=session, i=0, user_email=f"test_user_{uuid.uuid4()}@example.com"
|
||||
)
|
||||
user_id = new_user_response_json["user_id"]
|
||||
|
||||
await add_member_to_org(
|
||||
session, i=0, organization_id=organization_id, user_id=user_id
|
||||
)
|
||||
|
||||
response_json = await list_organization(session, i=0)
|
||||
print(len(response_json))
|
||||
|
||||
for orgs in response_json:
|
||||
tmp_organization_id = orgs["organization_id"]
|
||||
if (
|
||||
tmp_organization_id is not None
|
||||
and tmp_organization_id == organization_id
|
||||
):
|
||||
user_id = orgs["members"][0]["user_id"]
|
||||
|
||||
response_json = await list_organization(session, i=0)
|
||||
print(len(response_json))
|
||||
|
||||
await update_member_role(
|
||||
session,
|
||||
i=0,
|
||||
organization_id=organization_id,
|
||||
user_id=user_id,
|
||||
user_role="org_admin",
|
||||
)
|
||||
|
||||
response_json = await list_organization(session, i=0)
|
||||
print(len(response_json))
|
||||
|
||||
await delete_member_from_org(
|
||||
session, i=0, organization_id=organization_id, user_id=user_id
|
||||
)
|
||||
|
||||
response_json = await list_organization(session, i=0)
|
||||
print(len(response_json))
|
||||
|
|
@ -221,23 +221,6 @@ async def get_predict_spend_logs(session):
|
|||
return await response.json()
|
||||
|
||||
|
||||
async def get_spend_report(session, start_date, end_date):
|
||||
url = "http://0.0.0.0:4000/global/spend/report"
|
||||
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
|
||||
async with session.get(
|
||||
url, headers=headers, params={"start_date": start_date, "end_date": end_date}
|
||||
) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request did not return a 200 status code: {status}")
|
||||
return await response.json()
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="datetime in ci/cd gets set weirdly")
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_predicted_spend_logs():
|
||||
|
|
@ -308,37 +291,3 @@ async def test_spend_logs_high_traffic():
|
|||
raise Exception("it worked!")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_report_endpoint():
|
||||
async with aiohttp.ClientSession(
|
||||
timeout=aiohttp.ClientTimeout(total=600)
|
||||
) as session:
|
||||
import datetime
|
||||
|
||||
todays_date = datetime.date.today() + datetime.timedelta(days=1)
|
||||
todays_date = todays_date.strftime("%Y-%m-%d")
|
||||
|
||||
print("todays_date", todays_date)
|
||||
thirty_days_ago = (
|
||||
datetime.date.today() - datetime.timedelta(days=30)
|
||||
).strftime("%Y-%m-%d")
|
||||
spend_report = await get_spend_report(
|
||||
session=session, start_date=thirty_days_ago, end_date=todays_date
|
||||
)
|
||||
print("spend report", spend_report)
|
||||
|
||||
for row in spend_report:
|
||||
date = row["group_by_day"]
|
||||
teams = row["teams"]
|
||||
for team in teams:
|
||||
team_name = team["team_name"]
|
||||
total_spend = team["total_spend"]
|
||||
metadata = team["metadata"]
|
||||
|
||||
assert team_name is not None
|
||||
|
||||
print(f"Date: {date}")
|
||||
print(f"Team: {team_name}")
|
||||
print(f"Total Spend: {total_spend}")
|
||||
print("Metadata: ", metadata)
|
||||
print()
|
||||
|
|
|
|||
|
|
@ -690,40 +690,6 @@ async def test_member_delete(dimension):
|
|||
assert user_in_team is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_alias():
|
||||
"""
|
||||
- Create team w/ model alias
|
||||
- Create key for team
|
||||
- Check if key works
|
||||
"""
|
||||
async with aiohttp.ClientSession() as session:
|
||||
## Create admin
|
||||
admin_user = f"{uuid.uuid4()}"
|
||||
await new_user(session=session, i=0, user_id=admin_user)
|
||||
## Create normal user
|
||||
normal_user = f"{uuid.uuid4()}"
|
||||
await new_user(session=session, i=0, user_id=normal_user)
|
||||
## Create team with 1 admin and 1 user
|
||||
member_list = [
|
||||
{"role": "admin", "user_id": admin_user},
|
||||
{"role": "user", "user_id": normal_user},
|
||||
]
|
||||
team_data = await new_team(
|
||||
session=session,
|
||||
i=0,
|
||||
member_list=member_list,
|
||||
model_aliases={"cheap-model": "gpt-3.5-turbo"},
|
||||
)
|
||||
## Create key
|
||||
key_gen = await generate_key(
|
||||
session=session, i=0, team_id=team_data["team_id"], models=["gpt-3.5-turbo"]
|
||||
)
|
||||
key = key_gen["key"]
|
||||
## Test key
|
||||
response = await chat_completion(session=session, key=key, model="cheap-model")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_users_in_team_budget():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -40,51 +40,6 @@ async def new_user(
|
|||
return await response.json()
|
||||
|
||||
|
||||
async def generate_key(
|
||||
session,
|
||||
i,
|
||||
budget=None,
|
||||
budget_duration=None,
|
||||
models=["azure-models", "gpt-4", "dall-e-3"],
|
||||
max_parallel_requests: Optional[int] = None,
|
||||
user_id: Optional[str] = None,
|
||||
team_id: Optional[str] = None,
|
||||
metadata: Optional[dict] = None,
|
||||
calling_key="sk-1234",
|
||||
):
|
||||
url = "http://0.0.0.0:4000/key/generate"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {calling_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
data = {
|
||||
"models": models,
|
||||
"aliases": {"mistral-7b": "gpt-3.5-turbo"},
|
||||
"duration": None,
|
||||
"max_budget": budget,
|
||||
"budget_duration": budget_duration,
|
||||
"max_parallel_requests": max_parallel_requests,
|
||||
"user_id": user_id,
|
||||
"team_id": team_id,
|
||||
"metadata": metadata,
|
||||
}
|
||||
|
||||
print(f"data: {data}")
|
||||
|
||||
async with session.post(url, headers=headers, json=data) as response:
|
||||
status = response.status
|
||||
response_text = await response.text()
|
||||
|
||||
print(f"Response {i} (Status code: {status}):")
|
||||
print(response_text)
|
||||
print()
|
||||
|
||||
if status != 200:
|
||||
raise Exception(f"Request {i} did not return a 200 status code: {status}")
|
||||
|
||||
return await response.json()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_new():
|
||||
"""
|
||||
|
|
@ -260,62 +215,6 @@ async def test_global_proxy_budget_update():
|
|||
assert new_new_spend > new_spend
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_model_access():
|
||||
"""
|
||||
- Create user with model access
|
||||
- Create key with user
|
||||
- Call model that user has access to -> should work
|
||||
- Call wildcard model that user has access to -> should work
|
||||
- Call model that user does not have access to -> should fail
|
||||
- Call wildcard model that user does not have access to -> should fail
|
||||
"""
|
||||
import openai
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
get_user = f"krrish_{time.time()}@berri.ai"
|
||||
await new_user(
|
||||
session=session,
|
||||
i=0,
|
||||
user_id=get_user,
|
||||
models=["good-model", "anthropic/*"],
|
||||
)
|
||||
|
||||
result = await generate_key(
|
||||
session=session,
|
||||
i=0,
|
||||
user_id=get_user,
|
||||
models=[], # assign no models. Allow inheritance from user
|
||||
)
|
||||
key = result["key"]
|
||||
|
||||
await chat_completion(
|
||||
session=session,
|
||||
key=key,
|
||||
model="anthropic/claude-haiku-4-5-20251001",
|
||||
)
|
||||
|
||||
await chat_completion(
|
||||
session=session,
|
||||
key=key,
|
||||
model="good-model",
|
||||
)
|
||||
|
||||
with pytest.raises(openai.PermissionDeniedError):
|
||||
await chat_completion(
|
||||
session=session,
|
||||
key=key,
|
||||
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
)
|
||||
|
||||
with pytest.raises(openai.PermissionDeniedError):
|
||||
await chat_completion(
|
||||
session=session,
|
||||
key=key,
|
||||
model="groq/claude-3-5-haiku-20241022",
|
||||
)
|
||||
|
||||
|
||||
import json
|
||||
from litellm._uuid import uuid
|
||||
import pytest
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue