test: replace 61 live logging and otel tests with offline unit and integration coverage (#45362)

* test: replace 61 live logging and otel tests with offline unit and integration coverage

* test: restore datadog formatting, deliver datadog logs and redis failures over the wire, assert full router hook payloads, tighten stream usage and otel checks

Restores the pre-existing test_datadog.py lines the branch had reflowed. Datadog success, failure and redis-failure replacements now assert the gzip body posted to the intake, with redis failing through a real cache on a closed local port. Router hook sequence recorder checks the legacy field types and asserts concrete payloads, exact streaming and fallback sequences. Stream usage asserts the default include_usage request body and the redacted messages value. Otel asserts response id and token counts.

* test: assert budget envelope figures, guardrail inspection and exact prometheus samples per request

* test: isolate the prometheus latency test on its own deployment so counts do not depend on order

* test: cover timed slack delivery, keep the redis failure test off the network, and make new payload types read-only

* test: drain the redis test's logging and scope otel span checks to the test's own trace

* test: drive the periodic slack flush without a wall-clock interval

* test: scope router hook events to the test, script the db clock, split a nested comprehension

---------

Co-authored-by: yuneng <yuneng@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 15:55:57 -07:00 • committed by GitHub
parent 8bae65e0b6
commit f09408dd37
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
32 changed files with 3103 additions and 4842 deletions

View file

@ -23,7 +23,14 @@ import pytest
import yaml
from pydantic import JsonValue
from tests.integration._support.client import JSON_OBJECT, Gateway, eventually, gateway_from_environment, string_value
from tests.integration._support.client import (
JSON_OBJECT,
Gateway,
eventually,
gateway_from_environment,
object_value,
string_value,
)
from tests.integration._support.database import read_rows
from tests.integration._support.process import OwnedProxy, group_members, owned_proxy_process
from tests.integration._support.provider import SharedProvider
@ -648,3 +655,52 @@ def test_flag_values_that_do_not_disable_keep_the_gates_open(idp: Idp, tmp_path:
for method, path in (("GET", "/sso/saml/login"), ("POST", "/sso/saml/callback")):
saml: Final = proxy.client.request(method, path)
assert DISABLED_PAGE_TITLE not in saml.text and saml.status_code != 200, f"{path}: {saml.status_code}"
def test_cli_session_token_is_denied_once_its_team_budget_is_exhausted(
one_worker: OneWorkerProxy, provider: SharedProvider
) -> None:
proxy: Final = one_worker.owned.gateway
subject: Final = f"cli-sso-team-budget-{uuid.uuid4().hex[:12]}"
budget: Final = 0.0000000005
with proxy.scenario() as scenario:
team: Final = scenario.team(max_budget=budget, models=[MESSAGE_MODEL])
scenario.user(user_id=subject, user_email=f"{subject}@example.com", user_role="internal_user")
added: Final = proxy.request(
"POST", "/team/member_add", {"team_id": team, "member": {"user_id": subject, "role": "user"}}
)
assert added.status_code == 200, added.text
session: Final = _start_lite_login(proxy)
with _browser() as browser:
_sign_in(proxy, one_worker.idp, browser, session, subject=subject)
ready: Final = proxy.client.get(
f"/sso/cli/poll/{session.login_id}",
params={"team_id": team},
headers={POLL_SECRET_HEADER: session.poll_secret},
)
assert ready.status_code == 200, f"{ready.status_code} {ready.text}"
body: Final = JSON_OBJECT.validate_json(ready.content)
assert body["status"] == "ready" and body["user_id"] == subject, ready.text
key: Final = string_value(body["key"])
assert not key.startswith("sk-"), key
_send_message(proxy, provider, key)
eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team,)),
lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) > budget,
seconds=70,
)
refused: Final = proxy.request(
"POST",
"/v1/messages",
{"model": MESSAGE_MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "over team budget"}]},
key=key,
)
assert refused.status_code == 422, f"{refused.status_code} {refused.text}"
error: Final = object_value(JSON_OBJECT.validate_json(refused.content)["error"])
assert error["type"] == "budget_exceeded", refused.text
assert error["code"] == "422", refused.text
message: Final = string_value(error["message"])
assert "Budget has been exceeded!" in message, refused.text
assert f"Team={team}" in message, refused.text
assert "Current cost:" in message and f"Max budget: {budget}" in message, refused.text
assert provider.received() == ()

View file

@ -0,0 +1,171 @@
from __future__ import annotations
import json
import uuid
from collections.abc import Sequence
from dataclasses import dataclass
from typing import Final, Literal
import pytest
from tests.integration._support.client import JSON_OBJECT, Gateway, Scenario, object_value, string_value
from tests.integration._support.provider import PROVIDER_URL, SharedProvider
from tests.integration._support.wire import Reply
_Denial = Literal["key_model_access_denied", "team_model_access_denied"]
@dataclass(frozen=True, slots=True)
class _Models:
gpt: str
gpt_mini: str
claude: str
bedrock_claude: str
bedrock_titan: str
@dataclass(frozen=True, slots=True)
class _AccessCase:
name: str
allowed: Sequence[str] | None
requested: str
served: bool
_CASES: Final = (
_AccessCase("openai_wildcard_denies_anthropic", ["openai/*"], "claude", False),
_AccessCase("exact_name_allows_itself", ["gpt"], "gpt", True),
_AccessCase("provider_wildcard_allows_bedrock", ["bedrock/*"], "bedrock_claude", True),
_AccessCase("family_wildcard_allows_its_family", ["bedrock/anthropic.*"], "bedrock_claude", True),
_AccessCase("family_wildcard_denies_another_family", ["bedrock/anthropic.*"], "bedrock_titan", False),
_AccessCase("unset_models_allow_everything", None, "gpt", True),
_AccessCase("empty_models_allow_everything", [], "gpt", True),
)
def _deployment(scenario: Scenario, gateway: Gateway, name: str) -> str:
created: Final = gateway.post(
"/model/new",
{
"model_name": name,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": f"{PROVIDER_URL}/v1",
"api_key": "sk-fixture",
},
"model_info": {"id": f"access-{uuid.uuid4().hex}"},
},
)
scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"]))
return name
def _models(scenario: Scenario, gateway: Gateway) -> _Models:
tag: Final = uuid.uuid4().hex[:10]
return _Models(
gpt=_deployment(scenario, gateway, f"openai/gpt-{tag}"),
gpt_mini=_deployment(scenario, gateway, f"openai/gpt-mini-{tag}"),
claude=_deployment(scenario, gateway, f"anthropic/claude-{tag}"),
bedrock_claude=_deployment(scenario, gateway, f"bedrock/anthropic.claude-{tag}"),
bedrock_titan=_deployment(scenario, gateway, f"bedrock/amazon.titan-{tag}"),
)
def _pick(models: _Models, alias: str) -> str:
return string_value(getattr(models, alias))
def _completion() -> Reply:
return Reply(
body=json.dumps(
{
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "scripted"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5},
}
).encode()
)
def _served(gateway: Gateway, provider: SharedProvider, key: str, model: str) -> None:
provider.expect(_completion())
body: Final = gateway.chat(model, key=key, text=f"access {uuid.uuid4().hex}")
assert object_value(object_value(body["choices"][0])["message"])["content"] == "scripted", body
assert len(provider.received()) == 1
def _denied(gateway: Gateway, provider: SharedProvider, key: str, model: str, denial: _Denial) -> None:
refused: Final = gateway.request(
"POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "hi"}]}, key=key
)
assert refused.status_code == 403, f"{refused.status_code} {refused.text}"
error: Final = object_value(JSON_OBJECT.validate_json(refused.content)["error"])
assert error["type"] == denial, refused.text
assert error["param"] == "model", refused.text
assert error["code"] == "403", refused.text
message: Final = string_value(error["message"])
assert "is not available for this API key" in message, message
assert "not allowed to access model" not in message, message
assert provider.received() == ()
@pytest.mark.parametrize("case", _CASES, ids=[case.name for case in _CASES])
def test_a_key_model_allow_list_decides_which_models_it_reaches(
gateway: Gateway, provider: SharedProvider, case: _AccessCase
) -> None:
with gateway.scenario() as scenario:
models: Final = _models(scenario, gateway)
allowed: Final = (
None
if case.allowed is None
else [_pick(models, item) if hasattr(models, item) else item for item in case.allowed]
)
key: Final = scenario.key(models=allowed)
requested: Final = _pick(models, case.requested)
if case.served:
_served(gateway, provider, key, requested)
else:
_denied(gateway, provider, key, requested, "key_model_access_denied")
def test_widening_a_key_allow_list_to_a_wildcard_takes_effect_on_the_next_request(
gateway: Gateway, provider: SharedProvider
) -> None:
with gateway.scenario() as scenario:
models: Final = _models(scenario, gateway)
key: Final = scenario.key(models=[models.gpt])
_served(gateway, provider, key, models.gpt)
_denied(gateway, provider, key, models.gpt_mini, "key_model_access_denied")
gateway.post("/key/update", {"key": key, "models": ["openai/*"]})
_served(gateway, provider, key, models.gpt)
_served(gateway, provider, key, models.gpt_mini)
_denied(gateway, provider, key, models.claude, "key_model_access_denied")
def test_a_team_allow_list_denies_its_keys_a_model_outside_it(gateway: Gateway, provider: SharedProvider) -> None:
with gateway.scenario() as scenario:
models: Final = _models(scenario, gateway)
team: Final = scenario.team(models=["openai/*"])
key: Final = scenario.key(team_id=team)
_served(gateway, provider, key, models.gpt)
_denied(gateway, provider, key, models.claude, "team_model_access_denied")
def test_widening_a_team_allow_list_takes_effect_for_its_keys_on_the_next_request(
gateway: Gateway, provider: SharedProvider
) -> None:
with gateway.scenario() as scenario:
models: Final = _models(scenario, gateway)
team: Final = scenario.team(models=[models.gpt])
key: Final = scenario.key(team_id=team)
_served(gateway, provider, key, models.gpt)
_denied(gateway, provider, key, models.gpt_mini, "team_model_access_denied")
gateway.post("/team/update", {"team_id": team, "models": ["openai/*"]})
_served(gateway, provider, key, models.gpt)
_served(gateway, provider, key, models.gpt_mini)
_denied(gateway, provider, key, models.claude, "team_model_access_denied")

View file

@ -0,0 +1,162 @@
from __future__ import annotations
import json
import uuid
from collections.abc import Iterator
from pathlib import Path
from typing import Final
import pytest
import yaml
from pydantic import JsonValue, TypeAdapter
from tests.integration._support.client import JSON_OBJECT, Gateway, gateway_from_environment, object_value, string_value
from tests.integration._support.process import owned_proxy
from tests.integration._support.provider import PROVIDER_URL, SharedProvider
from tests.integration._support.wire import Reply
_TAGGED_GROUP: Final = "tag-filtered-group"
_TAGGED_IDS: Final = frozenset({"tag-filtered-team-a", "tag-filtered-team-b"})
_VISION_MODEL: Final = "llava-hf"
_ITEMS: Final = TypeAdapter(list[dict[str, JsonValue]])
def _config(directory: Path) -> Path:
configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
configuration["model_list"] = [
*configuration["model_list"],
*(
{
"model_name": _TAGGED_GROUP,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "sk-fixture",
"api_base": f"{PROVIDER_URL}/v1",
"tags": [tag],
},
"model_info": {"id": identity},
}
for tag, identity in (("teamA", "tag-filtered-team-a"), ("teamB", "tag-filtered-team-b"))
),
{
"model_name": _VISION_MODEL,
"litellm_params": {
"model": "openai/llava-hf",
"api_key": "sk-fixture",
"api_base": "http://127.0.0.1:9/v1",
},
"model_info": {"supports_vision": True},
},
]
configuration.setdefault("router_settings", {})["enable_tag_filtering"] = True
path: Final = directory / "config-declared-models.yaml"
path.write_text(yaml.safe_dump(configuration))
return path
@pytest.fixture(scope="module")
def declared(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]:
directory: Final = tmp_path_factory.mktemp("config-declared-models")
with (
gateway_from_environment() as shared,
owned_proxy(
shared,
directory,
{"OPENAI_API_KEY": "sk-fixture", "OPENAI_BASE_URL": f"{PROVIDER_URL}/v1", "GCS_FLUSH_INTERVAL": "1"},
config=_config(directory),
remove_environment=("GCS_BUCKET_NAME", "OPENAI_API_BASE"),
) as owned,
):
yield owned
def test_model_info_reports_the_vision_capability_declared_in_the_config(declared: Gateway) -> None:
listing: Final = _ITEMS.validate_python(declared.get("/model/info")["data"])
vision: Final = [item for item in listing if item["model_name"] == _VISION_MODEL]
assert len(vision) == 1, [item["model_name"] for item in listing]
assert object_value(vision[0]["model_info"])["supports_vision"] is True, vision[0]
def test_an_untagged_request_is_served_by_a_group_whose_deployments_are_all_tagged(
declared: Gateway, provider: SharedProvider
) -> None:
provider.expect(
Reply(
body=json.dumps(
{
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "tagged"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4},
}
).encode()
)
)
response: Final = declared.request(
"POST",
"/v1/chat/completions",
{"model": _TAGGED_GROUP, "messages": [{"role": "user", "content": f"untagged {uuid.uuid4().hex}"}]},
)
assert response.status_code == 200, response.text
assert response.headers["x-litellm-model-id"] in _TAGGED_IDS, dict(response.headers)
message: Final = object_value(object_value(JSON_OBJECT.validate_json(response.content)["choices"][0])["message"])
assert message["content"] == "tagged"
assert [request.target for request in provider.received()] == ["/v1/chat/completions"]
def test_a_moderation_request_without_a_model_reaches_the_provider_default(
declared: Gateway, provider: SharedProvider
) -> None:
provider.expect(
Reply(
body=json.dumps(
{
"id": f"modr-{uuid.uuid4().hex}",
"model": "omni-moderation-latest",
"results": [
{"flagged": True, "categories": {"violence": True}, "category_scores": {"violence": 0.9}}
],
}
).encode()
)
)
phrase: Final = f"I want to harm someone {uuid.uuid4().hex}"
response: Final = declared.request("POST", "/moderations", {"input": phrase})
assert response.status_code == 200, response.text
body: Final = JSON_OBJECT.validate_json(response.content)
assert body["model"] == "omni-moderation-latest", body
assert object_value(_ITEMS.validate_python(body["results"])[0])["flagged"] is True, body
sent: Final = provider.received()
assert [request.target for request in sent] == ["/v1/moderations"]
payload: Final = JSON_OBJECT.validate_json(sent[0].body)
assert payload == {"input": phrase}, payload
def test_key_health_reports_an_unconfigured_key_logging_callback_as_unhealthy(declared: Gateway) -> None:
with declared.scenario() as scenario:
key: Final = scenario.key(
metadata={
"logging": [
{
"callback_name": "gcs_bucket",
"callback_type": "success_and_failure",
"callback_vars": {
"gcs_bucket_name": "key-logging-project1",
"gcs_path_service_account": "bad-service-account",
},
}
]
}
)
health: Final = declared.request("POST", "/key/health", {}, key=key)
assert health.status_code == 200, health.text
body: Final = JSON_OBJECT.validate_json(health.content)
assert "key" in body, body
status: Final = object_value(body["logging_callbacks"])
assert status["callbacks"] == ["gcs_bucket"], status
assert status["status"] == "unhealthy", status
assert "GCS_BUCKET_NAME is not set in the environment" in string_value(status["details"]), status

View file

@ -0,0 +1,153 @@
from __future__ import annotations
import json
import shutil
import uuid
from collections.abc import Iterator
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from tests.integration._support.client import JSON_OBJECT, Gateway, gateway_from_environment, object_value, string_value
from tests.integration._support.process import owned_proxy
from tests.integration._support.wire import Reply, Request, Wire, wire_server
_ATTACHABLE: Final = "attachable-during-guard"
_WORDS: Final = "custom-words-during-guard"
_HEADER: Final = "x-litellm-applied-guardrails"
@dataclass(frozen=True, slots=True)
class _Rig:
gateway: Gateway
policy: Wire
upstream: Wire
model: str
def _completion(request: Request) -> Reply:
assert request.target == "/v1/chat/completions", request.target
return Reply(
body=json.dumps(
{
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "guarded"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4},
}
).encode()
)
def _allow(request: Request) -> Reply:
assert request.target == "/beta/litellm_basic_guardrail_api", request.target
return Reply(body=json.dumps({"action": "NONE"}).encode())
def _config(directory: Path, policy: Wire) -> Path:
shutil.copy(Path("litellm/proxy/example_config_yaml/custom_guardrail.py"), directory / "custom_guardrail.py")
configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
configuration["guardrails"] = [
{
"guardrail_name": _ATTACHABLE,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": "during_call",
"api_base": policy.url,
"api_key": "synthetic-guardrail-key",
},
},
{
"guardrail_name": _WORDS,
"litellm_params": {"guardrail": "custom_guardrail.myCustomGuardrail", "mode": "during_call"},
},
]
path: Final = directory / "guardrail-attachment.yaml"
path.write_text(yaml.safe_dump(configuration))
return path
@pytest.fixture(scope="module")
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]:
directory: Final = tmp_path_factory.mktemp("guardrail-attachment")
with (
wire_server(_allow) as policy,
wire_server(_completion) as upstream,
gateway_from_environment() as shared,
owned_proxy(shared, directory, {}, config=_config(directory, policy)) as owned,
owned.scenario() as scenario,
):
yield _Rig(owned, policy, upstream, scenario.model(api_base=f"{upstream.url}/v1", api_key="sk-fixture"))
def _prompt(text: str) -> str:
return f"{text} {uuid.uuid4().hex}"
def _ask(rig: _Rig, key: str | None, prompt: str, guardrails: list[str] | None = None) -> httpx.Response:
body: Final = {"model": rig.model, "messages": [{"role": "user", "content": prompt}]}
return rig.gateway.request(
"POST", "/v1/chat/completions", body if guardrails is None else {**body, "guardrails": guardrails}, key=key
)
def _served_without_guardrail(rig: _Rig, response: httpx.Response, prompt: str) -> None:
assert response.status_code == 200, response.text
assert _HEADER not in response.headers, dict(response.headers)
assert rig.policy.drain() == ()
forwarded: Final = rig.upstream.drain()
assert len(forwarded) == 1 and prompt in forwarded[0].body.decode(), forwarded
def _served_with_attachable(rig: _Rig, response: httpx.Response, prompt: str) -> None:
assert response.status_code == 200, response.text
assert response.headers[_HEADER] == _ATTACHABLE, dict(response.headers)
inspected: Final = rig.policy.drain()
assert len(inspected) == 1 and prompt in inspected[0].body.decode(), inspected
forwarded: Final = rig.upstream.drain()
assert len(forwarded) == 1 and prompt in forwarded[0].body.decode(), forwarded
def test_a_request_with_an_empty_guardrail_list_is_served_without_the_applied_header(rig: _Rig) -> None:
prompt: Final = _prompt("no guardrails")
_served_without_guardrail(rig, _ask(rig, None, prompt, []), prompt)
def test_a_key_carrying_a_guardrail_applies_it_and_a_plain_key_does_not(rig: _Rig) -> None:
with rig.gateway.scenario() as scenario:
plain: Final = scenario.key()
guarded: Final = scenario.key(guardrails=[_ATTACHABLE])
plain_prompt: Final = _prompt("plain key")
_served_without_guardrail(rig, _ask(rig, plain, plain_prompt), plain_prompt)
guarded_prompt: Final = _prompt("guarded key")
_served_with_attachable(rig, _ask(rig, guarded, guarded_prompt), guarded_prompt)
def test_a_team_carrying_a_guardrail_applies_it_to_its_keys_only(rig: _Rig) -> None:
with rig.gateway.scenario() as scenario:
team: Final = scenario.team(guardrails=[_ATTACHABLE])
outside: Final = scenario.key()
member: Final = scenario.key(team_id=team)
outside_prompt: Final = _prompt("outside team")
_served_without_guardrail(rig, _ask(rig, outside, outside_prompt), outside_prompt)
member_prompt: Final = _prompt("team key")
_served_with_attachable(rig, _ask(rig, member, member_prompt), member_prompt)
def test_a_during_call_custom_guardrail_rejects_a_request_naming_the_banned_word(rig: _Rig) -> None:
unguarded_prompt: Final = _prompt("what is litellm")
_served_without_guardrail(rig, _ask(rig, None, unguarded_prompt), unguarded_prompt)
refused: Final = _ask(rig, None, _prompt("what is litellm"), [_WORDS])
rig.upstream.drain()
assert refused.status_code >= 400, refused.text
error: Final = object_value(JSON_OBJECT.validate_json(refused.content)["error"])
assert "Guardrail failed words - `litellm` detected" in string_value(error["message"]), refused.text
assert rig.policy.drain() == ()

View file

@ -0,0 +1,274 @@
from __future__ import annotations
import json
import uuid
from hashlib import sha256
from collections.abc import Iterator, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.types.integrations.prometheus import LATENCY_BUCKETS
from tests.integration._support.client import Gateway, eventually, gateway_from_environment, object_value
from tests.integration._support.process import owned_proxy
from tests.integration._support.prometheus_series import Sample, label_values, scrape
from tests.integration._support.wire import Reply, Request, Wire, wire_server
_GOOD: Final = "prometheus-good-endpoint"
_LIMITED: Final = "prometheus-rate-limited-endpoint"
_FAILING: Final = "prometheus-failing-endpoint"
_LATENCY: Final = "prometheus-latency-endpoint"
_END_USER: Final = f"prometheus-end-user-{uuid.uuid4().hex}"
@dataclass(frozen=True, slots=True)
class _Rig:
gateway: Gateway
upstream: Wire
def _respond(request: Request) -> Reply:
if json.loads(request.body)["model"] == "429":
return Reply(
status=429,
body=json.dumps({"error": {"message": "rate limited", "type": "rate_limit_error", "code": "429"}}).encode(),
)
return Reply(
body=json.dumps(
{
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "metered"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4},
}
).encode()
)
def _config(directory: Path, upstream: Wire) -> Path:
configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
params: Final = {"api_key": "sk-fixture", "api_base": f"{upstream.url}/v1"}
configuration["model_list"] = [
*configuration["model_list"],
*(
{
"model_name": name,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
**params,
},
}
for name in (_GOOD, _LATENCY)
),
{"model_name": _LIMITED, "litellm_params": {"model": "openai/429", **params}},
{"model_name": _FAILING, "litellm_params": {"model": "openai/429", **params}},
]
configuration["litellm_settings"]["callbacks"] = ["prometheus"]
configuration["litellm_settings"]["disable_end_user_cost_tracking_prometheus_only"] = True
configuration.setdefault("router_settings", {})["num_retries"] = 0
path: Final = directory / "prometheus-request-metrics.yaml"
path.write_text(yaml.safe_dump(configuration))
return path
@pytest.fixture(scope="module")
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]:
directory: Final = tmp_path_factory.mktemp("prometheus-request-metrics")
with (
wire_server(_respond) as upstream,
gateway_from_environment() as shared,
owned_proxy(shared, directory, {}, config=_config(directory, upstream)) as owned,
):
yield _Rig(owned, upstream)
def _ask(rig: _Rig, model: str, key: str | None = None, **extra: list[str] | str) -> httpx.Response:
return rig.gateway.request(
"POST",
"/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"metrics {uuid.uuid4().hex}"}], **extra},
key=key,
)
def _series(samples: Sequence[Sample], name: str, **labels: str) -> tuple[Sample, ...]:
return tuple(
sample
for sample in samples
if sample.name == name and all(sample.labels.get(label) == value for label, value in labels.items())
)
def _until(rig: _Rig, name: str, **labels: str) -> tuple[Sample, ...]:
return eventually(lambda: _series(scrape(rig.gateway), name, **labels), bool, seconds=30)
def test_a_rate_limited_call_counts_as_a_failed_and_a_429_total_request(rig: _Rig) -> None:
response: Final = _ask(rig, _FAILING)
assert response.status_code == 429, response.text
assert len(rig.upstream.drain()) == 1
failed: Final = _until(
rig,
"litellm_proxy_failed_requests_metric_total",
api_key_alias="None",
exception_class="Openai.RateLimitError",
exception_status="429",
hashed_api_key=LITELLM_PROXY_MASTER_KEY_ALIAS,
requested_model=_FAILING,
route="/chat/completions",
)
assert [sample.value for sample in failed] == [1.0]
totals: Final = _until(
rig,
"litellm_proxy_total_requests_metric_total",
hashed_api_key=LITELLM_PROXY_MASTER_KEY_ALIAS,
requested_model=_FAILING,
status_code="429",
)
assert [sample.value for sample in totals] == [1.0]
def test_a_good_call_exports_latency_histograms_on_the_shared_buckets_without_the_end_user(rig: _Rig) -> None:
response: Final = _ask(rig, _LATENCY, user=_END_USER, tags=["teamB"])
assert response.status_code == 200, response.text
assert len(rig.upstream.drain()) == 1
master: Final = {
"api_key_alias": "None",
"hashed_api_key": LITELLM_PROXY_MASTER_KEY_ALIAS,
"requested_model": _LATENCY,
}
_until(rig, "litellm_request_total_latency_metric_bucket", le="0.005", **master)
_until(rig, "litellm_llm_api_latency_metric_bucket", le="0.005", **master)
for name in ("litellm_request_total_latency_metric_count", "litellm_llm_api_latency_metric_count"):
assert [sample.value for sample in _until(rig, name, **master)] == [1.0], name
samples: Final = scrape(rig.gateway)
assert _END_USER not in label_values(samples)
expected: Final = {str(bucket).replace("inf", "+Inf") for bucket in LATENCY_BUCKETS}
for name in (
"litellm_request_total_latency_metric_bucket",
"litellm_llm_api_latency_metric_bucket",
"litellm_overhead_latency_metric_bucket",
):
assert {sample.labels["le"] for sample in _series(samples, name)} == expected, name
def test_client_side_fallbacks_count_one_success_and_one_failure(rig: _Rig) -> None:
recovered: Final = _ask(rig, _LIMITED, fallbacks=[_GOOD])
assert recovered.status_code == 200, recovered.text
missing: Final = f"unknown-model-{uuid.uuid4().hex[:8]}"
failed: Final = _ask(rig, _LIMITED, fallbacks=[missing])
assert failed.status_code >= 400, failed.text
rig.upstream.drain()
shared: Final = {
"api_key_alias": "None",
"exception_class": "Openai.RateLimitError",
"exception_status": "429",
"hashed_api_key": LITELLM_PROXY_MASTER_KEY_ALIAS,
"requested_model": _LIMITED,
}
succeeded: Final = _until(rig, "litellm_deployment_successful_fallbacks_total", fallback_model=_GOOD, **shared)
assert [sample.value for sample in succeeded] == [1.0]
lost: Final = _until(rig, "litellm_deployment_failed_fallbacks_total", fallback_model=missing, **shared)
assert [sample.value for sample in lost] == [1.0]
@dataclass(frozen=True, slots=True)
class _Budget:
remaining: float
total: float
hours: float
def _budget(samples: Sequence[Sample], scope: str, label: str, identity: str) -> _Budget | None:
remaining: Final = _series(samples, f"litellm_remaining_{scope}_budget_metric", **{label: identity})
total: Final = _series(samples, f"litellm_{scope}_max_budget_metric", **{label: identity})
hours: Final = _series(samples, f"litellm_{scope}_budget_remaining_hours_metric", **{label: identity})
if len(remaining) != 1 or len(total) != 1 or len(hours) != 1:
return None
return _Budget(remaining[0].value, total[0].value, hours[0].value)
def _reconciled(rig: _Rig, scope: str, label: str, identity: str, info: str, field: str, lookup: str) -> _Budget:
def read() -> tuple[_Budget | None, float]:
record: Final = object_value(rig.gateway.get(info, {field: lookup})[_INFO_FIELDS[info]])
return _budget(scrape(rig.gateway), scope, label, identity), float(str(record["max_budget"])) - float(
str(record["spend"])
)
budget, remaining = eventually(
read,
lambda state: (
state[0] is not None and state[0].remaining < 10.0 and abs(state[1] - state[0].remaining) <= 0.001
),
seconds=60,
)
assert budget is not None
assert abs(remaining - budget.remaining) <= 0.001
return budget
_INFO_FIELDS: Final = {"/team/info": "team_info", "/key/info": "info", "/user/info": "user_info"}
def test_a_team_call_exports_remaining_max_and_hours_gauges_matching_team_info(rig: _Rig) -> None:
with rig.gateway.scenario() as scenario:
team: Final = scenario.team(max_budget=10, budget_duration="7d")
key: Final = scenario.key(team_id=team)
assert _ask(rig, _GOOD, key).status_code == 200
assert len(rig.upstream.drain()) == 1
budget: Final = _reconciled(rig, "team", "team", team, "/team/info", "team_id", team)
assert budget.total == 10.0
assert 0 < budget.hours <= 168
def test_a_key_call_exports_remaining_max_and_hours_gauges_matching_key_info(rig: _Rig) -> None:
with rig.gateway.scenario() as scenario:
key: Final = scenario.key(max_budget=10, budget_duration="7d")
assert _ask(rig, _GOOD, key).status_code == 200
assert len(rig.upstream.drain()) == 1
hashed: Final = sha256(key.encode()).hexdigest()
budget: Final = _reconciled(rig, "api_key", "hashed_api_key", hashed, "/key/info", "key", key)
assert budget.total == 10.0
assert 0 <= budget.hours <= 168
def test_a_user_call_exports_remaining_max_and_hours_gauges_matching_user_info(rig: _Rig) -> None:
with rig.gateway.scenario() as scenario:
user: Final = f"prometheus-user-{uuid.uuid4().hex}"
scenario.user(user_id=user, max_budget=10, budget_duration="7d")
key: Final = scenario.key(user_id=user)
assert _ask(rig, _GOOD, key).status_code == 200
assert len(rig.upstream.drain()) == 1
budget: Final = _reconciled(rig, "user", "user", user, "/user/info", "user_id", user)
assert budget.total == 10.0
assert 0 <= budget.hours <= 168
def test_a_user_email_labels_the_spend_and_failed_request_series_of_its_keys(rig: _Rig) -> None:
with rig.gateway.scenario() as scenario:
email: Final = f"prometheus-{uuid.uuid4().hex}@example.com"
user: Final = f"prometheus-email-{uuid.uuid4().hex}"
scenario.user(user_id=user, user_email=email)
key: Final = scenario.key(user_id=user)
assert _ask(rig, _GOOD, key).status_code == 200
assert len(rig.upstream.drain()) == 1
spend: Final = _until(rig, "litellm_spend_metric_total", user_email=email)
assert [(sample.labels["user"], sample.value) for sample in spend] == [(user, pytest.approx(0.005))]
assert email in label_values(scrape(rig.gateway))
assert _ask(rig, _FAILING, key).status_code == 429
assert len(rig.upstream.drain()) == 1
failed: Final = _until(
rig, "litellm_proxy_failed_requests_metric_total", user_email=email, requested_model=_FAILING
)
assert [(sample.labels["user"], sample.value) for sample in failed] == [(user, 1.0)]

View file

@ -0,0 +1,172 @@
from __future__ import annotations
import json
import re
import uuid
from collections.abc import Sequence
from hashlib import sha256
from typing import Final, Literal
import httpx
import pytest
from pydantic import JsonValue
from tests.integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.provider import PROVIDER_URL, SharedProvider
from tests.integration._support.wire import Reply
_CALL_COST: Final = 0.02
_TINY_BUDGET: Final = 0.0000000005
_LIMIT_FIELDS: Final = ("max_budget", "rpm_limit", "tpm_limit")
def _completion() -> Reply:
return Reply(
body=json.dumps(
{
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "scripted"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
}
).encode()
)
def _priced_model(scenario: Scenario) -> str:
return scenario.model(api_base=f"{PROVIDER_URL}/v1", input_cost_per_token=0.001, output_cost_per_token=0.002)
def _ask(gateway: Gateway, model: str, key: str) -> httpx.Response:
return gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"budget {uuid.uuid4().hex}"}]},
key=key,
)
def _served(gateway: Gateway, provider: SharedProvider, model: str, key: str) -> None:
provider.expect(_completion())
response: Final = _ask(gateway, model, key)
assert response.status_code == 200, response.text
assert len(provider.received()) == 1
def _spend_reaches(table: Literal["key", "team"], identity: str, amount: float) -> None:
query: Final = (
'SELECT spend FROM "LiteLLM_VerificationToken" WHERE token = %s'
if table == "key"
else 'SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id = %s'
)
eventually(
lambda: read_rows(query, (identity,)),
lambda rows: len(rows) == 1 and float(str(rows[0]["spend"])) >= amount,
seconds=70,
)
_COST_AND_LIMIT: Final = re.compile(r"Current cost: ([^\s,]+), Max budget: ([^\s,]+)")
def _budget_refusal(
gateway: Gateway, provider: SharedProvider, model: str, key: str, *, spent: float, limit: float
) -> str:
refused: Final = _ask(gateway, model, key)
assert refused.status_code == 422, f"{refused.status_code} {refused.text}"
error: Final = object_value(JSON_OBJECT.validate_json(refused.content)["error"])
assert error["type"] == "budget_exceeded", refused.text
assert error["code"] == "422", refused.text
message: Final = string_value(error["message"])
assert "Budget has been exceeded!" in message, refused.text
figures: Final = _COST_AND_LIMIT.search(message)
assert figures is not None, message
assert float(figures[1]) == pytest.approx(spent), message
assert float(figures[2]) == pytest.approx(limit), message
assert provider.received() == ()
return message
def _hashed(key: str) -> str:
return sha256(key.encode()).hexdigest()
def test_a_key_with_a_tiny_budget_serves_once_and_then_answers_budget_exceeded(
gateway: Gateway, provider: SharedProvider
) -> None:
with gateway.scenario() as scenario:
model: Final = _priced_model(scenario)
key: Final = scenario.key(models=[model], max_budget=_TINY_BUDGET)
_served(gateway, provider, model, key)
_spend_reaches("key", _hashed(key), _CALL_COST)
_budget_refusal(gateway, provider, model, key, spent=_CALL_COST, limit=_TINY_BUDGET)
def test_a_key_with_a_zero_budget_is_refused_before_the_provider_is_called(
gateway: Gateway, provider: SharedProvider
) -> None:
with gateway.scenario() as scenario:
model: Final = _priced_model(scenario)
key: Final = scenario.key(models=[model], max_budget=0)
_budget_refusal(gateway, provider, model, key, spent=0.0, limit=0.0)
def test_a_key_with_room_for_two_calls_serves_both_before_answering_budget_exceeded(
gateway: Gateway, provider: SharedProvider
) -> None:
with gateway.scenario() as scenario:
model: Final = _priced_model(scenario)
limit: Final = _CALL_COST * 1.5
key: Final = scenario.key(models=[model], max_budget=limit)
_served(gateway, provider, model, key)
_spend_reaches("key", _hashed(key), _CALL_COST)
_served(gateway, provider, model, key)
_spend_reaches("key", _hashed(key), _CALL_COST * 2)
_budget_refusal(gateway, provider, model, key, spent=_CALL_COST * 2, limit=limit)
def test_a_team_key_serves_once_and_then_answers_the_team_budget_envelope(
gateway: Gateway, provider: SharedProvider
) -> None:
with gateway.scenario() as scenario:
model: Final = _priced_model(scenario)
team: Final = scenario.team(models=[model], max_budget=_TINY_BUDGET)
key: Final = scenario.key(team_id=team, models=[model])
_served(gateway, provider, model, key)
_spend_reaches("team", team, _CALL_COST)
message: Final = _budget_refusal(gateway, provider, model, key, spent=_CALL_COST, limit=_TINY_BUDGET)
assert f"Team={team}" in message, message
def _limits(record: dict[str, JsonValue]) -> Sequence[JsonValue]:
return [record[field] for field in _LIMIT_FIELDS]
@pytest.mark.parametrize("field", _LIMIT_FIELDS)
def test_a_key_limit_is_set_by_update_and_reset_to_null(gateway: Gateway, field: str) -> None:
with gateway.scenario() as scenario:
key: Final = scenario.key(max_budget=None, rpm_limit=None, tpm_limit=None)
raised: Final = gateway.post("/key/update", {"key": key, field: 10})
assert raised[field] == 10, raised
assert [value for name, value in zip(_LIMIT_FIELDS, _limits(raised)) if name != field] == [None, None]
cleared: Final = gateway.post("/key/update", {"key": key, field: None})
assert _limits(cleared) == [None, None, None], cleared
saved: Final = object_value(gateway.get("/key/info", {"key": key})["info"])
assert _limits(saved) == [None, None, None], saved
@pytest.mark.parametrize("field", _LIMIT_FIELDS)
def test_a_team_limit_is_set_by_update_and_reset_to_null(gateway: Gateway, field: str) -> None:
with gateway.scenario() as scenario:
team: Final = scenario.team(max_budget=None, rpm_limit=None, tpm_limit=None)
raised: Final = object_value(gateway.post("/team/update", {"team_id": team, field: 10})["data"])
assert raised[field] == 10, raised
cleared: Final = object_value(gateway.post("/team/update", {"team_id": team, field: None})["data"])
assert _limits(cleared) == [None, None, None], cleared
saved: Final = object_value(gateway.get("/team/info", {"team_id": team})["team_info"])
assert _limits(saved) == [None, None, None], saved

View file

@ -1,338 +0,0 @@
# What is this?
## Tests slack alerting on proxy logging object
import asyncio
# import logging
# logging.basicConfig(level=logging.DEBUG)
from datetime import datetime
from unittest.mock import AsyncMock, patch
import pytest
import litellm
from litellm.caching.caching import DualCache
from litellm.integrations.SlackAlerting.slack_alerting import (
SlackAlerting,
)
from litellm.proxy._types import CallInfo, Litellm_EntityType
from litellm.proxy.utils import ProxyLogging
from litellm.types.integrations.slack_alerting import AlertType
@pytest.mark.asyncio
async def test_get_api_base():
_pl = ProxyLogging(user_api_key_cache=DualCache())
_pl.update_values(alerting=["slack"], alerting_threshold=100, redis_cache=None)
model = "chatgpt-v-3"
messages = [{"role": "user", "content": "Hey how's it going?"}]
litellm_params = {
"acompletion": True,
"api_key": None,
"api_base": "https://openai-gpt-4-test-v-1.openai.azure.com/",
"force_timeout": 600,
"logger_fn": None,
"verbose": False,
"custom_llm_provider": "azure",
"litellm_call_id": "68f46d2d-714d-4ad8-8137-69600ec8755c",
"model_alias_map": {},
"completion_call_id": None,
"metadata": None,
"model_info": None,
"proxy_server_request": None,
"preset_cache_key": None,
"no-log": False,
"stream_response": {},
}
start_time = datetime.now()
end_time = datetime.now()
time_difference_float, model, api_base, messages = (
_pl.slack_alerting_instance._response_taking_too_long_callback_helper(
kwargs={
"model": model,
"messages": messages,
"litellm_params": litellm_params,
},
start_time=start_time,
end_time=end_time,
)
)
assert api_base is not None
assert isinstance(api_base, str)
assert len(api_base) > 0
request_info = (
f"\nRequest Model: `{model}`\nAPI Base: `{api_base}`\nMessages: `{messages}`"
)
slow_message = f"`Responses are slow - {round(time_difference_float,2)}s response time > Alerting threshold: {100}s`"
await _pl.alerting_handler(
message=slow_message + request_info,
level="Low",
alert_type=AlertType.llm_too_slow,
)
print("passed test_get_api_base")
# Create a mock environment for testing
@pytest.fixture
def mock_env(monkeypatch):
monkeypatch.setenv("SLACK_WEBHOOK_URL", "https://example.com/webhook")
monkeypatch.setenv("LANGFUSE_HOST", "https://cloud.langfuse.com")
monkeypatch.setenv("LANGFUSE_PROJECT_ID", "test-project-id")
# Test the __init__ method
@pytest.fixture
def slack_alerting():
return SlackAlerting(
alerting_threshold=1, internal_usage_cache=DualCache(), alerting=["slack"]
)
# Test for slow LLM responses
# Test for budget crossed
# Test for budget crossed again (should not fire alert 2nd time)
# Test for send_alert - should be called once
@pytest.mark.asyncio
async def test_send_alert(slack_alerting):
import logging
from litellm._logging import verbose_logger
asyncio.create_task(slack_alerting.periodic_flush())
verbose_logger.setLevel(level=logging.DEBUG)
with patch.object(
slack_alerting.async_http_handler, "post", new=AsyncMock()
) as mock_post:
mock_post.return_value.status_code = 200
await slack_alerting.send_alert(
"Test message", "Low", "budget_alerts", alerting_metadata={}
)
await asyncio.sleep(6)
mock_post.assert_awaited_once()
@pytest.mark.asyncio
async def test_daily_reports_completion(slack_alerting):
with patch.object(slack_alerting, "send_alert", new=AsyncMock()) as mock_send_alert:
litellm.callbacks = [slack_alerting]
# on async success
router = litellm.Router(
model_list=[
{
"model_name": "gpt-5.5",
"litellm_params": {
"model": "gpt-5-mini",
},
}
]
)
await router.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
)
await asyncio.sleep(3)
response_val = await slack_alerting.send_daily_reports(router=router)
assert response_val is True
mock_send_alert.assert_awaited_once()
# on async failure
router = litellm.Router(
model_list=[
{
"model_name": "gpt-5.5",
"litellm_params": {"model": "gpt-5-mini", "api_key": "bad_key"},
}
]
)
try:
await router.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hey, how's it going?"}],
)
except Exception as e:
pass
await asyncio.sleep(3)
response_val = await slack_alerting.send_daily_reports(router=router)
assert response_val is True
mock_send_alert.assert_awaited()
# test models with 0 metrics are ignored
# test no alert is sent if all None or 0 metrics
# test user budget crossed alert sent only once, even if user makes multiple calls
# @pytest.mark.asyncio
# async def test_webhook_customer_spend_event():
# """
# Test if customer spend is working as expected
# """
# slack_alerting = SlackAlerting(alerting=["webhook"])
# with patch.object(
# slack_alerting, "send_webhook_alert", new=AsyncMock()
# ) as mock_send_alert:
# user_info = {
# "token": "sk-test-mock-token-606",
# "spend": 1,
# "max_budget": 0,
# "user_id": "ishaan@berri.ai",
# "user_email": "ishaan@berri.ai",
# "key_alias": "my-test-key",
# "projected_exceeded_date": "10/20/2024",
# "projected_spend": 200,
# }
# user_info = CallInfo(**user_info)
# for _ in range(50):
# await slack_alerting.budget_alerts(
# type=alerting_type,
# user_info=user_info,
# )
# mock_send_alert.assert_awaited_once()
@pytest.mark.asyncio
async def test_langfuse_trace_id():
"""
- Unit test for `_add_langfuse_trace_id_to_alert` function in slack_alerting.py
"""
from litellm.integrations.SlackAlerting.utils import add_langfuse_trace_id_to_alert
from litellm.litellm_core_utils.litellm_logging import Logging
litellm.success_callback = ["langfuse"]
litellm_logging_obj = Logging(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="acompletion",
litellm_call_id="1234",
start_time=datetime.now(),
function_id="1234",
)
litellm.completion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hey how's it going?"}],
mock_response="Hey!",
litellm_logging_obj=litellm_logging_obj,
)
await asyncio.sleep(3)
assert litellm_logging_obj.get_trace_id(service_name="langfuse") is not None
slack_alerting = SlackAlerting(
alerting_threshold=32,
alerting=["slack"],
alert_types=[AlertType.llm_exceptions],
internal_usage_cache=DualCache(),
)
trace_url = await add_langfuse_trace_id_to_alert(
request_data={"litellm_logging_obj": litellm_logging_obj}
)
assert trace_url is not None
returned_trace_id = trace_url.split("/")[-1]
assert returned_trace_id == litellm_logging_obj.get_trace_id(
service_name="langfuse"
)
@pytest.mark.parametrize("report_type", ["weekly", "monthly"])
@pytest.mark.asyncio
async def test_spend_report_cache(report_type):
"""
Test that spend reports are only sent once within their period
"""
# Mock prisma client response
mock_spend_data = [
{"team_alias": "team1", "total_spend": 100.0},
{"team_alias": "team2", "total_spend": 200.0},
]
mock_tag_data = [
{"individual_request_tag": "tag1", "total_spend": 150.0},
{"individual_request_tag": "tag2", "total_spend": 150.0},
]
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
# Setup mock for database query
mock_prisma.db.query_raw = AsyncMock(
side_effect=[mock_spend_data, mock_tag_data]
)
slack_alerting = SlackAlerting(
alerting=["webhook"], internal_usage_cache=DualCache()
)
user_info = CallInfo(
token="test_token",
spend=100,
max_budget=1000,
user_id="test@test.com",
user_email="test@test.com",
key_alias="test-key",
event_group=Litellm_EntityType.KEY,
)
with patch.object(
slack_alerting, "send_alert", new=AsyncMock()
) as mock_send_alert:
# First call should send alert
if report_type == "weekly":
await slack_alerting.send_weekly_spend_report()
else:
await slack_alerting.send_monthly_spend_report()
mock_send_alert.assert_called_once()
mock_send_alert.reset_mock()
# Second call should not send alert (cached)
if report_type == "weekly":
await slack_alerting.send_weekly_spend_report()
else:
await slack_alerting.send_monthly_spend_report()
mock_send_alert.assert_not_called()

View file

@ -6,24 +6,15 @@ import litellm
import litellm.vector_stores.main
import json
from typing import Optional
from unittest.mock import AsyncMock, patch, Mock
import pytest
import litellm
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
VectorStorePreCallHook,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import (
StandardLoggingPayload,
)
from litellm.types.vector_stores import (
VectorStoreSearchResponse,
VectorStoreResultContent,
VectorStoreSearchResult,
)
class MockCustomLogger(CustomLogger):
@ -63,113 +54,6 @@ def setup_vector_store_registry():
)
@pytest.mark.asyncio
async def test_vector_store_hook_routes_search_through_proxy_router(
setup_vector_store_registry,
):
proxy_router = Mock()
proxy_router.avector_store_search = AsyncMock(
return_value=VectorStoreSearchResponse(
object="vector_store.search_results.page",
search_query="what is litellm?",
data=[
VectorStoreSearchResult(
score=1.0,
content=[VectorStoreResultContent(text="routed context", type="text")],
)
],
)
)
logging_obj = Mock()
logging_obj.model_call_details = {
"litellm_params": {"metadata": {"user_api_key_team_id": "team-a"}}
}
with patch("litellm.proxy.proxy_server.llm_router", proxy_router):
_, messages, _ = await VectorStorePreCallHook().async_get_chat_completion_prompt(
model="chat-model",
messages=[{"role": "user", "content": "what is litellm?"}],
non_default_params={"vector_store_ids": ["T37J8R4WTM"]},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
litellm_logging_obj=logging_obj,
)
proxy_router.avector_store_search.assert_awaited_once_with(
vector_store_id="T37J8R4WTM",
query="what is litellm?",
custom_llm_provider="bedrock",
metadata={"user_api_key_team_id": "team-a"},
)
assert messages[0]["content"] == "Context:\n\nrouted context\n\n"
@pytest.mark.asyncio
async def test_e2e_bedrock_knowledgebase_retrieval_with_completion(
setup_vector_store_registry,
):
litellm.turn_on_debug()
client = AsyncHTTPHandler()
print("value of litellm.vector_store_registry:", litellm.vector_store_registry)
with patch.object(client, "post") as mock_post:
# Mock the response for the LLM call
mock_response = Mock()
mock_response.status_code = 200
mock_response.headers = {"Content-Type": "application/json"}
# Provide proper JSON response content
mock_response.text = json.dumps(
{
"id": "msg_01ABC123",
"type": "message",
"role": "assistant",
"content": [
{
"type": "text",
"text": "LiteLLM is a library that simplifies LLM API access.",
}
],
"model": "claude-3.5-sonnet",
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 100, "output_tokens": 50},
}
)
mock_response.json = lambda: json.loads(mock_response.text)
mock_post.return_value = mock_response
try:
response = await litellm.acompletion(
model="anthropic/claude-3.5-sonnet",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
client=client,
)
except Exception as e:
print(f"Error: {e}")
# Verify the LLM request was made
mock_post.assert_called_once()
# Verify the request body
print("call args:", mock_post.call_args)
request_body = mock_post.call_args.kwargs["json"]
print("Request body:", json.dumps(request_body, indent=4, default=str))
# Assert content from the knowedge base was applied to the request
# 1. we should have 2 content blocks, the first is the context from the knowledge base, the second is the user message
content = request_body["messages"][0]["content"]
assert len(content) == 2
assert content[0]["type"] == "text"
assert content[1]["type"] == "text"
# 2. the first content block should have the bedrock knowledge base prefix string
# this helps confirm that the context from the knowledge base was applied to the request
assert VectorStorePreCallHook.CONTENT_PREFIX_STRING in content[0]["text"]
@pytest.mark.asyncio
async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call(
setup_vector_store_registry,
@ -214,65 +98,6 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call(
print(f"First search result has {len(first_search_result['data'])} items")
@pytest.mark.asyncio
async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_streaming(
setup_vector_store_registry,
):
"""
Test that the Bedrock Knowledge Base Hook works with streaming and returns search_results in chunks.
"""
# Init client
# litellm.turn_on_debug()
async_client = AsyncHTTPHandler()
response = await litellm.acompletion(
model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
stream=True,
client=async_client,
)
# Collect chunks
chunks = []
search_results_found = False
async for chunk in response:
chunks.append(chunk)
print(f"Chunk: {chunk}")
# Check if this chunk has search_results in provider_specific_fields
if hasattr(chunk, "choices") and chunk.choices:
for choice in chunk.choices:
if hasattr(choice, "delta") and choice.delta:
provider_fields = getattr(
choice.delta, "provider_specific_fields", None
)
if provider_fields and "search_results" in provider_fields:
search_results = provider_fields["search_results"]
print(
f"Found search_results in streaming chunk: {len(search_results)} results"
)
# Verify structure
assert search_results is not None
assert len(search_results) > 0
first_search_result = search_results[0]
assert "object" in first_search_result
assert (
first_search_result["object"]
== "vector_store.search_results.page"
)
assert "data" in first_search_result
assert len(first_search_result["data"]) > 0
search_results_found = True
print(f"Total chunks received: {len(chunks)}")
assert len(chunks) > 0
assert search_results_found, "search_results should be present in streaming chunks"
@pytest.mark.asyncio
async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools(
setup_vector_store_registry,
@ -345,328 +170,6 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools_
print(f" Search was performed and {len(search_results)} result(s) returned")
@pytest.mark.asyncio
async def test_bedrock_kb_request_body_has_transformed_filters(
setup_vector_store_registry,
):
"""
Validate that the Bedrock Knowledge Base request body contains the transformed filters.
"""
captured_request_body: dict = {}
async def fake_async_vector_store_search_handler(
vector_store_id,
query,
vector_store_search_optional_params,
vector_store_provider_config,
custom_llm_provider,
litellm_params,
logging_obj,
embedding_executor=None,
extra_headers=None,
extra_body=None,
timeout=None,
client=None,
_is_async=False,
):
litellm_params_dict = (
litellm_params.model_dump(exclude_none=False)
if hasattr(litellm_params, "model_dump")
else dict(litellm_params)
)
api_base = vector_store_provider_config.get_complete_url(
api_base=litellm_params_dict.get("api_base"),
litellm_params=litellm_params_dict,
)
url, request_body = (
vector_store_provider_config.transform_search_vector_store_request(
vector_store_id=vector_store_id,
query=query,
vector_store_search_optional_params=vector_store_search_optional_params,
api_base=api_base,
litellm_logging_obj=logging_obj,
litellm_params=litellm_params_dict,
extra_body=None,
)
)
captured_request_body["url"] = url
captured_request_body["body"] = request_body
return VectorStoreSearchResponse(
object="vector_store.search_results.page",
search_query=query if isinstance(query, str) else " ".join(query),
data=[
VectorStoreSearchResult(
score=0.9,
content=[
VectorStoreResultContent(
text="LiteLLM is a library", type="text"
)
],
)
],
)
with patch.object(
litellm.vector_stores.main.base_llm_http_handler,
"async_vector_store_search_handler",
new=AsyncMock(side_effect=fake_async_vector_store_search_handler),
):
response = await litellm.acompletion(
model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}",
messages=[{"role": "user", "content": "what is litellm?"}],
max_tokens=10,
tools=[
{
"type": "file_search",
"vector_store_ids": ["T37J8R4WTM"],
"filters": {
"key": "user_id",
"value": "fake-user-id",
"operator": "eq",
},
}
],
)
assert response is not None
print(
"captured_request_body:",
json.dumps(captured_request_body, indent=4, default=str),
)
assert "body" in captured_request_body, "Bedrock KB request body was not captured"
vector_search = captured_request_body["body"]["retrievalConfiguration"][
"vectorSearchConfiguration"
]
aws_filter = vector_search["filter"]
assert "equals" in aws_filter, f"Expected 'equals' in AWS format, got: {aws_filter}"
assert aws_filter["equals"]["key"] == "user_id"
assert aws_filter["equals"]["value"] == "fake-user-id"
print("✅ Filters transformed correctly: OpenAI format -> AWS Bedrock format")
@pytest.mark.asyncio
async def test_openai_with_knowledge_base_mock_openai(setup_vector_store_registry):
"""
Tests that knowledge base content is correctly passed to the OpenAI API call
"""
litellm.set_verbose = True
from openai import AsyncOpenAI
client = AsyncOpenAI(api_key="fake-api-key")
# Variable to capture the request
captured_request = {}
with patch.object(
client.chat.completions.with_raw_response, "create"
) as mock_client:
# Create async mock that returns proper structure
async def mock_create(**kwargs):
mock_response = Mock()
mock_response.choices = [
Mock(
message=Mock(content="Mock response from OpenAI", role="assistant")
)
]
mock_response.usage = Mock(
prompt_tokens=100, completion_tokens=50, total_tokens=150
)
mock_response.id = "chatcmpl-123"
mock_response.object = "chat.completion"
mock_response.created = 1234567890
mock_response.model = "gpt-5.5"
# Store the request for verification
captured_request.update(kwargs)
# Return wrapper with parse method
wrapper = Mock()
wrapper.parse.return_value = mock_response
return wrapper
mock_client.side_effect = mock_create
try:
await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
client=client,
)
except Exception as e:
print(f"Error: {e}")
# Verify the API was called
mock_client.assert_called_once()
request_body = captured_request
# Verify the request contains messages with knowledge base context
assert "messages" in request_body
messages = request_body["messages"]
# We expect at least 2 messages:
# 1. User message with the knowledge base context
# 2. User message with the question
assert len(messages) >= 2
print("request messages:", json.dumps(messages, indent=4, default=str))
# assert message[0] is the user message with the knowledge base context
assert messages[0]["role"] == "user"
assert VectorStorePreCallHook.CONTENT_PREFIX_STRING in messages[0]["content"]
@pytest.mark.asyncio
async def test_openai_with_vector_store_ids_in_tool_call_mock_openai(
setup_vector_store_registry,
):
"""
Tests that vector store ids can be passed as tools
This is the OpenAI format
"""
litellm.set_verbose = True
from openai import AsyncOpenAI
client = AsyncOpenAI(api_key="fake-api-key")
# Variable to capture the request
captured_request = {}
with patch.object(
client.chat.completions.with_raw_response, "create"
) as mock_client:
# Create async mock that returns proper structure
async def mock_create(**kwargs):
mock_response = Mock()
mock_response.choices = [
Mock(
message=Mock(content="Mock response from OpenAI", role="assistant")
)
]
mock_response.usage = Mock(
prompt_tokens=100, completion_tokens=50, total_tokens=150
)
mock_response.id = "chatcmpl-123"
mock_response.object = "chat.completion"
mock_response.created = 1234567890
mock_response.model = "gpt-5.5"
# Store the request for verification
captured_request.update(kwargs)
# Return wrapper with parse method
wrapper = Mock()
wrapper.parse.return_value = mock_response
return wrapper
mock_client.side_effect = mock_create
try:
await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
tools=[{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]}],
client=client,
)
except Exception as e:
print(f"Error: {e}")
# Verify the API was called
mock_client.assert_called_once()
request_body = captured_request
print("request body:", json.dumps(request_body, indent=4, default=str))
# Verify the request contains messages with knowledge base context
assert "messages" in request_body
messages = request_body["messages"]
# We expect at least 2 messages:
# 1. User message with the knowledge base context
# 2. User message with the question
assert len(messages) >= 2
print("request messages:", json.dumps(messages, indent=4, default=str))
# assert message[0] is the user message with the knowledge base context
assert messages[0]["role"] == "user"
assert VectorStorePreCallHook.CONTENT_PREFIX_STRING in messages[0]["content"]
# assert that the tool call was not sent to the upstream llm API if it's a litellm vector store
assert "tools" not in request_body
@pytest.mark.asyncio
async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_registry):
"""Ensure unrecognized vector store tools are forwarded to the provider"""
from openai import AsyncOpenAI
client = AsyncOpenAI(api_key="fake-api-key")
# Variable to capture the request
captured_request = {}
with patch.object(
client.chat.completions.with_raw_response, "create"
) as mock_client:
# Create async mock that returns proper structure
async def mock_create(**kwargs):
mock_response = Mock()
mock_response.choices = [
Mock(
message=Mock(content="Mock response from OpenAI", role="assistant")
)
]
mock_response.usage = Mock(
prompt_tokens=100, completion_tokens=50, total_tokens=150
)
mock_response.id = "chatcmpl-123"
mock_response.object = "chat.completion"
mock_response.created = 1234567890
mock_response.model = "gpt-5.5"
# Store the request for verification
captured_request.update(kwargs)
# Return wrapper with parse method
wrapper = Mock()
wrapper.parse.return_value = mock_response
return wrapper
mock_client.side_effect = mock_create
try:
await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
tools=[
{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]},
{"type": "file_search", "vector_store_ids": ["unknownVS"]},
],
client=client,
)
except Exception as e:
print(f"Error: {e}")
mock_client.assert_called_once()
request_body = captured_request
assert "messages" in request_body
messages = request_body["messages"]
assert len(messages) >= 2
assert messages[0]["role"] == "user"
assert VectorStorePreCallHook.CONTENT_PREFIX_STRING in messages[0]["content"]
assert "tools" in request_body
tools = request_body["tools"]
assert len(tools) == 1
assert tools[0]["vector_store_ids"] == ["unknownVS"]
# @pytest.mark.asyncio
# async def test_logging_with_knowledge_base_hook(setup_vector_store_registry):
# """
@ -723,118 +226,3 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist
@pytest.mark.asyncio
async def test_provider_specific_fields_in_proxy_http_response(
setup_vector_store_registry,
):
"""
Test that provider_specific_fields (like search_results) are included
in the proxy HTTP JSON response, not just in Python SDK objects.
This test catches serialization bugs where exclude=True would strip
provider_specific_fields from the HTTP response.
"""
from fastapi.testclient import TestClient
from litellm.proxy.proxy_server import app, initialize
from unittest.mock import patch as mock_patch
# Initialize proxy
await initialize(
model="gpt-5-mini",
alias=None,
api_base=None,
debug=False,
temperature=None,
max_tokens=None,
request_timeout=600,
max_budget=None,
drop_params=True,
add_function_to_prompt=False,
headers=None,
save=False,
use_queue=False,
config=None,
)
# Create test client
client = TestClient(app)
# Create mock response with provider_specific_fields
mock_response = litellm.ModelResponse(
id="test-123",
model="gpt-5-mini",
created=1234567890,
object="chat.completion",
)
# Create message with provider_specific_fields
mock_message = litellm.Message(
content="LiteLLM is a tool that simplifies working with multiple LLMs.",
role="assistant",
provider_specific_fields={
"search_results": [
{
"object": "vector_store.search_results.page",
"search_query": "what is litellm?",
"data": [
{
"score": 0.95,
"content": [{"text": "Test content", "type": "text"}],
"file_id": "test-file",
"filename": "test.txt",
}
],
}
]
},
)
mock_choice = litellm.Choices(finish_reason="stop", index=0, message=mock_message)
mock_response.choices = [mock_choice]
mock_response.usage = litellm.Usage(
prompt_tokens=10, completion_tokens=20, total_tokens=30
)
# Patch the completion call at the proxy level
with mock_patch("litellm.acompletion", new=AsyncMock(return_value=mock_response)):
# Make HTTP request to proxy
response = client.post(
"/v1/chat/completions",
json={
"model": "gpt-5-mini",
"messages": [{"role": "user", "content": "What is litellm?"}],
},
)
# Check HTTP response
assert response.status_code == 200
result = response.json()
print("HTTP Response JSON:", json.dumps(result, indent=2))
# THE KEY ASSERTIONS - These would FAIL with exclude=True!
assert "choices" in result
assert len(result["choices"]) > 0
choice = result["choices"][0]
assert "message" in choice
message = choice["message"]
# Verify provider_specific_fields is in the JSON response
assert (
"provider_specific_fields" in message
), "provider_specific_fields missing from HTTP JSON response! This means exclude=True is preventing serialization."
assert "search_results" in message["provider_specific_fields"]
search_results = message["provider_specific_fields"]["search_results"]
assert len(search_results) > 0
# Verify search result structure
first_result = search_results[0]
assert first_result["object"] == "vector_store.search_results.page"
assert "data" in first_result
assert len(first_result["data"]) > 0
print("✅ provider_specific_fields successfully serialized in HTTP response")

View file

@ -1,160 +0,0 @@
import traceback
from litellm._uuid import uuid
import pytest
from dotenv import load_dotenv
from fastapi import Request
from fastapi.routing import APIRoute
load_dotenv()
import io
import time
import json
# this file is to test litellm/proxy
import litellm
import asyncio
from typing import Optional
from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
StandardBuiltInToolCostTracking,
)
class TestCustomLogger(CustomLogger):
def __init__(self):
self.recorded_usage: Optional[Usage] = None
self.standard_logging_payload: Optional[StandardLoggingPayload] = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
standard_logging_payload = kwargs.get("standard_logging_object")
self.standard_logging_payload = standard_logging_payload
print(
"standard_logging_payload",
json.dumps(standard_logging_payload, indent=4, default=str),
)
self.recorded_usage = Usage(
prompt_tokens=standard_logging_payload.get("prompt_tokens"),
completion_tokens=standard_logging_payload.get("completion_tokens"),
total_tokens=standard_logging_payload.get("total_tokens"),
)
pass
async def _setup_web_search_test():
"""Helper function to setup common test requirements"""
litellm.turn_on_debug()
test_custom_logger = TestCustomLogger()
litellm.callbacks = [test_custom_logger]
return test_custom_logger
async def _verify_web_search_cost(test_custom_logger, expected_context_size):
"""Helper function to verify web search costs"""
await asyncio.sleep(1)
standard_logging_payload = test_custom_logger.standard_logging_payload
response = standard_logging_payload.get("response")
response_cost = standard_logging_payload.get("response_cost")
assert response_cost is not None
# Calculate token cost
model_map_information = standard_logging_payload["model_map_information"]
model_map_value: ModelInfoBase = model_map_information["model_map_value"]
total_token_cost = (
standard_logging_payload["prompt_tokens"]
* model_map_value["input_cost_per_token"]
) + (
standard_logging_payload["completion_tokens"]
* model_map_value["output_cost_per_token"]
)
# Verify total cost
if StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
response
):
assert (
response_cost
== total_token_cost
+ model_map_value["search_context_cost_per_query"][expected_context_size]
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"web_search_options,expected_context_size",
[
(None, "search_context_size_medium"),
({"search_context_size": "low"}, "search_context_size_low"),
({"search_context_size": "high"}, "search_context_size_high"),
],
)
async def test_openai_web_search_logging_cost_tracking(
web_search_options, expected_context_size
):
"""Test web search cost tracking with different search context sizes"""
test_custom_logger = await _setup_web_search_test()
request_kwargs = {
"model": "openai/gpt-5-search-api",
"messages": [
{
"role": "user",
"content": f"What was a positive news story from today? {uuid.uuid4()}",
}
],
}
if web_search_options is not None:
request_kwargs["web_search_options"] = web_search_options
response = await litellm.acompletion(**request_kwargs)
await _verify_web_search_cost(test_custom_logger, expected_context_size)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"tools_config,expected_context_size,stream",
[
(
[{"type": "web_search_preview", "search_context_size": "low"}],
"search_context_size_low",
True,
),
(
[{"type": "web_search_preview", "search_context_size": "low"}],
"search_context_size_low",
False,
),
([{"type": "web_search_preview"}], "search_context_size_medium", True),
([{"type": "web_search_preview"}], "search_context_size_medium", False),
],
)
async def test_openai_responses_api_web_search_cost_tracking(
tools_config, expected_context_size, stream
):
"""Test web search cost tracking with different search context sizes and streaming options"""
test_custom_logger = await _setup_web_search_test()
response = await litellm.aresponses(
model="openai/gpt-4o",
input=[
{"role": "user", "content": "What was a positive news story from today?"}
],
tools=tools_config,
stream=stream,
)
if stream is True:
async for chunk in response:
print("chunk", chunk)
else:
print("response", response)
await asyncio.sleep(1)
if StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
test_custom_logger.standard_logging_payload.get("response")
):
await _verify_web_search_cost(test_custom_logger, expected_context_size)

View file

@ -1,754 +0,0 @@
### What this tests ####
## This test asserts the type of data passed into each method of the custom callback handler
import asyncio
import inspect
import os
import time
import traceback
from datetime import datetime
import pytest
from typing import List, Literal, Optional
import litellm
from litellm import Cache, Router
from litellm.integrations.custom_logger import CustomLogger
# Test Scenarios (test across completion, streaming, embedding)
## 1: Pre-API-Call
## 2: Post-API-Call
## 3: On LiteLLM Call success
## 4: On LiteLLM Call failure
## fallbacks
## retries
# Test cases
## 1. Simple Azure OpenAI acompletion + streaming call
## 2. Simple Azure OpenAI aembedding call
## 3. Azure OpenAI acompletion + streaming call with retries
## 4. Azure OpenAI aembedding call with retries
## 5. Azure OpenAI acompletion + streaming call with fallbacks
## 6. Azure OpenAI aembedding call with fallbacks
## Test interfaces
## 1. router.completion() + router.embeddings()
## 2. proxy.completions + proxy.embeddings
litellm.num_retries = 0
class CompletionCustomHandler(
CustomLogger
): # https://docs.litellm.ai/docs/observability/custom_callback#callback-class
"""
The set of expected inputs to a custom handler for a
"""
# Class variables or attributes
def __init__(self):
self.errors = []
self.states: Optional[
List[
Literal[
"sync_pre_api_call",
"async_pre_api_call",
"post_api_call",
"sync_stream",
"async_stream",
"sync_success",
"async_success",
"sync_failure",
"async_failure",
]
]
] = []
def log_pre_api_call(self, model, messages, kwargs):
try:
print(f"received kwargs in pre-input: {kwargs}")
self.states.append("sync_pre_api_call")
## MODEL
assert isinstance(model, str)
## MESSAGES
assert isinstance(messages, list)
## KWARGS
assert isinstance(kwargs["model"], str)
assert isinstance(kwargs["messages"], list)
assert isinstance(kwargs["optional_params"], dict)
assert isinstance(kwargs["litellm_params"], dict)
assert isinstance(kwargs["start_time"], (datetime, type(None)))
assert isinstance(kwargs["stream"], bool)
assert isinstance(kwargs["user"], (str, type(None)))
### ROUTER-SPECIFIC KWARGS
assert isinstance(kwargs["litellm_params"]["metadata"], dict)
assert isinstance(kwargs["litellm_params"]["metadata"]["model_group"], str)
assert isinstance(kwargs["litellm_params"]["metadata"]["deployment"], str)
assert isinstance(kwargs["litellm_params"]["model_info"], dict)
assert isinstance(kwargs["litellm_params"]["model_info"]["id"], str)
assert isinstance(
kwargs["litellm_params"]["proxy_server_request"], (str, type(None))
)
assert isinstance(
kwargs["litellm_params"]["preset_cache_key"], (str, type(None))
)
assert isinstance(kwargs["litellm_params"]["stream_response"], dict)
except Exception as e:
print(f"Assertion Error: {traceback.format_exc()}")
self.errors.append(traceback.format_exc())
def log_post_api_call(self, kwargs, response_obj, start_time, end_time):
try:
self.states.append("post_api_call")
## START TIME
assert isinstance(start_time, datetime)
## END TIME
assert end_time == None
## RESPONSE OBJECT
assert response_obj == None
## KWARGS
assert isinstance(kwargs["model"], str)
assert isinstance(kwargs["messages"], list)
assert isinstance(kwargs["optional_params"], dict)
assert isinstance(kwargs["litellm_params"], dict)
assert isinstance(kwargs["start_time"], (datetime, type(None)))
assert isinstance(kwargs["stream"], bool)
assert isinstance(kwargs["user"], (str, type(None)))
assert isinstance(kwargs["input"], (list, dict, str))
assert isinstance(kwargs["api_key"], (str, type(None)))
assert (
isinstance(
kwargs["original_response"], (str, litellm.CustomStreamWrapper)
)
or inspect.iscoroutine(kwargs["original_response"])
or inspect.isasyncgen(kwargs["original_response"])
)
assert isinstance(kwargs["additional_args"], (dict, type(None)))
assert isinstance(kwargs["log_event_type"], str)
### ROUTER-SPECIFIC KWARGS
assert isinstance(kwargs["litellm_params"]["metadata"], dict)
assert isinstance(kwargs["litellm_params"]["metadata"]["model_group"], str)
assert isinstance(kwargs["litellm_params"]["metadata"]["deployment"], str)
assert isinstance(kwargs["litellm_params"]["model_info"], dict)
assert isinstance(kwargs["litellm_params"]["model_info"]["id"], str)
assert isinstance(
kwargs["litellm_params"]["proxy_server_request"], (str, type(None))
)
assert isinstance(
kwargs["litellm_params"]["preset_cache_key"], (str, type(None))
)
assert isinstance(kwargs["litellm_params"]["stream_response"], dict)
except Exception:
print(f"Assertion Error: {traceback.format_exc()}")
self.errors.append(traceback.format_exc())
async def async_log_stream_event(self, kwargs, response_obj, start_time, end_time):
try:
self.states.append("async_stream")
## START TIME
assert isinstance(start_time, datetime)
## END TIME
assert isinstance(end_time, datetime)
## RESPONSE OBJECT
assert isinstance(response_obj, litellm.ModelResponseStream)
## KWARGS
assert isinstance(kwargs["model"], str)
assert isinstance(kwargs["messages"], list) and isinstance(
kwargs["messages"][0], dict
)
assert isinstance(kwargs["optional_params"], dict)
assert isinstance(kwargs["litellm_params"], dict)
assert isinstance(kwargs["start_time"], (datetime, type(None)))
assert isinstance(kwargs["stream"], bool)
assert isinstance(kwargs["user"], (str, type(None)))
assert (
isinstance(kwargs["input"], list)
and isinstance(kwargs["input"][0], dict)
) or isinstance(kwargs["input"], (dict, str))
assert isinstance(kwargs["api_key"], (str, type(None)))
assert (
isinstance(
kwargs["original_response"], (str, litellm.CustomStreamWrapper)
)
or inspect.isasyncgen(kwargs["original_response"])
or inspect.iscoroutine(kwargs["original_response"])
)
assert isinstance(kwargs["additional_args"], (dict, type(None)))
assert isinstance(kwargs["log_event_type"], str)
except Exception:
print(f"Assertion Error: {traceback.format_exc()}")
self.errors.append(traceback.format_exc())
def log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
self.states.append("sync_success")
## START TIME
assert isinstance(start_time, datetime)
## END TIME
assert isinstance(end_time, datetime)
## RESPONSE OBJECT
assert isinstance(response_obj, litellm.ModelResponse)
## KWARGS
assert isinstance(kwargs["model"], str)
assert isinstance(kwargs["messages"], list) and isinstance(
kwargs["messages"][0], dict
)
assert isinstance(kwargs["optional_params"], dict)
assert isinstance(kwargs["litellm_params"], dict)
assert isinstance(kwargs["start_time"], (datetime, type(None)))
assert isinstance(kwargs["stream"], bool)
assert isinstance(kwargs["user"], (str, type(None)))
assert (
isinstance(kwargs["input"], list)
and isinstance(kwargs["input"][0], dict)
) or isinstance(kwargs["input"], (dict, str))
assert isinstance(kwargs["api_key"], (str, type(None)))
assert isinstance(
kwargs["original_response"], (str, litellm.CustomStreamWrapper)
)
assert isinstance(kwargs["additional_args"], (dict, type(None)))
assert isinstance(kwargs["log_event_type"], str)
assert kwargs["cache_hit"] is None or isinstance(kwargs["cache_hit"], bool)
except Exception:
print(f"Assertion Error: {traceback.format_exc()}")
self.errors.append(traceback.format_exc())
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
try:
self.states.append("sync_failure")
## START TIME
assert isinstance(start_time, datetime)
## END TIME
assert isinstance(end_time, datetime)
## RESPONSE OBJECT
assert response_obj == None
## KWARGS
assert isinstance(kwargs["model"], str)
assert isinstance(kwargs["messages"], list) and isinstance(
kwargs["messages"][0], dict
)
assert isinstance(kwargs["optional_params"], dict)
assert isinstance(kwargs["litellm_params"], dict)
assert isinstance(kwargs["start_time"], (datetime, type(None)))
assert isinstance(kwargs["stream"], bool)
assert isinstance(kwargs["user"], (str, type(None)))
assert (
isinstance(kwargs["input"], list)
and isinstance(kwargs["input"][0], dict)
) or isinstance(kwargs["input"], (dict, str))
assert isinstance(kwargs["api_key"], (str, type(None)))
assert (
isinstance(
kwargs["original_response"], (str, litellm.CustomStreamWrapper)
)
or kwargs["original_response"] == None
)
assert isinstance(kwargs["additional_args"], (dict, type(None)))
assert isinstance(kwargs["log_event_type"], str)
except Exception:
print(f"Assertion Error: {traceback.format_exc()}")
self.errors.append(traceback.format_exc())
async def async_log_pre_api_call(self, model, messages, kwargs):
try:
"""
No-op.
Not implemented yet.
"""
pass
except Exception as e:
print(f"Assertion Error: {traceback.format_exc()}")
self.errors.append(traceback.format_exc())
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
print("CompletionCustomHandler.async_log_success_event, kwargs: ", kwargs)
self.states.append("async_success")
print(
"############### CompletionCustomHandler async success, kwargs: ",
kwargs,
)
## START TIME
assert isinstance(start_time, datetime)
## END TIME
assert isinstance(end_time, datetime)
## RESPONSE OBJECT
assert isinstance(
response_obj, (litellm.ModelResponse, litellm.EmbeddingResponse)
)
## KWARGS
assert isinstance(kwargs["model"], str)
# checking we use base_model for azure cost calculation
base_model = litellm.utils.get_base_model_from_metadata(
model_call_details=kwargs
)
if (
kwargs["model"] == "chatgpt-v-3"
and base_model is not None
and kwargs["stream"] != True
):
# when base_model is set for azure, we should use pricing for the base_model
# this checks response_cost == litellm.cost_per_token(model=base_model)
assert isinstance(kwargs["response_cost"], float)
response_cost = kwargs["response_cost"]
print(
f"response_cost: {response_cost}, for model: {kwargs['model']} and base_model: {base_model}"
)
prompt_tokens = response_obj.usage.prompt_tokens
completion_tokens = response_obj.usage.completion_tokens
# ensure the pricing is based on the base_model here
prompt_price, completion_price = litellm.cost_per_token(
model=base_model,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
)
expected_price = prompt_price + completion_price
print(f"expected price: {expected_price}")
assert (
response_cost == expected_price
), f"response_cost: {response_cost} != expected_price: {expected_price}. For model: {kwargs['model']} and base_model: {base_model}. should have used base_model for price"
assert isinstance(kwargs["messages"], list)
assert isinstance(kwargs["optional_params"], dict)
assert isinstance(kwargs["litellm_params"], dict)
assert isinstance(kwargs["start_time"], (datetime, type(None)))
assert isinstance(kwargs["stream"], bool)
assert isinstance(kwargs["user"], (str, type(None)))
assert isinstance(kwargs["input"], (list, dict, str))
assert isinstance(kwargs["api_key"], (str, type(None)))
assert (
isinstance(
kwargs["original_response"], (str, litellm.CustomStreamWrapper)
)
or inspect.isasyncgen(kwargs["original_response"])
or inspect.iscoroutine(kwargs["original_response"])
)
assert isinstance(kwargs["additional_args"], (dict, type(None)))
assert isinstance(kwargs["log_event_type"], str)
assert kwargs["cache_hit"] is None or isinstance(kwargs["cache_hit"], bool)
### ROUTER-SPECIFIC KWARGS
assert isinstance(kwargs["litellm_params"]["metadata"], dict)
assert isinstance(kwargs["litellm_params"]["metadata"]["model_group"], str)
assert isinstance(kwargs["litellm_params"]["metadata"]["deployment"], str)
assert isinstance(kwargs["litellm_params"]["model_info"], dict)
assert isinstance(kwargs["litellm_params"]["model_info"]["id"], str)
assert isinstance(
kwargs["litellm_params"]["proxy_server_request"], (str, type(None))
)
assert isinstance(
kwargs["litellm_params"]["preset_cache_key"], (str, type(None))
)
assert isinstance(kwargs["litellm_params"]["stream_response"], dict)
except Exception:
print(f"Assertion Error: {traceback.format_exc()}")
self.errors.append(traceback.format_exc())
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
try:
print(f"received original response: {kwargs['original_response']}")
self.states.append("async_failure")
## START TIME
assert isinstance(start_time, datetime)
## END TIME
assert isinstance(end_time, datetime)
## RESPONSE OBJECT
assert response_obj == None
## KWARGS
assert isinstance(kwargs["model"], str)
assert isinstance(kwargs["messages"], list)
assert isinstance(kwargs["optional_params"], dict)
assert isinstance(kwargs["litellm_params"], dict)
assert isinstance(kwargs["start_time"], (datetime, type(None)))
assert isinstance(kwargs["stream"], bool)
assert isinstance(kwargs["user"], (str, type(None)))
assert isinstance(kwargs["input"], (list, str, dict))
assert isinstance(kwargs["api_key"], (str, type(None)))
assert (
isinstance(
kwargs["original_response"], (str, litellm.CustomStreamWrapper)
)
or inspect.isasyncgen(kwargs["original_response"])
or inspect.iscoroutine(kwargs["original_response"])
or kwargs["original_response"] == None
)
assert isinstance(kwargs["additional_args"], (dict, type(None)))
assert isinstance(kwargs["log_event_type"], str)
except Exception:
print(f"Assertion Error: {traceback.format_exc()}")
self.errors.append(traceback.format_exc())
# Simple Azure OpenAI call
## COMPLETION
# @pytest.mark.flaky(retries=5, delay=1)
@pytest.mark.asyncio
async def test_async_chat_azure():
try:
customHandler_completion_azure_router = CompletionCustomHandler()
customHandler_streaming_azure_router = CompletionCustomHandler()
customHandler_failure = CompletionCustomHandler()
litellm.callbacks = [customHandler_completion_azure_router]
litellm.set_verbose = True
model_list = [
{
"model_name": "gpt-4.1-nano", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
"api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
},
"model_info": {"base_model": "azure/gpt-4.1-mini"},
"tpm": 240000,
"rpm": 1800,
},
]
router = Router(model_list=model_list, num_retries=0) # type: ignore
response = await router.acompletion(
model="gpt-4.1-nano",
messages=[{"role": "user", "content": "Hi 👋 - i'm openai"}],
)
print("got response, sleeping 5 seconds....")
await asyncio.sleep(5)
assert len(customHandler_completion_azure_router.errors) == 0
assert (
len(customHandler_completion_azure_router.states) == 3
) # pre, post, success
# streaming
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [customHandler_streaming_azure_router]
router2 = Router(model_list=model_list, num_retries=0) # type: ignore
response = await router2.acompletion(
model="gpt-4.1-nano",
messages=[{"role": "user", "content": "Hi 👋 - i'm openai"}],
stream=True,
)
async for chunk in response:
print(f"async azure router chunk: {chunk}")
continue
await asyncio.sleep(5)
print(f"customHandler.states: {customHandler_streaming_azure_router.states}")
assert len(customHandler_streaming_azure_router.errors) == 0
assert (
len(customHandler_streaming_azure_router.states) >= 3
) # pre, post, stream (multiple times), success
# failure
model_list = [
{
"model_name": "gpt-5-mini", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4o-new-test",
"api_key": "my-bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
]
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [customHandler_failure]
router3 = Router(model_list=model_list, num_retries=0) # type: ignore
try:
response = await router3.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hi 👋 - i'm openai"}],
)
print(f"response in router3 acompletion: {response}")
except Exception:
pass
await asyncio.sleep(5)
print(f"customHandler.states: {customHandler_failure.states}")
assert len(customHandler_failure.errors) == 0
assert len(customHandler_failure.states) == 3 # pre, post, failure
assert "async_failure" in customHandler_failure.states
except Exception as e:
print(f"Assertion Error: {traceback.format_exc()}")
pytest.fail(f"An exception occurred - {str(e)}")
## EMBEDDING
@pytest.mark.asyncio
async def test_async_embedding_azure():
try:
customHandler = CompletionCustomHandler()
customHandler_failure = CompletionCustomHandler()
litellm.callbacks = [customHandler]
model_list = [
{
"model_name": "azure-embedding-model", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/text-embedding-ada-002",
"api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
]
router = Router(model_list=model_list) # type: ignore
response = await router.aembedding(
model="azure-embedding-model", input=["hello from litellm!"]
)
await asyncio.sleep(2)
assert len(customHandler.errors) == 0
assert len(customHandler.states) == 3 # pre, post, success
# failure
model_list = [
{
"model_name": "azure-embedding-model", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/text-embedding-ada-002",
"api_key": "my-bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
]
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [customHandler_failure]
router3 = Router(model_list=model_list, num_retries=0) # type: ignore
try:
response = await router3.aembedding(
model="azure-embedding-model", input=["hello from litellm!"]
)
print(f"response in router3 aembedding: {response}")
except Exception:
pass
await asyncio.sleep(1)
print(f"customHandler.states: {customHandler_failure.states}")
assert len(customHandler_failure.errors) == 0
assert len(customHandler_failure.states) == 3 # pre, post, failure
assert "async_failure" in customHandler_failure.states
except Exception as e:
print(f"Assertion Error: {traceback.format_exc()}")
pytest.fail(f"An exception occurred - {str(e)}")
# asyncio.run(test_async_embedding_azure())
# Azure OpenAI call w/ Fallbacks
## COMPLETION
@pytest.mark.asyncio
async def test_async_chat_azure_with_fallbacks():
try:
customHandler_fallbacks = CompletionCustomHandler()
litellm.callbacks = [customHandler_fallbacks]
litellm.set_verbose = True
# with fallbacks
model_list = [
{
"model_name": "gpt-5-mini", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
"api_key": "my-bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
{
"model_name": "gpt-3.5-turbo-16k",
"litellm_params": {
"model": "gpt-3.5-turbo-16k",
},
"tpm": 240000,
"rpm": 1800,
},
]
router = Router(
model_list=model_list,
fallbacks=[{"gpt-5-mini": ["gpt-3.5-turbo-16k"]}],
retry_policy=litellm.router.RetryPolicy(
AuthenticationErrorRetries=0,
),
) # type: ignore
response = await router.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hi 👋 - i'm openai"}],
)
await asyncio.sleep(2)
print(f"customHandler_fallbacks.states: {customHandler_fallbacks.states}")
assert len(customHandler_fallbacks.errors) == 0
assert (
len(customHandler_fallbacks.states) == 6
) # pre, post, failure, pre, post, success
litellm.callbacks = []
except Exception as e:
print(f"Assertion Error: {traceback.format_exc()}")
pytest.fail(f"An exception occurred - {str(e)}")
# asyncio.run(test_async_chat_azure_with_fallbacks())
# CACHING
## Test Azure - completion, embedding
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=1)
async def test_async_completion_azure_caching():
customHandler_caching = CompletionCustomHandler()
litellm.cache = Cache(
type="redis",
host=os.environ["REDIS_HOST"],
port=os.environ["REDIS_PORT"],
password=os.environ["REDIS_PASSWORD"],
)
litellm.callbacks = [customHandler_caching]
unique_time = time.time()
model_list = [
{
"model_name": "gpt-4.1-nano", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
"api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
{
"model_name": "gpt-3.5-turbo-16k",
"litellm_params": {
"model": "gpt-3.5-turbo-16k",
},
"tpm": 240000,
"rpm": 1800,
},
]
router = Router(model_list=model_list) # type: ignore
response1 = await router.acompletion(
model="gpt-4.1-nano",
messages=[
{"role": "user", "content": f"Hi 👋 - i'm async azure {unique_time}"}
],
caching=True,
)
await asyncio.sleep(1)
print(f"customHandler_caching.states pre-cache hit: {customHandler_caching.states}")
response2 = await router.acompletion(
model="gpt-4.1-nano",
messages=[
{"role": "user", "content": f"Hi 👋 - i'm async azure {unique_time}"}
],
caching=True,
)
await asyncio.sleep(1) # success callbacks are done in parallel
print(
f"customHandler_caching.states post-cache hit: {customHandler_caching.states}"
)
assert len(customHandler_caching.errors) == 0
assert len(customHandler_caching.states) == 4 # pre, post, success, success
@pytest.mark.asyncio
async def test_async_completion_azure_caching_streaming():
import uuid
litellm.set_verbose = True
customHandler_caching = CompletionCustomHandler()
litellm.cache = Cache(
type="redis",
host=os.environ["REDIS_HOST"],
port=os.environ["REDIS_PORT"],
password=os.environ["REDIS_PASSWORD"],
)
litellm.callbacks = [customHandler_caching]
unique_time = uuid.uuid4()
# Use Router instead of direct litellm.acompletion to get router-specific metadata
model_list = [
{
"model_name": "gpt-4.1-nano",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": os.getenv("AZURE_AI_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_AI_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
]
router = Router(model_list=model_list)
response1 = await router.acompletion(
model="gpt-4.1-nano",
messages=[
{"role": "user", "content": f"Hi 👋 - i'm async azure {unique_time}"}
],
caching=True,
stream=True,
)
async for chunk in response1:
print(f"chunk in response1: {chunk}")
await asyncio.sleep(1)
initial_customhandler_caching_states = len(customHandler_caching.states)
print(f"customHandler_caching.states pre-cache hit: {customHandler_caching.states}")
response2 = await router.acompletion(
model="gpt-4.1-nano",
messages=[
{"role": "user", "content": f"Hi 👋 - i'm async azure {unique_time}"}
],
caching=True,
stream=True,
)
async for chunk in response2:
print(f"chunk in response2: {chunk}")
await asyncio.sleep(1) # success callbacks are done in parallel
print(
f"customHandler_caching.states post-cache hit: {customHandler_caching.states}"
)
assert len(customHandler_caching.errors) == 0
assert (
len(customHandler_caching.states) > initial_customhandler_caching_states
) # pre, post, streaming .., success, success
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=2)
async def test_async_embedding_azure_caching():
print("Testing custom callback input - Azure Caching")
customHandler_caching = CompletionCustomHandler()
litellm.cache = Cache(
type="redis",
host=os.environ["REDIS_HOST"],
port=os.environ["REDIS_PORT"],
password=os.environ["REDIS_PASSWORD"],
)
router = Router(
model_list=[
{
"model_name": "text-embedding-3-small",
"litellm_params": {
"model": "openai/text-embedding-3-small",
},
}
]
)
litellm.callbacks = [customHandler_caching]
unique_time = time.time()
response1 = await router.aembedding(
model="text-embedding-3-small",
input=[f"good morning from litellm1 {unique_time}"],
caching=True,
)
await asyncio.sleep(1) # set cache is async for aembedding()
response2 = await router.aembedding(
model="text-embedding-3-small",
input=[f"good morning from litellm1 {unique_time}"],
caching=True,
)
await asyncio.sleep(1) # success callbacks are done in parallel
print(customHandler_caching.states)
print(customHandler_caching.errors)
assert len(customHandler_caching.errors) == 0
assert len(customHandler_caching.states) == 4 # pre, post, success, success

View file

@ -1,226 +0,0 @@
import asyncio
import gzip
import io
import json
import logging
import os
from datetime import datetime as datetime_class
from unittest.mock import AsyncMock
import pytest
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.datadog.datadog import *
from litellm.types.utils import (
StandardLoggingHiddenParams,
StandardLoggingMetadata,
StandardLoggingModelInformation,
StandardLoggingPayload,
)
verbose_logger.setLevel(logging.DEBUG)
def create_standard_logging_payload() -> StandardLoggingPayload:
return StandardLoggingPayload(
id="test_id",
call_type="completion",
response_cost=0.1,
response_cost_failure_debug_info=None,
status="success",
total_tokens=30,
prompt_tokens=20,
completion_tokens=10,
startTime=1234567890.0,
endTime=1234567891.0,
completionStartTime=1234567890.5,
model_map_information=StandardLoggingModelInformation(
model_map_key="gpt-4.1-mini", model_map_value=None
),
model="gpt-4.1-mini",
model_id="model-123",
model_group="openai-gpt",
api_base="https://api.openai.com",
metadata=StandardLoggingMetadata(
user_api_key_hash="test_hash",
user_api_key_org_id=None,
user_api_key_alias="test_alias",
user_api_key_team_id="test_team",
user_api_key_user_id="test_user",
user_api_key_team_alias="test_team_alias",
spend_logs_metadata=None,
requester_ip_address="127.0.0.1",
requester_metadata=None,
),
cache_hit=False,
cache_key=None,
saved_cache_cost=0.0,
request_tags=[],
end_user=None,
requester_ip_address="127.0.0.1",
messages=[{"role": "user", "content": "Hello, world!"}],
response={"choices": [{"message": {"content": "Hi there!"}}]},
error_str=None,
model_parameters={"stream": True},
hidden_params=StandardLoggingHiddenParams(
model_id="model-123",
cache_key=None,
api_base="https://api.openai.com",
response_cost="0.1",
additional_headers=None,
),
)
@pytest.mark.asyncio
async def test_create_datadog_logging_payload():
"""Test creating a DataDog logging payload from a standard logging object"""
dd_logger = DataDogLogger()
standard_payload = create_standard_logging_payload()
# Create mock kwargs with the standard logging object
kwargs = {"standard_logging_object": standard_payload}
# Test payload creation
dd_payload = dd_logger.create_datadog_logging_payload(
kwargs=kwargs,
response_obj=None,
start_time=datetime_class.now(),
end_time=datetime_class.now(),
)
# Verify payload structure
assert dd_payload["ddsource"] == os.getenv("DD_SOURCE", "litellm")
assert dd_payload["service"] == "litellm-server"
assert dd_payload["status"] == DataDogStatus.INFO
# verify the message field == standard_payload
dict_payload = json.loads(dd_payload["message"])
assert dict_payload == standard_payload
@pytest.mark.asyncio
async def test_datadog_failure_logging():
"""Test logging a failure event to DataDog"""
dd_logger = DataDogLogger()
standard_payload = create_standard_logging_payload()
standard_payload["status"] = "failure" # Set status to failure
standard_payload["error_str"] = "Test error"
kwargs = {"standard_logging_object": standard_payload}
dd_payload = dd_logger.create_datadog_logging_payload(
kwargs=kwargs,
response_obj=None,
start_time=datetime_class.now(),
end_time=datetime_class.now(),
)
assert (
dd_payload["status"] == DataDogStatus.ERROR
) # Verify failure maps to warning status
# verify the message field == standard_payload
dict_payload = json.loads(dd_payload["message"])
assert dict_payload == standard_payload
# verify error_str is in the message field
assert "error_str" in dict_payload
assert dict_payload["error_str"] == "Test error"
@pytest.mark.asyncio
async def test_datadog_log_redis_failures():
"""
Test that poorly configured Redis is logged as Warning on DataDog
"""
try:
from litellm.caching.caching import Cache
from litellm.integrations.datadog.datadog import DataDogLogger
litellm.cache = Cache(
type="redis", host="badhost", port="6379", password="badpassword"
)
os.environ["DD_SITE"] = "https://fake.datadoghq.com"
os.environ["DD_API_KEY"] = "anything"
dd_logger = DataDogLogger()
litellm.callbacks = [dd_logger]
litellm.service_callback = ["datadog"]
litellm.set_verbose = True
# Create a mock for the async_client's post method
mock_post = AsyncMock()
mock_post.return_value.status_code = 202
mock_post.return_value.text = "Accepted"
dd_logger.async_client.post = mock_post
# Make the completion call
for _ in range(3):
response = await litellm.acompletion(
model="gpt-4.1-mini",
messages=[{"role": "user", "content": "what llm are u"}],
max_tokens=10,
temperature=0.2,
mock_response="Accepted",
)
print(response)
# Wait for 5 seconds
await asyncio.sleep(6)
# Assert that the mock was called
assert mock_post.called, "HTTP request was not made"
# Get the arguments of the last call
args, kwargs = mock_post.call_args
print("CAll args and kwargs", args, kwargs)
# For example, checking if the URL is correct
assert kwargs["url"].endswith("/api/v2/logs"), "Incorrect DataDog endpoint"
body = kwargs["data"]
# use gzip to unzip the body
with gzip.open(io.BytesIO(body), "rb") as f:
body = f.read().decode("utf-8")
print(body)
# body is string parse it to dict
body = json.loads(body)
print(body)
failure_events = [log for log in body if log["status"] == "warning"]
assert len(failure_events) > 0, "No failure events logged"
print("ALL FAILURE/WARN EVENTS", failure_events)
for event in failure_events:
message = json.loads(event["message"])
assert (
event["status"] == "warning"
), f"Event status is not 'warning': {event['status']}"
assert (
message["service"] == "redis"
), f"Service is not 'redis': {message['service']}"
assert "error" in message, "No 'error' field in the message"
assert message["error"], "Error field is empty"
except Exception as e:
pytest.fail(f"Test failed with exception: {str(e)}")

View file

@ -1,11 +1,9 @@
import io
import asyncio
import gzip
import json
import logging
import time
from unittest.mock import AsyncMock, patch
import pytest
@ -13,215 +11,7 @@ import pytest
import litellm
from litellm import completion
from litellm._logging import verbose_logger
from litellm.proxy.utils import log_db_metrics, ServiceTypes
from litellm.proxy.db.prisma_client import _PrismaDrainTracker, _TrackedPrismaEngine
from datetime import datetime
from types import SimpleNamespace
import httpx
from prisma.errors import ClientNotConnectedError
async def _run_prisma_query() -> None:
engine = _TrackedPrismaEngine(SimpleNamespace(query=AsyncMock(return_value={})), _PrismaDrainTracker())
await engine.query("{}", tx_id=None)
# Test async function to decorate
@log_db_metrics
async def sample_db_function(*args, **kwargs):
await _run_prisma_query()
return "success"
@log_db_metrics
async def sample_proxy_function(*args, **kwargs):
return "success"
@pytest.mark.asyncio
async def test_log_db_metrics_success():
# Mock the proxy_logging_obj
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
# Setup mock
mock_proxy_logging.service_logging_obj.async_service_success_hook = AsyncMock()
# Call the decorated function
result = await sample_db_function(parent_otel_span="test_span")
# Assertions
assert result == "success"
mock_proxy_logging.service_logging_obj.async_service_success_hook.assert_called_once()
call_args = (
mock_proxy_logging.service_logging_obj.async_service_success_hook.call_args[
1
]
)
assert call_args["service"] == ServiceTypes.DB
assert call_args["call_type"] == "sample_db_function"
assert call_args["parent_otel_span"] == "test_span"
assert isinstance(call_args["duration"], float)
assert isinstance(call_args["start_time"], datetime)
assert isinstance(call_args["end_time"], datetime)
assert call_args["event_metadata"] is None
@pytest.mark.asyncio
async def test_log_db_metrics_event_metadata_is_safe():
"""event_metadata must surface only the table name, never the raw
kwargs/args which carry live clients (Prisma, OTel spans) and secrets.
Regression guard for #28909: a previous version dumped function_kwargs and
function_args onto the span.
"""
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
mock_proxy_logging.service_logging_obj.async_service_success_hook = AsyncMock()
@log_db_metrics
async def db_call(**kwargs):
await _run_prisma_query()
return "success"
await db_call(
parent_otel_span="test_span",
table_name="LiteLLM_SpendLogs",
token="sk-secret-should-not-leak",
prisma_client=object(),
)
await asyncio.sleep(0)
call_args = (
mock_proxy_logging.service_logging_obj.async_service_success_hook.call_args[
1
]
)
assert call_args["event_metadata"] == {"table_name": "LiteLLM_SpendLogs"}
@pytest.mark.asyncio
async def test_log_db_metrics_duration():
# Mock the proxy_logging_obj
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
# Setup mock
mock_proxy_logging.service_logging_obj.async_service_success_hook = AsyncMock()
# Add a delay to the function to test duration
@log_db_metrics
async def delayed_function(**kwargs):
await _run_prisma_query()
await asyncio.sleep(1) # 1 second delay
return "success"
# Call the decorated function
start = time.time()
result = await delayed_function(parent_otel_span="test_span")
end = time.time()
# Get the actual duration
actual_duration = end - start
# Get the logged duration from the mock call
call_args = (
mock_proxy_logging.service_logging_obj.async_service_success_hook.call_args[
1
]
)
logged_duration = call_args["duration"]
# Assert the logged duration is approximately equal to actual duration (within 0.1 seconds)
assert abs(logged_duration - actual_duration) < 0.1
assert result == "success"
@pytest.mark.asyncio
async def test_log_db_metrics_failure():
"""
should log a failure if a prisma error is raised
"""
# Mock the proxy_logging_obj
from prisma.errors import ClientNotConnectedError
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
# Setup mock
mock_proxy_logging.service_logging_obj.async_service_failure_hook = AsyncMock()
# Create a failing function
@log_db_metrics
async def failing_function(**kwargs):
raise ClientNotConnectedError()
# Call the decorated function and expect it to raise
with pytest.raises(ClientNotConnectedError) as exc_info:
await failing_function(parent_otel_span="test_span")
# Assertions
assert "Client is not connected to the query engine" in str(exc_info.value)
mock_proxy_logging.service_logging_obj.async_service_failure_hook.assert_called_once()
call_args = (
mock_proxy_logging.service_logging_obj.async_service_failure_hook.call_args[
1
]
)
assert call_args["service"] == ServiceTypes.DB
assert call_args["call_type"] == "failing_function"
assert call_args["parent_otel_span"] == "test_span"
assert isinstance(call_args["duration"], float)
assert isinstance(call_args["start_time"], datetime)
assert isinstance(call_args["end_time"], datetime)
assert isinstance(call_args["error"], ClientNotConnectedError)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"exception,should_log",
[
(ValueError("Generic error"), False),
(KeyError("Missing key"), False),
(TypeError("Type error"), False),
(httpx.ConnectError("Failed to connect"), True),
(httpx.TimeoutException("Request timed out"), True),
(ClientNotConnectedError(), True), # Prisma error
],
)
async def test_log_db_metrics_failure_error_types(exception, should_log):
"""
Why Test?
Users were seeing that non-DB errors were being logged as DB Service Failures
Example a failure to read a value from cache was being logged as a DB Service Failure
Parameterized test to verify:
- DB-related errors (Prisma, httpx) are logged as service failures
- Non-DB errors (ValueError, KeyError, etc.) are not logged
"""
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
mock_proxy_logging.service_logging_obj.async_service_failure_hook = AsyncMock()
@log_db_metrics
async def failing_function(**kwargs):
raise exception
# Call the function and expect it to raise the exception
with pytest.raises(type(exception)):
await failing_function(parent_otel_span="test_span")
if should_log:
# Assert failure was logged for DB-related errors
mock_proxy_logging.service_logging_obj.async_service_failure_hook.assert_called_once()
call_args = mock_proxy_logging.service_logging_obj.async_service_failure_hook.call_args[
1
]
assert call_args["service"] == ServiceTypes.DB
assert call_args["call_type"] == "failing_function"
assert call_args["parent_otel_span"] == "test_span"
assert isinstance(call_args["duration"], float)
assert isinstance(call_args["start_time"], datetime)
assert isinstance(call_args["end_time"], datetime)
assert isinstance(call_args["error"], type(exception))
else:
# Assert failure was NOT logged for non-DB errors
mock_proxy_logging.service_logging_obj.async_service_failure_hook.assert_not_called()
from litellm.proxy.utils import ServiceTypes
@pytest.mark.asyncio

View file

@ -1,100 +0,0 @@
import traceback
from litellm._uuid import uuid
import pytest
from dotenv import load_dotenv
from fastapi import Request
from fastapi.routing import APIRoute
load_dotenv()
import io
import time
import json
import litellm
from litellm.router import Router
import asyncio
from typing import Optional
from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase
from litellm.integrations.custom_logger import CustomLogger
class TestCustomLogger(CustomLogger):
def __init__(self):
self.recorded_usage: Optional[Usage] = None
self.standard_logging_payload: Optional[StandardLoggingPayload] = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
standard_logging_payload = kwargs.get("standard_logging_object")
self.standard_logging_payload = standard_logging_payload
print(
"standard_logging_payload",
json.dumps(standard_logging_payload, indent=4, default=str),
)
pass
@pytest.mark.asyncio
@pytest.mark.parametrize(
"model", [None, "omni-moderation-latest", "router-internal-moderation-model"]
)
async def test_moderations_api_logging(model):
"""
When moderations API is called, it should log the event on standard_logging_payload
"""
custom_logger = TestCustomLogger()
litellm.logging_callback_manager.add_litellm_callback(custom_logger)
MODEL_GROUP = "internal-moderation-model"
router = Router(
model_list=[
{
"model_name": MODEL_GROUP,
"litellm_params": {
"model": "openai/omni-moderation-latest",
},
}
]
)
input_content = "Hello, how are you?"
if model == "router-internal-moderation-model":
response = await router.amoderation(
input=input_content,
model=MODEL_GROUP,
)
else:
response = await litellm.amoderation(
input=input_content,
model=model,
)
print("response", json.dumps(response, indent=4, default=str))
await asyncio.sleep(2)
assert custom_logger.standard_logging_payload is not None
# validate the standard_logging_payload
standard_logging_payload: StandardLoggingPayload = (
custom_logger.standard_logging_payload
)
assert (
standard_logging_payload["call_type"]
== litellm.utils.CallTypes.amoderation.value
)
assert standard_logging_payload["status"] == "success"
assert (
standard_logging_payload["custom_llm_provider"]
== litellm.LlmProviders.OPENAI.value
)
# assert the logged input == input
assert standard_logging_payload["messages"][0]["content"] == input_content
# assert the logged response == response user received client side
assert dict(standard_logging_payload["response"]) == response.model_dump()
# if router used, validate model_group is logged as expected
if model == "router-internal-moderation-model":
assert standard_logging_payload["model_group"] == MODEL_GROUP

View file

@ -1,133 +0,0 @@
import pytest
import litellm
import asyncio
import logging
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from litellm._logging import verbose_logger
from litellm.integrations.opentelemetry import (
OpenTelemetry,
OpenTelemetryConfig,
)
verbose_logger.setLevel(logging.DEBUG)
EXPECTED_SPAN_NAMES = ["litellm_request", "raw_gen_ai_request"]
exporter = InMemorySpanExporter()
@pytest.mark.asyncio
@pytest.mark.parametrize("streaming", [True, False])
async def test_async_otel_callback(streaming):
litellm.set_verbose = True
# Clear exporter at the start to ensure clean state
exporter.clear()
litellm.callbacks = [OpenTelemetry(config=OpenTelemetryConfig(exporter=exporter))]
response = await litellm.acompletion(
model="gpt-4.1-mini",
messages=[{"role": "user", "content": "hi"}],
temperature=0.1,
user="OTEL_USER",
stream=streaming,
)
if streaming is True:
async for chunk in response:
print("chunk", chunk)
await asyncio.sleep(4)
spans = exporter.get_finished_spans()
print("spans", spans)
assert len(spans) == 2
_span_names = [span.name for span in spans]
print("recorded span names", _span_names)
assert set(_span_names) == set(EXPECTED_SPAN_NAMES)
# print the value of a span
for span in spans:
print("span name", span.name)
print("span attributes", span.attributes)
if span.name == "litellm_request":
validate_litellm_request(span)
# Additional specific checks
assert span._attributes["gen_ai.request.model"] == "gpt-4.1-mini"
assert span._attributes["gen_ai.system"] == "openai"
assert span._attributes["gen_ai.request.temperature"] == 0.1
assert span._attributes["llm.is_streaming"] == str(streaming)
assert span._attributes["llm.user"] == "OTEL_USER"
elif span.name == "raw_gen_ai_request":
if streaming is True:
validate_raw_gen_ai_request_openai_streaming(span)
else:
validate_raw_gen_ai_request_openai_non_streaming(span)
# clear in memory exporter
exporter.clear()
def validate_litellm_request(span):
expected_attributes = [
"gen_ai.request.model",
"gen_ai.system",
"gen_ai.request.temperature",
"llm.is_streaming",
"llm.user",
"gen_ai.response.id",
"gen_ai.response.model",
"gen_ai.usage.total_tokens",
"gen_ai.usage.output_tokens",
"gen_ai.usage.input_tokens",
]
# get the str of all the span attributes
print("span attributes", span._attributes)
for attr in expected_attributes:
value = span._attributes[attr]
print("value", value)
assert value is not None, f"Attribute {attr} has None value"
def validate_raw_gen_ai_request_openai_non_streaming(span):
expected_attributes = [
"llm.openai.messages",
"llm.openai.temperature",
"llm.openai.user",
"llm.openai.extra_body",
"llm.openai.id",
"llm.openai.choices",
"llm.openai.created",
"llm.openai.model",
"llm.openai.object",
"llm.openai.service_tier",
"llm.openai.system_fingerprint",
"llm.openai.usage",
]
print("span attributes", span._attributes)
for attr in span._attributes:
print(attr)
for attr in expected_attributes:
assert span._attributes[attr] is not None, f"Attribute {attr} has None"
def validate_raw_gen_ai_request_openai_streaming(span):
expected_attributes = [
"llm.openai.messages",
"llm.openai.temperature",
"llm.openai.user",
"llm.openai.extra_body",
"llm.openai.model",
]
print("span attributes", span._attributes)
for attr in span._attributes:
print(attr)
for attr in expected_attributes:
assert span._attributes[attr] is not None, f"Attribute {attr} has None"

View file

@ -1,157 +0,0 @@
import traceback
from litellm._uuid import uuid
import pytest
from dotenv import load_dotenv
from fastapi import Request
from fastapi.routing import APIRoute
load_dotenv()
import io
import time
import json
# this file is to test litellm/proxy
import litellm
import asyncio
from typing import Optional
from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase
from litellm.integrations.custom_logger import CustomLogger
class TestCustomLogger(CustomLogger):
def __init__(self):
self.recorded_usage: Optional[Usage] = None
self.standard_logging_payload: Optional[StandardLoggingPayload] = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
standard_logging_payload = kwargs.get("standard_logging_object")
self.standard_logging_payload = standard_logging_payload
print(
"standard_logging_payload",
json.dumps(standard_logging_payload, indent=4, default=str),
)
self.recorded_usage = Usage(
prompt_tokens=standard_logging_payload.get("prompt_tokens"),
completion_tokens=standard_logging_payload.get("completion_tokens"),
total_tokens=standard_logging_payload.get("total_tokens"),
)
pass
@pytest.mark.asyncio
async def test_stream_token_counting_gpt_4o():
"""
When stream_options={"include_usage": True} logging callback tracks Usage == Usage from llm API
"""
custom_logger = TestCustomLogger()
litellm.logging_callback_manager.add_litellm_callback(custom_logger)
response = await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "Hello, how are you?" * 100}],
stream=True,
stream_options={"include_usage": True},
)
actual_usage = None
async for chunk in response:
if "usage" in chunk:
actual_usage = chunk["usage"]
print("chunk.usage", json.dumps(chunk["usage"], indent=4, default=str))
pass
await asyncio.sleep(2)
print("\n\n\n\n\n")
print(
"recorded_usage",
json.dumps(custom_logger.recorded_usage, indent=4, default=str),
)
print("\n\n\n\n\n")
assert actual_usage.prompt_tokens == custom_logger.recorded_usage.prompt_tokens
assert (
actual_usage.completion_tokens == custom_logger.recorded_usage.completion_tokens
)
assert actual_usage.total_tokens == custom_logger.recorded_usage.total_tokens
@pytest.mark.asyncio
async def test_stream_token_counting_without_include_usage():
"""
When stream_options={"include_usage": True} is not passed, the usage tracked == usage from llm api chunk
by default, litellm passes `include_usage=True` for OpenAI API
"""
custom_logger = TestCustomLogger()
litellm.logging_callback_manager.add_litellm_callback(custom_logger)
response = await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "Hello, how are you?" * 100}],
stream=True,
)
actual_usage = None
async for chunk in response:
if "usage" in chunk:
actual_usage = chunk["usage"]
print("chunk.usage", json.dumps(chunk["usage"], indent=4, default=str))
pass
await asyncio.sleep(2)
print("\n\n\n\n\n")
print(
"recorded_usage",
json.dumps(custom_logger.recorded_usage, indent=4, default=str),
)
print("\n\n\n\n\n")
assert actual_usage.prompt_tokens == custom_logger.recorded_usage.prompt_tokens
assert (
actual_usage.completion_tokens == custom_logger.recorded_usage.completion_tokens
)
assert actual_usage.total_tokens == custom_logger.recorded_usage.total_tokens
@pytest.mark.asyncio
async def test_stream_token_counting_with_redaction():
"""
When litellm.turn_off_message_logging=True is used, the usage tracked == usage from llm api chunk
"""
litellm.turn_off_message_logging = True
custom_logger = TestCustomLogger()
litellm.logging_callback_manager.add_litellm_callback(custom_logger)
response = await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "Hello, how are you?" * 100}],
stream=True,
)
actual_usage = None
async for chunk in response:
if "usage" in chunk:
actual_usage = chunk["usage"]
print("chunk.usage", json.dumps(chunk["usage"], indent=4, default=str))
pass
await asyncio.sleep(2)
print("\n\n\n\n\n")
print(
"recorded_usage",
json.dumps(custom_logger.recorded_usage, indent=4, default=str),
)
print("\n\n\n\n\n")
assert actual_usage.prompt_tokens == custom_logger.recorded_usage.prompt_tokens
assert (
actual_usage.completion_tokens == custom_logger.recorded_usage.completion_tokens
)
assert actual_usage.total_tokens == custom_logger.recorded_usage.total_tokens

View file

@ -1,557 +0,0 @@
import os
import asyncio
import json
import secrets
import uuid
from typing import Any, Optional
import aiohttp
import openai
import pytest
from httpx import AsyncClient
PROXY_BASE = "http://0.0.0.0:4000"
MASTER_HEADERS = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
CLI_SSO_MODEL = "fake-openai-endpoint"
async def make_calls_until_budget_exceeded(session, key: str, call_function, **kwargs):
"""Helper function to make API calls until budget is exceeded. Verify that the budget is exceeded error is returned."""
MAX_CALLS = 200
call_count = 0
try:
while call_count < MAX_CALLS:
await call_function(session=session, key=key, **kwargs)
call_count += 1
await asyncio.sleep(0.1) # allow spend tracking to catch up
pytest.fail(f"Budget was not exceeded after {MAX_CALLS} calls")
except openai.APIStatusError as e:
print("vars: ", vars(e))
print("e.body: ", e.body)
error_dict = e.body
print("error_dict: ", error_dict)
# Check error structure and values that should be consistent
assert (
error_dict["code"] == "422"
), f"Expected error code 422, got: {error_dict['code']}"
assert (
error_dict["type"] == "budget_exceeded"
), f"Expected error type budget_exceeded, got: {error_dict['type']}"
# Check message contains required parts without checking specific values
message = error_dict["message"]
assert (
"Budget has been exceeded!" in message
), f"Expected message to start with 'Budget has been exceeded!', got: {message}"
assert (
"Current cost:" in message
), f"Expected message to contain 'Current cost:', got: {message}"
assert (
"Max budget:" in message
), f"Expected message to contain 'Max budget:', got: {message}"
return call_count
async def generate_key(
session,
max_budget=None,
):
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {
"max_budget": max_budget,
}
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
async def chat_completion(session, key: str, model: str):
"""Make a chat completion request using OpenAI SDK"""
from openai import AsyncOpenAI
from litellm._uuid import uuid
client = AsyncOpenAI(
api_key=key, base_url="http://0.0.0.0:4000/v1" # Point to our local proxy
)
response = await client.chat.completions.create(
model=model,
messages=[{"role": "user", "content": f"Say hello! {uuid.uuid4()}" * 100}],
)
return response
@pytest.mark.asyncio
async def test_chat_completion_low_budget():
"""
Test budget enforcement for chat completions:
1. Create key with $0.01 budget
2. Make chat completion calls until budget exceeded
3. Verify budget exceeded error
"""
async with aiohttp.ClientSession() as session:
# Create key with $0.01 budget
key_gen = await generate_key(session=session, max_budget=0.0000000005)
print("response from key generation: ", key_gen)
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"
@pytest.mark.asyncio
async def test_chat_completion_zero_budget():
"""
Test budget enforcement for chat completions:
1. Create key with $0.01 budget
2. Make chat completion calls until budget exceeded
3. Verify budget exceeded error
"""
async with aiohttp.ClientSession() as session:
# Create key with $0.01 budget
key_gen = await generate_key(session=session, max_budget=0.000000000)
print("response from key generation: ", key_gen)
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 no calls before budget exceeded"
@pytest.mark.asyncio
async def test_chat_completion_high_budget():
"""
Test budget enforcement for chat completions:
1. Create key with $0.01 budget
2. Make chat completion calls until budget exceeded
3. Verify budget exceeded error
"""
async with aiohttp.ClientSession() as session:
# Create key with $0.01 budget
key_gen = await generate_key(session=session, max_budget=0.001)
print("response from key generation: ", key_gen)
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"
@pytest.mark.parametrize(
"field",
[
"max_budget",
"rpm_limit",
"tpm_limit",
],
)
@pytest.mark.asyncio
async def test_key_limit_modifications(field):
# Create initial key
client = AsyncClient(base_url="http://0.0.0.0:4000")
key_data = {"max_budget": None, "rpm_limit": None, "tpm_limit": None}
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}
response = await client.post("/key/generate", json=key_data, headers=headers)
assert response.status_code == 200
generate_key_response = response.json()
print("generate_key_response: ", json.dumps(generate_key_response, indent=4))
key_id = generate_key_response["key"]
# Update key with any non-null value for the field
update_data = {"key": key_id}
update_data[field] = 10 # Any non-null value works
print("update_data: ", json.dumps(update_data, indent=4))
response = await client.post(f"/key/update", json=update_data, headers=headers)
assert response.status_code == 200
assert response.json()[field] is not None
# Reset limit to null
print(f"resetting {field} to null")
update_data[field] = None
response = await client.post(f"/key/update", json=update_data, headers=headers)
print("response: ", json.dumps(response.json(), indent=4))
assert response.status_code == 200
assert response.json()[field] is None
@pytest.mark.parametrize(
"field",
[
"max_budget",
],
)
@pytest.mark.asyncio
async def test_team_limit_modifications(field):
# Create initial team
client = AsyncClient(base_url="http://0.0.0.0:4000")
team_data = {"max_budget": None, "rpm_limit": None, "tpm_limit": None}
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}
response = await client.post("/team/new", json=team_data, headers=headers)
print("response: ", json.dumps(response.json(), indent=4))
assert response.status_code == 200
team_id = response.json()["team_id"]
# Update team with any non-null value for the field
update_data = {"team_id": team_id}
update_data[field] = 10 # Any non-null value works
response = await client.post(f"/team/update", json=update_data, headers=headers)
print("response: ", json.dumps(response.json(), indent=4))
assert response.status_code == 200
assert response.json()["data"][field] is not None
# Reset limit to null
print(f"resetting {field} to null")
update_data[field] = None
response = await client.post(f"/team/update", json=update_data, headers=headers)
print("response: ", json.dumps(response.json(), indent=4))
assert response.status_code == 200
assert response.json()["data"][field] is None
async def generate_team_key(
session,
team_id: str,
max_budget: Optional[float] = None,
):
"""Helper function to generate a key for a specific team"""
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data: dict[str, Any] = {"team_id": team_id}
if max_budget is not None:
data["max_budget"] = max_budget
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
async def create_team(
session,
max_budget=None,
models: Optional[list[str]] = None,
team_alias: Optional[str] = None,
):
"""Helper function to create a new team"""
url = f"{PROXY_BASE}/team/new"
data: dict[str, Any] = {"max_budget": max_budget}
if models is not None:
data["models"] = models
if team_alias is not None:
data["team_alias"] = team_alias
async with session.post(url, headers=MASTER_HEADERS, json=data) as response:
return await response.json()
async def create_user(
session,
*,
user_id: str,
user_email: str,
teams: list[str],
models: list[str],
):
url = f"{PROXY_BASE}/user/new"
data = {
"user_id": user_id,
"user_email": user_email,
"teams": teams,
"models": models,
"auto_create_key": False,
}
async with session.post(url, headers=MASTER_HEADERS, json=data) as response:
return await response.json()
async def add_team_member(
session,
*,
team_id: str,
user_id: str,
user_email: str,
):
url = f"{PROXY_BASE}/team/member_add"
data = {
"team_id": team_id,
"member": [{"user_id": user_id, "user_email": user_email, "role": "user"}],
}
async with session.post(url, headers=MASTER_HEADERS, json=data) as response:
return await response.json()
async def obtain_cli_sso_token_via_poll_flow(
session,
*,
user_id: str,
user_email: str,
team_id: str,
team_alias: str,
models: list[str],
) -> str:
"""
Obtain a CLI SSO JWT through the same HTTP flow as `lite login`:
/sso/cli/start -> (SSO callback) -> /sso/cli/complete -> /sso/cli/poll.
When the proxy SSO session cache is not shared with the test runner (otel CI
uses an isolated in-container cache), falls back to minting the identical JWT
that /sso/cli/poll would return.
"""
async with session.post(f"{PROXY_BASE}/sso/cli/start") as resp:
resp.raise_for_status()
start = await resp.json()
login_id = start["login_id"]
poll_secret = start["poll_secret"]
user_code = start["user_code"]
browser_complete_token = secrets.token_urlsafe(32)
seeded = await _seed_cli_sso_flow_in_shared_redis(
login_id=login_id,
user_id=user_id,
user_email=user_email,
team_id=team_id,
team_alias=team_alias,
models=models,
browser_complete_token=browser_complete_token,
)
if not seeded:
pytest.skip("Shared Redis not available; skipping full poll-flow test")
async with session.post(
f"{PROXY_BASE}/sso/cli/complete/{login_id}",
data={
"user_code": user_code,
"browser_complete_token": browser_complete_token,
},
headers={"Content-Type": "application/x-www-form-urlencoded"},
) as resp:
assert resp.status == 200, await resp.text()
poll_headers = {
"x-litellm-cli-poll-secret": poll_secret,
}
async with session.get(
f"{PROXY_BASE}/sso/cli/poll/{login_id}",
params={"team_id": team_id},
headers=poll_headers,
) as resp:
poll = await resp.json()
assert poll.get("status") == "ready", poll
assert "key" in poll, poll
return poll["key"]
async def _seed_cli_sso_flow_in_shared_redis(
*,
login_id: str,
user_id: str,
user_email: str,
team_id: str,
team_alias: str,
models: list[str],
browser_complete_token: str,
) -> bool:
"""Seed the CLI SSO flow in Redis when tests share the proxy's Redis instance."""
import ast
import json
import os
try:
import redis
except ImportError:
return False
host = os.getenv("REDIS_HOST")
if not host:
return False
try:
client = redis.Redis(
host=host,
port=int(os.getenv("REDIS_PORT", "6379")),
password=os.getenv("REDIS_PASSWORD") or None,
decode_responses=True,
)
client.ping()
except Exception:
return False
from litellm.proxy.management_endpoints.ui_sso import (
_get_cli_sso_flow_cache_key,
_hash_cli_sso_secret,
)
cache_key = _get_cli_sso_flow_cache_key(login_id)
raw_flow = client.get(cache_key)
if raw_flow is None:
return False
try:
flow = ast.literal_eval(raw_flow)
except (SyntaxError, ValueError):
return False
if not isinstance(flow, dict):
return False
updated_flow = {
**flow,
"sso_complete": True,
"user_code_verified": False,
"session_data": {
"user_id": user_id,
"user_role": "internal_user",
"models": models,
"user_email": user_email,
"teams": [team_id],
"team_details": [{"team_id": team_id, "team_alias": team_alias}],
},
"browser_complete_token_hash": _hash_cli_sso_secret(browser_complete_token),
}
client.setex(cache_key, 600, json.dumps(updated_flow))
return True
async def make_calls_until_team_budget_exceeded_cli_sso(
session,
token: str,
team_id: str,
model: str,
):
"""Like make_calls_until_budget_exceeded but asserts team budget blocked the CLI SSO token."""
MAX_CALLS = 200
call_count = 0
try:
while call_count < MAX_CALLS:
await chat_completion(session=session, key=token, model=model)
call_count += 1
await asyncio.sleep(0.1)
pytest.fail(f"Budget was not exceeded after {MAX_CALLS} calls")
except openai.APIStatusError as e:
error_dict = e.body
assert error_dict["code"] == "422"
assert error_dict["type"] == "budget_exceeded"
message = error_dict["message"]
assert "Budget has been exceeded!" in message
assert "Team=" in message, f"Expected team budget error, got: {message}"
assert team_id in message, f"Expected team id in error, got: {message}"
return call_count
@pytest.mark.asyncio
async def test_team_budget_enforcement():
"""
Test budget enforcement for team-wide budgets:
1. Create team with low budget
2. Create key for that team
3. Make calls until team budget exceeded
4. Verify budget exceeded error
"""
async with aiohttp.ClientSession() as session:
# Create team with 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"
@pytest.mark.asyncio
async def test_team_budget_enforcement_cli_sso_token():
"""
Team budget enforcement for CLI SSO session tokens (lite login JWT).
1. Create team with a tiny max_budget and a user on that team
2. Obtain a CLI SSO JWT (HTTP poll flow when Redis is shared, else mint)
3. Make chat completion calls until the team budget is exceeded
4. Verify HTTP 422 budget_exceeded names the team
"""
user_id = f"cli-budget-user-{uuid.uuid4().hex[:8]}"
user_email = f"{user_id}@example.com"
team_alias = f"cli-budget-team-{uuid.uuid4().hex[:8]}"
async with aiohttp.ClientSession() as session:
team_response = await create_team(
session=session,
max_budget=0.0000000005,
models=[CLI_SSO_MODEL],
team_alias=team_alias,
)
team_id = team_response["team_id"]
await create_user(
session,
user_id=user_id,
user_email=user_email,
teams=[team_id],
models=[CLI_SSO_MODEL],
)
await add_team_member(
session,
team_id=team_id,
user_id=user_id,
user_email=user_email,
)
cli_token = await obtain_cli_sso_token_via_poll_flow(
session,
user_id=user_id,
user_email=user_email,
team_id=team_id,
team_alias=team_alias,
models=[CLI_SSO_MODEL],
)
assert not cli_token.startswith(
"sk-"
), "CLI SSO token must not be a virtual key"
calls_made = await make_calls_until_team_budget_exceeded_cli_sso(
session=session,
token=cli_token,
team_id=team_id,
model=CLI_SSO_MODEL,
)
assert (
calls_made > 0
), "Should make at least one successful call before team budget exceeded"
# Verify it was the team budget that was exceeded

View file

@ -1,304 +0,0 @@
import os
import pytest
import asyncio
import aiohttp
import json
from httpx import AsyncClient
from openai import PermissionDeniedError
from typing import Any, Optional, List, Literal
# The proxy strips client-supplied `mock_response` unless the calling key or
# team has this admin-metadata flag set. See `_UNTRUSTED_ROOT_CONTROL_FIELDS`
# in litellm/proxy/litellm_pre_call_utils.py.
_ALLOW_CLIENT_MOCK_METADATA = {"allow_client_mock_response": True}
async def generate_key(
session, models: Optional[List[str]] = None, team_id: Optional[str] = None
):
"""Helper function to generate a key with specific model access controls"""
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data: dict = {"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA)}
if models is not None:
data["models"] = models
if team_id is not None:
data["team_id"] = team_id
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
async def generate_team(session, models: Optional[List[str]] = None):
"""Helper function to generate a team with specific model access"""
url = "http://0.0.0.0:4000/team/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data: dict = {"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA)}
if models is not None:
data["models"] = models
async with session.post(url, headers=headers, json=data) as response:
return await response.json()
async def mock_chat_completion(session, key: str, model: str):
"""Make a chat completion request using OpenAI SDK"""
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=model,
messages=[{"role": "user", "content": f"Say hello! {uuid.uuid4()}"}],
extra_body={
"mock_response": "mock_response",
},
)
return response
@pytest.mark.parametrize(
"key_models, test_model, expect_success",
[
(["openai/*"], "anthropic/claude-2", False), # Non-matching model
(["gpt-5.5"], "gpt-5.5", True), # Exact model match
(["bedrock/*"], "bedrock/anthropic.claude-3", True), # Bedrock wildcard
(["bedrock/anthropic.*"], "bedrock/anthropic.claude-3", True), # Pattern match
(["bedrock/anthropic.*"], "bedrock/amazon.titan", False), # Pattern non-match
(None, "gpt-5.5", True), # No model restrictions
([], "gpt-5.5", True), # Empty model list
],
)
@pytest.mark.asyncio
async def test_model_access_patterns(key_models, test_model, expect_success):
"""
Test model access patterns for API keys:
1. Create key with specific model access pattern
2. Attempt to make completion with test model
3. Verify access is granted/denied as expected
"""
async with aiohttp.ClientSession() as session:
# Generate key with specified model access
key_gen = await generate_key(session=session, models=key_models)
key = key_gen["key"]
try:
response = await mock_chat_completion(
session=session,
key=key,
model=test_model,
)
if not expect_success:
pytest.fail(f"Expected request to fail for model {test_model}")
assert (
response is not None
), "Should get valid response when access is allowed"
except Exception as e:
if expect_success:
pytest.fail(f"Expected request to succeed but got error: {e}")
_error_body = e.body
# Assert error structure and values
assert _error_body["type"] == "key_model_access_denied"
assert _error_body["param"] == "model"
assert _error_body["code"] == "403"
assert "is not available for this API key" in _error_body["message"]
@pytest.mark.asyncio
async def test_model_access_update():
"""
Test updating model access for an existing key:
1. Create key with restricted model access
2. Verify access patterns
3. Update key with new model access
4. Verify new access patterns
"""
client = AsyncClient(base_url="http://0.0.0.0:4000")
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}
# Create initial key with restricted access
response = await client.post(
"/key/generate",
json={
"models": ["openai/gpt-5.5"],
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
key_data = response.json()
key = key_data["key"]
# Test initial access
async with aiohttp.ClientSession() as session:
# Should work with gpt-5.5
await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5")
# Should fail with gpt-5-mini
with pytest.raises(PermissionDeniedError) as exc_info:
await mock_chat_completion(
session=session, key=key, model="openai/gpt-5-mini"
)
_validate_model_access_exception(
exc_info.value, expected_type="key_model_access_denied"
)
# Update key with new model access
response = await client.post(
"/key/update", json={"key": key, "models": ["openai/*"]}, headers=headers
)
assert response.status_code == 200
# Test updated access
async with aiohttp.ClientSession() as session:
# Both models should now work
await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5")
await mock_chat_completion(
session=session, key=key, model="openai/gpt-5-mini"
)
# Non-OpenAI model should still fail
with pytest.raises(PermissionDeniedError) as exc_info:
await mock_chat_completion(
session=session, key=key, model="anthropic/claude-2"
)
_validate_model_access_exception(
exc_info.value, expected_type="key_model_access_denied"
)
@pytest.mark.parametrize(
"team_models, test_model, expect_success",
[
(["openai/*"], "anthropic/claude-2", False), # Non-matching model
],
)
@pytest.mark.asyncio
async def test_team_model_access_patterns(team_models, test_model, expect_success):
"""
Test model access patterns for team-based API keys:
1. Create team with specific model access pattern
2. Generate key for that team
3. Attempt to make completion with test model
4. Verify access is granted/denied as expected
"""
client = AsyncClient(base_url="http://0.0.0.0:4000")
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}
async with aiohttp.ClientSession() as session:
try:
team_gen = await generate_team(session=session, models=team_models)
print("created team", team_gen)
team_id = team_gen["team_id"]
key_gen = await generate_key(session=session, team_id=team_id)
print("created key", key_gen)
key = key_gen["key"]
response = await mock_chat_completion(
session=session,
key=key,
model=test_model,
)
if not expect_success:
pytest.fail(f"Expected request to fail for model {test_model}")
assert (
response is not None
), "Should get valid response when access is allowed"
except Exception as e:
if expect_success:
pytest.fail(f"Expected request to succeed but got error: {e}")
_validate_model_access_exception(
e, expected_type="team_model_access_denied"
)
@pytest.mark.asyncio
async def test_team_model_access_update():
"""
Test updating model access for a team:
1. Create team with restricted model access
2. Verify access patterns
3. Update team with new model access
4. Verify new access patterns
"""
client = AsyncClient(base_url="http://0.0.0.0:4000")
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}
# Create initial team with restricted access
response = await client.post(
"/team/new",
json={
"models": ["openai/gpt-5.5"],
"name": "test-team",
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
team_data = response.json()
team_id = team_data["team_id"]
# Generate a key for this team
response = await client.post(
"/key/generate",
json={
"team_id": team_id,
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
key = response.json()["key"]
# Test initial access
async with aiohttp.ClientSession() as session:
# Should work with gpt-5.5
await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5")
# Should fail with gpt-5-mini
with pytest.raises(PermissionDeniedError) as exc_info:
await mock_chat_completion(
session=session, key=key, model="openai/gpt-5-mini"
)
_validate_model_access_exception(
exc_info.value, expected_type="team_model_access_denied"
)
# Update team with new model access
response = await client.post(
"/team/update",
json={"team_id": team_id, "models": ["openai/*"]},
headers=headers,
)
assert response.status_code == 200
# Test updated access
async with aiohttp.ClientSession() as session:
# Both models should now work
await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5")
await mock_chat_completion(
session=session, key=key, model="openai/gpt-5-mini"
)
# Non-OpenAI model should still fail
with pytest.raises(PermissionDeniedError) as exc_info:
await mock_chat_completion(
session=session, key=key, model="anthropic/claude-2"
)
_validate_model_access_exception(
exc_info.value, expected_type="team_model_access_denied"
)
def _validate_model_access_exception(
e: Exception,
expected_type: Literal["key_model_access_denied", "team_model_access_denied"],
):
_error_body = e.body
# Assert error structure and values
assert _error_body["type"] == expected_type
assert _error_body["param"] == "model"
assert _error_body["code"] == "403"
assert "is not available for this API key" in _error_body["message"]
assert "not allowed to access model" not in _error_body["message"]

View file

@ -48,98 +48,6 @@ async def chat_completion(
return await response.json(), response_headers
async def generate_key(
session, guardrails: Optional[List] = None, team_id: Optional[str] = None
):
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {}
if guardrails:
data["guardrails"] = guardrails
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()
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_no_llm_guard_triggered():
"""
- Tests a request where no content mod is triggered
- Assert that the guardrails applied are returned in the response headers
"""
async with aiohttp.ClientSession() as session:
response, headers = await chat_completion(
session,
os.environ["LITELLM_MASTER_KEY"],
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello what's the weather"}],
guardrails=[],
)
await asyncio.sleep(3)
print("response=", response, "response headers", headers)
assert "x-litellm-applied-guardrails" not in headers
@pytest.mark.asyncio
async def test_guardrails_with_api_key_controls():
"""
- Make two API Keys
- Key 1 with no guardrails
- Key 2 with guardrails
- Request to Key 1 -> should be success with no guardrails
- Request to Key 2 -> should be error since guardrails are triggered
"""
async with aiohttp.ClientSession() as session:
key_with_guardrails = await generate_key(
session=session,
guardrails=[
"bedrock-pre-guard",
],
)
key_with_guardrails = key_with_guardrails["key"]
key_without_guardrails = await generate_key(session=session, guardrails=None)
key_without_guardrails = key_without_guardrails["key"]
# test no guardrails triggered for key without guardrails
response, headers = await chat_completion(
session,
key_without_guardrails,
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello what's the weather"}],
)
await asyncio.sleep(3)
print("response=", response, "response headers", headers)
assert "x-litellm-applied-guardrails" not in headers
# test guardrails triggered for key with guardrails
response, headers = await chat_completion(
session,
key_with_guardrails,
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello my name is ishaan@berri.ai"}],
)
assert "x-litellm-applied-guardrails" in headers
assert headers["x-litellm-applied-guardrails"] == "bedrock-pre-guard"
@pytest.mark.asyncio
async def test_bedrock_guardrail_triggered():
"""
@ -160,105 +68,6 @@ async def test_bedrock_guardrail_triggered():
assert "Violated guardrail policy" in str(e)
@pytest.mark.asyncio
async def test_custom_guardrail_during_call_triggered():
"""
- Tests a request where our bedrock guardrail should be triggered
- Assert that the guardrails applied are returned in the response headers
"""
async with aiohttp.ClientSession() as session:
with pytest.raises(Exception, match="Guardrail failed words - `litellm` detected") as exc_info:
response, headers = await chat_completion(
session,
os.environ["LITELLM_MASTER_KEY"],
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello do you like litellm?"}],
guardrails=["custom-during-guard"],
)
e = exc_info.value
print(e)
assert "Guardrail failed words - `litellm` detected" in str(e)
async def create_team(session, guardrails: Optional[List] = None):
url = "http://0.0.0.0:4000/team/new"
headers = {"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}", "Content-Type": "application/json"}
data = {"guardrails": guardrails}
print("request data=", data)
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
print(response_text)
print()
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
return await response.json()
@pytest.mark.asyncio
async def test_guardrails_with_team_controls():
"""
- Create a team with guardrails
- Make two API Keys
- Key 1 not associated with team
- Key 2 associated with team (inherits team guardrails)
- Request with Key 1 -> should be success with no guardrails
- Request with Key 2 -> should error since team guardrails are triggered
"""
async with aiohttp.ClientSession() as session:
# Create team with guardrails
team = await create_team(
session=session,
guardrails=[
"bedrock-pre-guard",
],
)
print("team=", team)
team_id = team["team_id"]
# Create key with team association
key_with_team = await generate_key(session=session, team_id=team_id)
key_with_team = key_with_team["key"]
# Create key without team
key_without_team = await generate_key(
session=session,
)
key_without_team = key_without_team["key"]
# Test no guardrails triggered for key without a team
response, headers = await chat_completion(
session,
key_without_team,
model="fake-openai-endpoint",
messages=[{"role": "user", "content": "Hello my name is ishaan@berri.ai"}],
)
await asyncio.sleep(3)
print("response=", response, "response headers", headers)
assert "x-litellm-applied-guardrails" not in headers
response, headers = await chat_completion(
session,
key_with_team,
model="fake-openai-endpoint",
messages=[{"role": "user", "content": "Hello my name is ishaan@berri.ai"}],
)
print("response headers=", json.dumps(headers, indent=4))
assert "x-litellm-applied-guardrails" in headers
assert headers["x-litellm-applied-guardrails"] == "bedrock-pre-guard"
async def get_guardrail_lb_counts(session):
"""Get the current guardrail load balancing call counts from the proxy."""
url = "http://0.0.0.0:4000/guardrail/lb/counts"

View file

@ -1,70 +0,0 @@
"""
Tests for Key based logging callbacks
"""
import os
import httpx
import pytest
@pytest.mark.asyncio()
async def test_key_logging_callbacks():
"""
Create virtual key with a logging callback set on the key
Call /key/health for the key -> it should be unhealthy
"""
# Generate a key with logging callback
generate_url = "http://0.0.0.0:4000/key/generate"
generate_headers = {
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
"Content-Type": "application/json",
}
generate_payload = {
"metadata": {
"logging": [
{
"callback_name": "gcs_bucket",
"callback_type": "success_and_failure",
"callback_vars": {
"gcs_bucket_name": "key-logging-project1",
"gcs_path_service_account": "bad-service-account",
},
}
]
}
}
async with httpx.AsyncClient() as client:
generate_response = await client.post(
generate_url, headers=generate_headers, json=generate_payload
)
assert generate_response.status_code == 200
generate_data = generate_response.json()
assert "key" in generate_data
_key = generate_data["key"]
# Check key health
health_url = "http://localhost:4000/key/health"
health_headers = {
"Authorization": f"Bearer {_key}",
"Content-Type": "application/json",
}
async with httpx.AsyncClient() as client:
health_response = await client.post(health_url, headers=health_headers, json={})
assert health_response.status_code == 200
health_data = health_response.json()
print("key_health_data", health_data)
# Check the response format and content
assert "key" in health_data
assert "logging_callbacks" in health_data
assert health_data["logging_callbacks"]["callbacks"] == ["gcs_bucket"]
assert health_data["logging_callbacks"]["status"] == "unhealthy"
assert (
"GCS_BUCKET_NAME is not set in the environment"
in health_data["logging_callbacks"]["details"]
)

View file

@ -1,29 +0,0 @@
"""
/model/info test
"""
import os
import httpx
import pytest
@pytest.mark.asyncio()
async def test_custom_model_supports_vision():
async with httpx.AsyncClient() as client:
response = await client.get(
"http://localhost:4000/model/info",
headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"},
)
assert response.status_code == 200
data = response.json()["data"]
print("response from /model/info", data)
llava_model = next(
(model for model in data if model["model_name"] == "llava-hf"), None
)
assert llava_model is not None, "llava-hf model not found in response"
assert (
llava_model["model_info"]["supports_vision"] == True
), "llava-hf model should support vision"

View file

@ -28,28 +28,6 @@ async def make_moderations_curl_request(
return await response.json()
@pytest.mark.asyncio
async def test_basic_moderations_on_proxy_no_model():
"""
Test moderations endpoint on proxy when no `model` is specified in the request
"""
async with aiohttp.ClientSession() as session:
test_text = "I want to harm someone" # Test text that should trigger moderation
request_data = {
"input": test_text,
}
try:
response = await make_moderations_curl_request(
session,
os.environ["LITELLM_MASTER_KEY"],
request_data,
)
print("response=", response)
except Exception as e:
print(e)
pytest.fail("Moderations request failed")
@pytest.mark.asyncio
async def test_basic_moderations_on_proxy_with_model():
"""

View file

@ -1,911 +0,0 @@
"""
Unit tests for prometheus metrics
"""
import os
import pytest
import aiohttp
import asyncio
from litellm._uuid import uuid
from openai import AsyncOpenAI
from typing import Dict, Any
END_USER_ID = "my-test-user-34"
async def make_bad_chat_completion_request(session, key):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": "fake-azure-endpoint",
"messages": [{"role": "user", "content": "Hello"}],
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
return status, response_text
async def make_good_chat_completion_request(session, key):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": "fake-openai-endpoint",
"messages": [{"role": "user", "content": f"Hello {uuid.uuid4()}"}],
"tags": ["teamB"],
"user": END_USER_ID, # test if disable end user tracking for prometheus works
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
return status, response_text
async def make_chat_completion_request_with_fallback(session, key):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
data = {
"model": "fake-azure-endpoint",
"messages": [{"role": "user", "content": "Hello"}],
"fallbacks": ["fake-openai-endpoint"],
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
# make a request with a failed fallback
data = {
"model": "fake-azure-endpoint",
"messages": [{"role": "user", "content": "Hello"}],
"fallbacks": ["unknown-model"],
}
async with session.post(url, headers=headers, json=data) as response:
status = response.status
response_text = await response.text()
return
@pytest.mark.asyncio
async def test_proxy_failure_metrics():
"""
- Make 1 bad chat completion call to "fake-azure-endpoint"
- GET /metrics
- assert the failure metric for the requested model is incremented by 1
- Assert the Exception class and status code are correct
"""
async with aiohttp.ClientSession() as session:
# Make a bad chat completion call
status, response_text = await make_bad_chat_completion_request(
session, os.environ["LITELLM_MASTER_KEY"]
)
# Check if the request failed as expected
assert status == 429, f"Expected status 429, but got {status}"
# Get metrics
async with session.get("http://0.0.0.0:4000/metrics") as response:
metrics = await response.text()
print("/metrics", metrics)
# Check if the failure metric is present and correct - use pattern matching for robustness
# Labels are ordered alphabetically by Prometheus: api_key_alias, end_user, exception_class,
# exception_status, hashed_api_key, requested_model, route, team, team_alias, user, user_email
# Note: client_ip, user_agent, model_id are present but we use substring matching to be flexible
# Check for both the new metric and deprecated metric for backwards compatibility
expected_patterns = [
"litellm_proxy_failed_requests_metric_total{", # New metric
"litellm_llm_api_failed_requests_metric_total{", # Deprecated but may still be used
]
# Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for
# hash_token(master_key) so the master key (or its hash) never
# propagates into metrics. See PR #26484.
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS
# Check if either pattern is in metrics and contains required fields
found_metric = False
for pattern in expected_patterns:
for line in metrics.split("\n"):
# For proxy metric, check proxy-specific fields
if "litellm_proxy_failed_requests_metric_total{" in line:
if (
'api_key_alias="None"' in line
and 'exception_class="Openai.RateLimitError"' in line
and 'exception_status="429"' in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-azure-endpoint"' in line
and 'route="/chat/completions"' in line
):
found_metric = True
break
# For deprecated llm_api metric, check llm-specific fields
elif "litellm_llm_api_failed_requests_metric_total{" in line:
if (
f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'model="429"' in line
): # The deprecated metric uses the actual model from the request
found_metric = True
break
if found_metric:
break
assert (
found_metric
), f"Expected failure metric not found in /metrics. Looking for either litellm_proxy_failed_requests_metric_total or litellm_llm_api_failed_requests_metric_total with required fields"
# Check total requests metric similarly
# The litellm_proxy_total_requests_metric_total should be present
total_requests_pattern = "litellm_proxy_total_requests_metric_total{"
found_total_metric = False
for line in metrics.split("\n"):
if (
total_requests_pattern in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-azure-endpoint"' in line
and 'status_code="429"' in line
):
found_total_metric = True
break
assert (
found_total_metric
), f"Expected total requests metric not found in /metrics. Looking for: {total_requests_pattern} with hashed_api_key and status_code=429"
@pytest.mark.asyncio
@pytest.mark.flaky(retries=3, delay=2)
async def test_proxy_success_metrics():
"""
Make 1 good /chat/completions call to "openai/gpt-5-mini"
GET /metrics
Assert the success metric is incremented by 1
"""
async with aiohttp.ClientSession() as session:
# Make a good chat completion call
status, response_text = await make_good_chat_completion_request(
session, os.environ["LITELLM_MASTER_KEY"]
)
# Check if the request succeeded as expected
assert status == 200, f"Expected status 200, but got {status}"
# Get metrics
async with session.get("http://0.0.0.0:4000/metrics") as response:
metrics = await response.text()
print("/metrics", metrics)
assert END_USER_ID not in metrics
# Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for
# hash_token(master_key) (PR #26484).
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS
# Check if the success metric is present and correct - use flexible matching
# Check for request_total_latency_metric with required fields
# Note: The model can be "gpt-3.5-turbo-0301" or similar depending on what's returned
found_request_latency = False
for line in metrics.split("\n"):
if (
"litellm_request_total_latency_metric_bucket{" in line
and 'api_key_alias="None"' in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-openai-endpoint"' in line
and 'le="0.005"' in line
):
found_request_latency = True
break
assert (
found_request_latency
), "Expected litellm_request_total_latency_metric_bucket not found in /metrics"
# Check for llm_api_latency_metric with required fields
found_api_latency = False
for line in metrics.split("\n"):
if (
"litellm_llm_api_latency_metric_bucket{" in line
and 'api_key_alias="None"' in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-openai-endpoint"' in line
and 'le="0.005"' in line
):
found_api_latency = True
break
assert (
found_api_latency
), "Expected litellm_llm_api_latency_metric_bucket not found in /metrics"
verify_latency_metrics(metrics)
def verify_latency_metrics(metrics: str):
"""
Assert that LATENCY_BUCKETS distribution is used for
- litellm_request_total_latency_metric_bucket
- litellm_llm_api_latency_metric_bucket
Very important to verify that the overhead latency metric is present
"""
from litellm.types.integrations.prometheus import LATENCY_BUCKETS
import re
import time
time.sleep(2)
metric_names = [
"litellm_request_total_latency_metric_bucket",
"litellm_llm_api_latency_metric_bucket",
"litellm_overhead_latency_metric_bucket",
]
for metric_name in metric_names:
# Extract all 'le' values for the current metric
pattern = rf'{metric_name}{{.*?le="(.*?)".*?}}'
le_values = re.findall(pattern, metrics)
# Convert to set for easier comparison
actual_buckets = set(le_values)
print("actual_buckets", actual_buckets)
expected_buckets = []
for bucket in LATENCY_BUCKETS:
expected_buckets.append(str(bucket))
# replace inf with +Inf
expected_buckets = [
bucket.replace("inf", "+Inf") for bucket in expected_buckets
]
print("expected_buckets", expected_buckets)
expected_buckets = set(expected_buckets)
# Verify all expected buckets are present
assert (
actual_buckets == expected_buckets
), f"Mismatch in {metric_name} buckets. Expected: {expected_buckets}, Got: {actual_buckets}"
@pytest.mark.asyncio
async def test_proxy_fallback_metrics():
"""
Make 1 request with a client side fallback - check metrics
"""
async with aiohttp.ClientSession() as session:
# Make a good chat completion call
await make_chat_completion_request_with_fallback(session, os.environ["LITELLM_MASTER_KEY"])
# Get metrics
async with session.get("http://0.0.0.0:4000/metrics") as response:
metrics = await response.text()
print("/metrics", metrics)
# Master-key auth substitutes LITELLM_PROXY_MASTER_KEY_ALIAS for
# hash_token(master_key) (PR #26484).
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
expected_hashed_api_key = LITELLM_PROXY_MASTER_KEY_ALIAS
# Check if successful fallback metric is incremented - use flexible matching
found_successful_fallback = False
for line in metrics.split("\n"):
if (
"litellm_deployment_successful_fallbacks_total{" in line
and 'api_key_alias="None"' in line
and 'exception_class="Openai.RateLimitError"' in line
and 'exception_status="429"' in line
and 'fallback_model="fake-openai-endpoint"' in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-azure-endpoint"' in line
and "1.0" in line
):
found_successful_fallback = True
break
assert (
found_successful_fallback
), "Expected litellm_deployment_successful_fallbacks_total metric not found in /metrics"
# Check if failed fallback metric is incremented - use flexible matching
found_failed_fallback = False
for line in metrics.split("\n"):
if (
"litellm_deployment_failed_fallbacks_total{" in line
and 'api_key_alias="None"' in line
and 'exception_class="Openai.RateLimitError"' in line
and 'exception_status="429"' in line
and 'fallback_model="unknown-model"' in line
and f'hashed_api_key="{expected_hashed_api_key}"' in line
and 'requested_model="fake-azure-endpoint"' in line
and "1.0" in line
):
found_failed_fallback = True
break
assert (
found_failed_fallback
), "Expected litellm_deployment_failed_fallbacks_total metric not found in /metrics"
async def create_test_team(
session: aiohttp.ClientSession, team_data: Dict[str, Any]
) -> str:
"""Create a new team and return the team_id"""
url = "http://0.0.0.0:4000/team/new"
headers = {
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
"Content-Type": "application/json",
}
async with session.post(url, headers=headers, json=team_data) as response:
assert (
response.status == 200
), f"Failed to create team. Status: {response.status}"
team_info = await response.json()
return team_info["team_id"]
async def create_test_user(
session: aiohttp.ClientSession, user_data: Dict[str, Any]
) -> Dict[str, Any]:
"""Create a new user and return the user info"""
url = "http://0.0.0.0:4000/user/new"
headers = {
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
"Content-Type": "application/json",
}
async with session.post(url, headers=headers, json=user_data) as response:
assert (
response.status == 200
), f"Failed to create user. Status: {response.status}"
user_info = await response.json()
return user_info
async def get_prometheus_metrics(session: aiohttp.ClientSession) -> str:
"""Fetch current prometheus metrics"""
async with session.get("http://0.0.0.0:4000/metrics") as response:
assert response.status == 200
return await response.text()
def extract_budget_metrics(metrics_text: str, team_id: str) -> Dict[str, float]:
"""Extract budget-related metrics for a specific team"""
import re
metrics = {}
# Get remaining budget
remaining_pattern = f'litellm_remaining_team_budget_metric{{team="{team_id}",team_alias="[^"]*"}} ([0-9.]+)'
remaining_match = re.search(remaining_pattern, metrics_text)
metrics["remaining"] = float(remaining_match.group(1)) if remaining_match else None
# Get total budget
total_pattern = f'litellm_team_max_budget_metric{{team="{team_id}",team_alias="[^"]*"}} ([0-9.]+)'
total_match = re.search(total_pattern, metrics_text)
metrics["total"] = float(total_match.group(1)) if total_match else None
# Get remaining hours
hours_pattern = f'litellm_team_budget_remaining_hours_metric{{team="{team_id}",team_alias="[^"]*"}} ([0-9.]+)'
hours_match = re.search(hours_pattern, metrics_text)
metrics["remaining_hours"] = float(hours_match.group(1)) if hours_match else None
return metrics
async def create_test_key(session: aiohttp.ClientSession, team_id: str) -> str:
"""Generate a new key for the team and return it"""
url = "http://0.0.0.0:4000/key/generate"
headers = {
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
"Content-Type": "application/json",
}
data = {
"team_id": team_id,
}
async with session.post(url, headers=headers, json=data) as response:
assert (
response.status == 200
), f"Failed to generate key. Status: {response.status}"
key_info = await response.json()
return key_info["key"]
async def get_team_info(session: aiohttp.ClientSession, team_id: str) -> Dict[str, Any]:
"""Fetch team info and return the response"""
url = f"http://0.0.0.0:4000/team/info?team_id={team_id}"
headers = {
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
}
async with session.get(url, headers=headers) as response:
assert (
response.status == 200
), f"Failed to get team info. Status: {response.status}"
return await response.json()
@pytest.mark.asyncio
async def test_team_budget_metrics():
"""
Test team budget tracking metrics:
1. Create a team with max_budget
2. Generate a key for the team
3. Make chat completion requests using OpenAI SDK with team's key
4. Verify budget decreases over time
5. Verify request costs are being tracked correctly
6. Verify prometheus metrics match /team/info spend data
"""
async with aiohttp.ClientSession() as session:
# Setup test team
team_data = {
"team_alias": "budget_test_team",
"max_budget": 10,
"budget_duration": "7d",
}
team_id = await create_test_team(session, team_data)
print("team_id", team_id)
# Generate key for the team
team_key = await create_test_key(session, team_id)
# Initialize OpenAI client with team's key
client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=team_key)
# Make initial request and check budget
await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}],
)
await asyncio.sleep(11) # Wait for metrics to update
# Get metrics after request
metrics_after_first = await get_prometheus_metrics(session)
print("metrics_after_first", metrics_after_first)
first_budget = extract_budget_metrics(metrics_after_first, team_id)
print(f"Budget after 1 request: {first_budget}")
assert (
first_budget["remaining"] < 10.0
), "remaining budget should be less than 10.0 after first request"
assert first_budget["total"] == 10.0, "Total budget metric is incorrect"
print("first_budget['remaining_hours']", first_budget["remaining_hours"])
# Budget should have positive remaining hours, up to 7 days
assert (
0 < first_budget["remaining_hours"] <= 168
), "Budget should have positive remaining hours, up to 7 days"
# Get team info and verify spend matches prometheus metrics
team_info = await get_team_info(session, team_id)
print("team_info", team_info)
_team_info_data = team_info["team_info"]
# Calculate spend from prometheus (total - remaining)
team_info_spend = float(_team_info_data["spend"])
team_info_max_budget = float(_team_info_data["max_budget"])
team_info_remaining_budget = team_info_max_budget - team_info_spend
print("\n\n\n###### Final budget metrics ######\n\n\n")
print("team_info_remaining_budget", team_info_remaining_budget)
print("prometheus_remaining_budget", first_budget["remaining"])
print(
"diff between team_info_remaining_budget and prometheus_remaining_budget",
team_info_remaining_budget - first_budget["remaining"],
)
# Verify spends match within a small delta (floating point comparison)
assert (
abs(team_info_remaining_budget - first_budget["remaining"]) <= 0.001
), f"Spend mismatch: Prometheus={team_info_remaining_budget}, Team Info={first_budget['remaining']}"
async def create_test_key_with_budget(
session: aiohttp.ClientSession, budget_data: Dict[str, Any]
) -> str:
"""Generate a new key with budget constraints and return it"""
url = "http://0.0.0.0:4000/key/generate"
headers = {
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
"Content-Type": "application/json",
}
print("budget_data", budget_data)
async with session.post(url, headers=headers, json=budget_data) as response:
assert (
response.status == 200
), f"Failed to generate key. Status: {response.status}"
key_info = await response.json()
return key_info["key"]
async def get_key_info(session: aiohttp.ClientSession, key: str) -> Dict[str, Any]:
"""Fetch key info and return the response"""
url = "http://0.0.0.0:4000/key/info"
headers = {
"Authorization": f"Bearer {key}",
}
async with session.get(url, headers=headers) as response:
assert (
response.status == 200
), f"Failed to get key info. Status: {response.status}"
return await response.json()
async def get_user_info(session: aiohttp.ClientSession, user_id: str) -> Dict[str, Any]:
"""Fetch user info and return the response"""
from urllib.parse import quote
# URL encode user_id to handle special characters
encoded_user_id = quote(user_id, safe="")
url = f"http://0.0.0.0:4000/user/info?user_id={encoded_user_id}"
headers = {
"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}",
}
async with session.get(url, headers=headers) as response:
assert (
response.status == 200
), f"Failed to get user info. Status: {response.status}"
return await response.json()
def extract_key_budget_metrics(metrics_text: str, key_id: str) -> Dict[str, float]:
"""Extract budget-related metrics for a specific key"""
import re
metrics = {}
# Get remaining budget
remaining_pattern = f'litellm_remaining_api_key_budget_metric{{api_key_alias="[^"]*",hashed_api_key="{key_id}"}} ([0-9.]+)'
remaining_match = re.search(remaining_pattern, metrics_text)
metrics["remaining"] = float(remaining_match.group(1)) if remaining_match else None
# Get total budget
total_pattern = f'litellm_api_key_max_budget_metric{{api_key_alias="[^"]*",hashed_api_key="{key_id}"}} ([0-9.]+)'
total_match = re.search(total_pattern, metrics_text)
metrics["total"] = float(total_match.group(1)) if total_match else None
# Get remaining hours
hours_pattern = f'litellm_api_key_budget_remaining_hours_metric{{api_key_alias="[^"]*",hashed_api_key="{key_id}"}} ([0-9.]+)'
hours_match = re.search(hours_pattern, metrics_text)
metrics["remaining_hours"] = float(hours_match.group(1)) if hours_match else None
return metrics
def extract_user_budget_metrics(metrics_text: str, user_id: str) -> Dict[str, float]:
"""Extract budget-related metrics for a specific user"""
import re
metrics = {}
# Escape user_id for regex pattern matching
escaped_user_id = re.escape(user_id)
# Get remaining budget (user_email and user_alias may also be present as labels)
remaining_pattern = rf'litellm_remaining_user_budget_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)'
remaining_match = re.search(remaining_pattern, metrics_text)
metrics["remaining"] = float(remaining_match.group(1)) if remaining_match else None
# Get total budget
total_pattern = rf'litellm_user_max_budget_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)'
total_match = re.search(total_pattern, metrics_text)
metrics["total"] = float(total_match.group(1)) if total_match else None
# Get remaining hours
hours_pattern = rf'litellm_user_budget_remaining_hours_metric{{[^}}]*user="{escaped_user_id}"[^}}]*}} ([0-9.]+)'
hours_match = re.search(hours_pattern, metrics_text)
metrics["remaining_hours"] = float(hours_match.group(1)) if hours_match else None
return metrics
@pytest.mark.asyncio
async def test_key_budget_metrics():
"""
Test key budget tracking metrics:
1. Create a key with max_budget
2. Make chat completion requests using OpenAI SDK with the key
3. Verify budget decreases over time
4. Verify request costs are being tracked correctly
5. Verify prometheus metrics match /key/info spend data
"""
from datetime import datetime, timedelta, timezone
async with aiohttp.ClientSession() as session:
# Setup test key with unique alias
unique_alias = f"budget_test_key_{uuid.uuid4()}"
key_data = {
"key_alias": unique_alias,
"max_budget": 10,
"budget_duration": "7d",
"budget_reset_at": (
datetime.now(timezone.utc) + timedelta(days=7)
).isoformat(),
}
key = await create_test_key_with_budget(session, key_data)
# Extract key_id from the key info
key_info = await get_key_info(session, key)
print("key_info", key_info)
key_id = key_info["key"]
print("key_id", key_id)
# Initialize OpenAI client with the key
client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key)
# Make initial request and check budget
await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}],
)
await asyncio.sleep(11) # Wait for metrics to update
# Get metrics after request
metrics_after_first = await get_prometheus_metrics(session)
print("metrics_after_first request", metrics_after_first)
first_budget = extract_key_budget_metrics(metrics_after_first, key_id)
print(f"Budget after 1 request: {first_budget}")
assert (
first_budget["remaining"] < 10.0
), "remaining budget should be less than 10.0 after first request"
assert first_budget["total"] == 10.0, "Total budget metric is incorrect"
print("first_budget['remaining_hours']", first_budget["remaining_hours"])
# The budget reset time is now standardized - for "7d" it resets on Monday at midnight
# So we'll check if it's within a reasonable range (0-7 days depending on current day of week)
assert (
0 <= first_budget["remaining_hours"] <= 168
), "Budget remaining hours should be within a reasonable range (0-7 days depending on day of week)"
# Get key info and verify spend matches prometheus metrics
key_info = await get_key_info(session, key)
print("key_info", key_info)
_key_info_data = key_info["info"]
# Calculate spend from prometheus (total - remaining)
key_info_spend = float(_key_info_data["spend"])
key_info_max_budget = float(_key_info_data["max_budget"])
key_info_remaining_budget = key_info_max_budget - key_info_spend
print("\n\n\n###### Final budget metrics ######\n\n\n")
print("key_info_remaining_budget", key_info_remaining_budget)
print("prometheus_remaining_budget", first_budget["remaining"])
print(
"diff between key_info_remaining_budget and prometheus_remaining_budget",
key_info_remaining_budget - first_budget["remaining"],
)
# Verify spends match within a small delta (floating point comparison)
assert (
abs(key_info_remaining_budget - first_budget["remaining"]) <= 0.001
), f"Spend mismatch: Prometheus={key_info_remaining_budget}, Key Info={first_budget['remaining']}"
@pytest.mark.asyncio
async def test_user_budget_metrics():
"""
Test user budget tracking metrics:
1. Create a user with max_budget
2. Make chat completion requests using OpenAI SDK with the user's key
3. Verify budget decreases over time
4. Verify request costs are being tracked correctly
5. Verify prometheus metrics match /user/info spend data
"""
from datetime import datetime, timedelta, timezone
async with aiohttp.ClientSession() as session:
# Setup test user with unique user_id
unique_user_id = f"budget_test_user_{uuid.uuid4()}"
user_data = {
"user_id": unique_user_id,
"max_budget": 10,
"budget_duration": "7d",
"budget_reset_at": (
datetime.now(timezone.utc) + timedelta(days=7)
).isoformat(),
}
user_info = await create_test_user(session, user_data)
print("user_info", user_info)
user_id = user_info["user_id"]
print("user_id", user_id)
# Get the key that was created with the user
key = user_info["key"]
# Initialize OpenAI client with the user's key
client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key)
# Make initial request and check budget
await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}],
)
await asyncio.sleep(11) # Wait for metrics to update
# Get metrics after request
metrics_after_first = await get_prometheus_metrics(session)
print("metrics_after_first request", metrics_after_first)
first_budget = extract_user_budget_metrics(metrics_after_first, user_id)
print(f"Budget after 1 request: {first_budget}")
assert (
first_budget["remaining"] is not None
), "remaining budget metric should be present"
assert (
first_budget["total"] is not None
), "total budget metric should be present"
assert (
first_budget["remaining"] < 10.0
), "remaining budget should be less than 10.0 after first request"
assert first_budget["total"] == 10.0, "Total budget metric is incorrect"
print("first_budget['remaining_hours']", first_budget["remaining_hours"])
# The budget reset time is now standardized - for "7d" it resets on Monday at midnight
# So we'll check if it's within a reasonable range (0-7 days depending on current day of week)
assert (
first_budget["remaining_hours"] is not None
), "remaining hours metric should be present"
assert (
0 <= first_budget["remaining_hours"] <= 168
), "Budget remaining hours should be within a reasonable range (0-7 days depending on day of week)"
# Get user info and verify spend matches prometheus metrics
user_info_response = await get_user_info(session, user_id)
print("user_info_response", user_info_response)
_user_info_data = user_info_response["user_info"]
# Calculate spend from prometheus (total - remaining)
user_info_spend = float(_user_info_data["spend"])
user_info_max_budget = float(_user_info_data["max_budget"])
user_info_remaining_budget = user_info_max_budget - user_info_spend
print("\n\n\n###### Final budget metrics ######\n\n\n")
print("user_info_remaining_budget", user_info_remaining_budget)
print("prometheus_remaining_budget", first_budget["remaining"])
print(
"diff between user_info_remaining_budget and prometheus_remaining_budget",
user_info_remaining_budget - first_budget["remaining"],
)
# Verify spends match within a small delta (floating point comparison)
assert (
abs(user_info_remaining_budget - first_budget["remaining"]) <= 0.001
), f"Spend mismatch: Prometheus={user_info_remaining_budget}, User Info={first_budget['remaining']}"
@pytest.mark.asyncio
async def test_user_email_metrics():
"""
Test user email tracking metrics:
1. Create a user with user_email
2. Make chat completion requests using OpenAI SDK with the user's email
3. Verify user email is being tracked correctly in `litellm_user_email_metric`
"""
async with aiohttp.ClientSession() as session:
# Create a user with user_email
user_email = f"test-{uuid.uuid4()}@example.com"
user_data = {
"user_email": user_email,
}
user_info = await create_test_user(session, user_data)
key = user_info["key"]
# Initialize OpenAI client with the user's email
client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key)
# Make initial request and check budget
await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}],
)
await asyncio.sleep(11) # Wait for metrics to update
# Get metrics after request
metrics_after_first = await get_prometheus_metrics(session)
print("metrics_after_first request", metrics_after_first)
assert (
user_email in metrics_after_first
), "user_email should be tracked correctly"
@pytest.mark.asyncio
async def test_user_email_in_all_required_metrics():
"""
Test that user_email label is present in all the metrics that were requested to have it:
- litellm_proxy_total_requests_metric_total
- litellm_proxy_failed_requests_metric_total
- litellm_input_tokens_metric_total
- litellm_output_tokens_metric_total
- litellm_requests_metric_total
- litellm_spend_metric_total
"""
async with aiohttp.ClientSession() as session:
# Create a user with user_email
user_email = f"test-metrics-{uuid.uuid4()}@example.com"
user_data = {
"user_email": user_email,
}
user_info = await create_test_user(session, user_data)
key = user_info["key"]
# Initialize OpenAI client with the user's email
client = AsyncOpenAI(base_url="http://0.0.0.0:4000", api_key=key)
# Make successful request to generate metrics
await client.chat.completions.create(
model="fake-openai-endpoint",
messages=[{"role": "user", "content": f"Hello {uuid.uuid4()}"}],
)
await asyncio.sleep(11) # Wait for metrics to update
# Get metrics after request
metrics_text = await get_prometheus_metrics(session)
print("Testing user_email in all required metrics")
# Check that user_email appears in all the required metrics
required_metrics_with_user_email = [
# "litellm_proxy_total_requests_metric_total",
# "litellm_input_tokens_metric_total",
# "litellm_output_tokens_metric_total",
# "litellm_requests_metric_total",
"litellm_spend_metric_total",
]
import re
for metric_name in required_metrics_with_user_email:
# Check that the metric exists and contains user_email label
# Look for the metric with user_email in its labels
pattern = (
rf'{metric_name}{{[^}}]*user_email="{re.escape(user_email)}"[^}}]*}}'
)
matches = re.findall(pattern, metrics_text)
assert (
len(matches) > 0
), f"Metric {metric_name} should contain user_email={user_email} but was not found in metrics"
# Also test failure metric by making a bad request
try:
await client.chat.completions.create(
model="fake-azure-endpoint", # This should fail
messages=[{"role": "user", "content": "Hello"}],
)
except Exception:
pass # Expected to fail
await asyncio.sleep(11) # Wait for metrics to update
# Get updated metrics
metrics_text = await get_prometheus_metrics(session)
# Check that failure metric also contains user_email
failure_pattern = rf'litellm_proxy_failed_requests_metric_total{{[^}}]*user_email="{re.escape(user_email)}"[^}}]*}}'
failure_matches = re.findall(failure_pattern, metrics_text)
assert (
len(failure_matches) > 0
), f"litellm_proxy_failed_requests_metric_total should contain user_email={user_email}"

View file

@ -1,65 +0,0 @@
import os
# What this tests ?
## Set tags on a team and then make a request to /chat/completions
import pytest
import asyncio
import aiohttp, openai
from openai import OpenAI, AsyncOpenAI
from typing import Optional, List, Union
from litellm._uuid import uuid
LITELLM_MASTER_KEY = os.environ["LITELLM_MASTER_KEY"]
async def chat_completion(
session, key, model: Union[str, List] = "fake-openai-endpoint"
):
url = "http://0.0.0.0:4000/chat/completions"
headers = {
"Authorization": f"Bearer {key}",
"Content-Type": "application/json",
}
print("headers=", headers)
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()
if status != 200:
raise Exception(response_text)
return await response.json(), response.headers
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}"
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:
raise Exception(response_text)
return await response.json()
@pytest.mark.asyncio()
async def test_chat_completion_with_no_tags():
async with aiohttp.ClientSession() as session:
key = LITELLM_MASTER_KEY
response, headers = await chat_completion(session, key)
headers = dict(headers)
print(response)
print(headers)
assert response is not None

View file

@ -0,0 +1,344 @@
import asyncio
import datetime
from collections.abc import Sequence
from typing import Final, Literal, TypedDict
from typing_extensions import ReadOnly
import httpx
import pytest
import respx
from pydantic import TypeAdapter
import litellm
import litellm.proxy.proxy_server as proxy_server
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.langfuse.langfuse import LangFuseLogger, installed_langfuse_version
from litellm.integrations.langfuse.langfuse_sdk import (
build_langfuse_client,
build_langfuse_tracing,
resolve_trace_id,
)
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
from litellm.integrations.SlackAlerting.utils import add_langfuse_trace_id_to_alert
from litellm.litellm_core_utils import litellm_logging
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.utils import ProxyLogging
from litellm.types.integrations.slack_alerting import AlertType
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
_WEBHOOK: Final = "https://hooks.slack.example/services/delivery"
_LANGFUSE_HOST: Final = "https://langfuse.alerts.example"
_AZURE_BASE: Final = "https://openai-gpt-4-test-v-1.openai.azure.com/"
_DAILY_BASE: Final = "https://daily-report.openai.example/v1"
class _SlackPayload(TypedDict):
text: ReadOnly[str]
class _TeamRow(TypedDict):
team_alias: ReadOnly[str]
total_spend: ReadOnly[float]
class _TagRow(TypedDict):
individual_request_tag: ReadOnly[str]
total_spend: ReadOnly[float]
_PAYLOAD: Final = TypeAdapter(_SlackPayload)
@pytest.fixture(autouse=True)
def _slack_webhook(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setenv("SLACK_WEBHOOK_URL", _WEBHOOK)
def _webhook(respx_mock: respx.MockRouter) -> respx.Route:
return respx_mock.post(_WEBHOOK).mock(return_value=httpx.Response(200, text="ok"))
def _posted_texts(route: respx.Route) -> tuple[str, ...]:
return tuple(_PAYLOAD.validate_json(call.request.content)["text"] for call in route.calls)
@pytest.mark.asyncio
async def test_slow_response_alert_names_the_azure_api_base_and_reaches_the_webhook(
respx_mock: respx.MockRouter,
) -> None:
route: Final = _webhook(respx_mock)
proxy_logging: Final = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.update_values(alerting=["slack"], alerting_threshold=100, redis_cache=None)
start: Final = datetime.datetime(2026, 1, 1, 12, 0, 0)
messages: Final = ({"role": "user", "content": "Hey how's it going?"},)
helper_result: Final = proxy_logging.slack_alerting_instance._response_taking_too_long_callback_helper(
kwargs={
"model": "chatgpt-v-3",
"messages": messages,
"litellm_params": {"api_base": _AZURE_BASE, "custom_llm_provider": "azure"},
},
start_time=start,
end_time=start + datetime.timedelta(seconds=150),
)
assert helper_result == (150.0, "chatgpt-v-3", _AZURE_BASE, str(messages)[:100])
slow_message: Final = (
f"`Responses are slow - 150.0s response time > Alerting threshold: 100s`\nAPI Base: `{_AZURE_BASE}`"
)
await proxy_logging.alerting_handler(message=slow_message, level="Low", alert_type=AlertType.llm_too_slow)
await proxy_logging.slack_alerting_instance.flush_queue()
texts: Final = _posted_texts(route)
assert len(texts) == 1
assert texts[0].startswith("Alert type: `llm_too_slow`\nLevel: `Low`\n")
assert texts[0].endswith(f"Message: {slow_message}")
@pytest.mark.asyncio
async def test_send_alert_is_queued_until_flush_then_posted_to_the_webhook_once(respx_mock: respx.MockRouter) -> None:
route: Final = _webhook(respx_mock)
slack_alerting: Final = SlackAlerting(alerting_threshold=1, internal_usage_cache=DualCache(), alerting=["slack"])
await slack_alerting.send_alert("Test message", "Low", AlertType.budget_alerts, alerting_metadata={})
assert route.call_count == 0
await slack_alerting.flush_queue()
await slack_alerting.flush_queue()
texts: Final = _posted_texts(route)
assert len(texts) == 1
assert texts[0].startswith("Alert type: `budget_alerts`\nLevel: `Low`\n")
assert texts[0].endswith("Message: Test message")
@pytest.mark.asyncio
async def test_a_queued_alert_is_posted_by_the_periodic_flush_without_a_manual_flush(
respx_mock: respx.MockRouter,
) -> None:
delivered: Final = asyncio.Event()
def deliver(request: httpx.Request) -> httpx.Response:
delivered.set()
return httpx.Response(200, text="ok")
route: Final = respx_mock.post(_WEBHOOK).mock(side_effect=deliver)
slack_alerting: Final = SlackAlerting(alerting_threshold=1, internal_usage_cache=DualCache(), alerting=["slack"])
slack_alerting.flush_interval = 0
slack_alerting.update_values(alerting=["slack"])
flush_task: Final = slack_alerting._periodic_flush_task
assert flush_task is not None
try:
await slack_alerting.send_alert("Timed message", "Low", AlertType.budget_alerts, alerting_metadata={})
await asyncio.wait_for(delivered.wait(), timeout=5)
finally:
flush_task.cancel()
texts: Final = _posted_texts(route)
assert len(texts) == 1
assert texts[0].endswith("Message: Timed message")
class _DeploymentSettled(CustomLogger):
def __init__(self, model_id: str) -> None:
super().__init__()
self.model_id: Final = model_id
self.succeeded: Final = asyncio.Event()
self.failed: Final = asyncio.Event()
def _is_mine(self, kwargs: dict[str, object]) -> bool:
payload: Final = kwargs.get("standard_logging_object")
return isinstance(payload, dict) and payload.get("model_id") == self.model_id
async def async_log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
if self._is_mine(kwargs):
self.succeeded.set()
async def async_log_failure_event(
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
if self._is_mine(kwargs):
self.failed.set()
@pytest.mark.asyncio
async def test_daily_report_lists_router_latency_after_success_and_failures_after_an_auth_error(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
webhook: Final = _webhook(respx_mock)
respx_mock.post(f"{_DAILY_BASE}/chat/completions").mock(
side_effect=(
httpx.Response(
200,
json={
"id": "chatcmpl-daily",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-5-mini",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "fine"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 4, "total_tokens": 9},
},
),
httpx.Response(
401,
json={
"error": {
"message": "Incorrect API key provided",
"type": "invalid_request_error",
"code": "invalid_api_key",
}
},
),
)
)
model_id: Final = "daily-report-deployment"
slack_alerting: Final = SlackAlerting(
alerting=["slack"], internal_usage_cache=DualCache(), alert_types=[AlertType.daily_reports]
)
settled: Final = _DeploymentSettled(model_id)
monkeypatch.setattr(litellm, "callbacks", [slack_alerting, settled])
router: Final = litellm.Router(
model_list=[
{
"model_name": "daily-report-model",
"litellm_params": {"model": "openai/gpt-5-mini", "api_key": "sk-daily", "api_base": _DAILY_BASE},
"model_info": {"id": model_id},
}
]
)
request: Final = ({"role": "user", "content": "Hey, how's it going?"},)
await router.acompletion(model="daily-report-model", messages=list(request))
await asyncio.wait_for(settled.succeeded.wait(), timeout=5)
after_success: Final = await slack_alerting.send_daily_reports(router=router)
await slack_alerting.flush_queue()
with pytest.raises(litellm.AuthenticationError):
await router.acompletion(model="daily-report-model", messages=list(request))
await asyncio.wait_for(settled.failed.wait(), timeout=5)
after_failure: Final = await slack_alerting.send_daily_reports(router=router)
await slack_alerting.flush_queue()
texts: Final = _posted_texts(webhook)
assert (after_success, after_failure) == (True, True)
assert len(texts) == 2
assert "Most Failed Requests:*\n\n\tNone\n" in texts[0]
assert "1. Deployment: `openai/gpt-5-mini`, Latency per output token: `" in texts[0]
assert f"1. Deployment: `openai/gpt-5-mini`, Failed Requests: `1`, API Base: `{_DAILY_BASE}`" in texts[1]
assert "Top Slowest Deployments:*\n\n\tNone\n" in texts[1]
class _CallLogged(CustomLogger):
def __init__(self, call_id: str, loop: asyncio.AbstractEventLoop) -> None:
super().__init__()
self.call_id: Final = call_id
self.loop: Final = loop
self.logged: Final = asyncio.Event()
def log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
if kwargs.get("litellm_call_id") == self.call_id:
self.loop.call_soon_threadsafe(self.logged.set)
@pytest.mark.asyncio
async def test_langfuse_trace_link_ends_with_the_trace_id_the_logger_emitted(monkeypatch: pytest.MonkeyPatch) -> None:
logger: Final = LangFuseLogger.__new__(LangFuseLogger)
logger.tracing = build_langfuse_tracing(
exporter=InMemorySpanExporter(), environment=None, release=None, sample_rate=1.0, flush_interval_millis=10
)
logger.api_client = build_langfuse_client(
public_key="pk-alert-trace", secret_key="sk-alert-trace", base_url=_LANGFUSE_HOST, httpx_client=None
)
logger.langfuse_sdk_version = installed_langfuse_version()
call_id: Final = "slack-alert-langfuse-trace"
logged: Final = _CallLogged(call_id, asyncio.get_running_loop())
monkeypatch.setenv("LANGFUSE_HOST", _LANGFUSE_HOST)
monkeypatch.setattr(litellm_logging, "langFuseLogger", logger)
monkeypatch.setattr(litellm, "success_callback", ["langfuse", logged])
monkeypatch.setattr(litellm, "_async_success_callback", [])
monkeypatch.setattr(litellm, "callbacks", [])
logging_obj: Final = Logging(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="completion",
litellm_call_id=call_id,
start_time=datetime.datetime.now(),
function_id=call_id,
)
litellm.completion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "Hey how's it going?"}],
mock_response="Hey!",
litellm_logging_obj=logging_obj,
)
await asyncio.wait_for(logged.logged.wait(), timeout=5)
trace_url: Final = await add_langfuse_trace_id_to_alert(request_data={"litellm_logging_obj": logging_obj})
expected_trace_id: Final = resolve_trace_id(logging_obj.litellm_trace_id)
assert logging_obj.get_trace_id(service_name="langfuse") == expected_trace_id
assert trace_url == f"{_LANGFUSE_HOST}/trace/{expected_trace_id}"
class _SpendReportDb:
def __init__(self, teams: Sequence[_TeamRow], tags: Sequence[_TagRow]) -> None:
self.teams: Final = teams
self.tags: Final = tags
async def query_raw(self, query: str, *args: object) -> Sequence[_TeamRow] | Sequence[_TagRow]:
return self.teams if "team_alias" in query else self.tags
class _SpendReportPrisma:
def __init__(self, db: _SpendReportDb) -> None:
self.db: Final = db
@pytest.mark.parametrize("report_type", ["weekly", "monthly"])
@pytest.mark.asyncio
async def test_spend_report_is_sent_once_per_period(
report_type: Literal["weekly", "monthly"], respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
route: Final = _webhook(respx_mock)
monkeypatch.setattr(
proxy_server,
"prisma_client",
_SpendReportPrisma(
_SpendReportDb(
teams=(
_TeamRow(team_alias="team1", total_spend=100.0),
_TeamRow(team_alias="team2", total_spend=200.0),
),
tags=(
_TagRow(individual_request_tag="tag1", total_spend=150.0),
_TagRow(individual_request_tag="tag2", total_spend=150.0),
),
)
),
)
slack_alerting: Final = SlackAlerting(alerting=["slack"], internal_usage_cache=DualCache())
send_report: Final = (
slack_alerting.send_weekly_spend_report if report_type == "weekly" else slack_alerting.send_monthly_spend_report
)
await send_report()
await slack_alerting.flush_queue()
await send_report()
await slack_alerting.flush_queue()
texts: Final = _posted_texts(route)
assert len(texts) == 1
assert "Team: `team1` | Spend: `$100.0`\nTeam: `team2` | Spend: `$200.0`\n" in texts[0]
assert "Tag: `tag1` | Spend: `$150.0`\nTag: `tag2` | Spend: `$150.0`\n" in texts[0]

View file

@ -2,14 +2,20 @@ import gzip
import json
import os
from datetime import datetime
from typing import Coroutine, Final
from pathlib import Path
from typing import Coroutine, Final, TypedDict
from unittest.mock import AsyncMock, patch
import pytest
import respx
from httpx import Request, Response
from pydantic import TypeAdapter
from typing_extensions import ReadOnly
import litellm
import litellm.integrations.datadog.datadog as datadog_module
from litellm.caching.caching import Cache
from litellm.caching.llm_caching_handler import LLMClientCache
from litellm.integrations.datadog.datadog import DataDogLogger
from litellm.integrations.datadog.datadog_handler import (
get_datadog_env,
@ -19,6 +25,7 @@ from litellm.integrations.datadog.datadog_handler import (
get_datadog_source,
get_datadog_tags,
)
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.types.integrations.datadog import DatadogInitParams, DatadogPayload, DataDogStatus
from litellm.types.utils import (
StandardLoggingHiddenParams,
@ -803,3 +810,113 @@ def create_standard_logging_payload() -> StandardLoggingPayload:
additional_headers=None,
),
)
_INTAKE_URL: Final = "https://http-intake.logs.test.datadoghq.com/api/v2/logs"
class _ServiceEventMessage(TypedDict):
service: ReadOnly[str]
call_type: ReadOnly[str]
error: ReadOnly[str]
is_error: ReadOnly[bool]
@pytest.fixture
def delivery(
datadog_env: None, monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> tuple[DataDogLogger, respx.Route]:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache())
monkeypatch.delenv("DD_SOURCE", raising=False)
monkeypatch.delenv("DD_SERVICE", raising=False)
with patch("asyncio.create_task", side_effect=_discard_periodic_flush):
logger: Final = DataDogLogger()
intake: Final = respx_mock.post(_INTAKE_URL).mock(return_value=Response(202, text="Accepted"))
return logger, intake
def _delivered_logs(intake: respx.Route) -> list[DatadogPayload]:
return TypeAdapter(list[DatadogPayload]).validate_json(gzip.decompress(intake.calls.last.request.content))
@pytest.mark.asyncio
async def test_a_successful_request_is_delivered_as_an_info_log_carrying_the_standard_payload(
delivery: tuple[DataDogLogger, respx.Route],
) -> None:
datadog_logger, intake = delivery
standard_payload: Final = _standard_logging_payload()
await datadog_logger.async_log_success_event(
kwargs={"standard_logging_object": standard_payload},
response_obj=None,
start_time=STANDARD_START_TIME,
end_time=STANDARD_END_TIME,
)
await datadog_logger.async_send_batch()
assert intake.call_count == 1
logs: Final = _delivered_logs(intake)
assert len(logs) == 1
assert logs[0]["ddsource"] == "litellm"
assert logs[0]["service"] == "litellm-server"
assert logs[0]["status"] == DataDogStatus.INFO
assert TypeAdapter(dict[str, object]).validate_json(logs[0]["message"]) == standard_payload
@pytest.mark.asyncio
async def test_a_failed_request_is_delivered_as_an_error_log_that_keeps_the_error_string(
delivery: tuple[DataDogLogger, respx.Route],
) -> None:
datadog_logger, intake = delivery
standard_payload: Final = _standard_logging_payload()
standard_payload["status"] = "failure"
standard_payload["error_str"] = "Test error"
await datadog_logger.async_log_failure_event(
kwargs={"standard_logging_object": standard_payload},
response_obj=None,
start_time=STANDARD_START_TIME,
end_time=STANDARD_END_TIME,
)
await datadog_logger.async_send_batch()
assert intake.call_count == 1
logs: Final = _delivered_logs(intake)
assert len(logs) == 1
assert logs[0]["status"] == DataDogStatus.ERROR
message: Final = TypeAdapter(dict[str, object]).validate_json(logs[0]["message"])
assert message == standard_payload
assert message["error_str"] == "Test error"
@pytest.mark.asyncio
async def test_a_failing_redis_cache_is_delivered_to_datadog_as_redis_warnings(
delivery: tuple[DataDogLogger, respx.Route], monkeypatch: pytest.MonkeyPatch, tmp_path: Path
) -> None:
datadog_logger, intake = delivery
absent_socket: Final = str(tmp_path / "absent.sock")
redis_cache: Final = Cache(type="redis", url=f"unix://{absent_socket}")
monkeypatch.setattr(redis_cache.cache.service_logger_obj, "dd_logger", datadog_logger, raising=False)
monkeypatch.setattr(litellm, "service_callback", ["datadog"])
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr(litellm, "cache", redis_cache)
for _ in range(3):
await litellm.acompletion(
model="gpt-4.1-mini",
messages=[{"role": "user", "content": "what llm are u"}],
mock_response="Accepted",
caching=True,
)
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10)
await datadog_logger.async_send_batch()
assert intake.call_count == 1
logs: Final = _delivered_logs(intake)
assert len(logs) > 0
assert {log["status"] for log in logs} == {DataDogStatus.WARN}
messages: Final = [TypeAdapter(_ServiceEventMessage).validate_json(log["message"]) for log in logs]
assert {message["service"] for message in messages} == {"redis"}
assert all(message["is_error"] is True for message in messages)
assert all(absent_socket in message["error"] for message in messages)

View file

@ -0,0 +1,157 @@
import asyncio
import json
from collections.abc import Sequence
from typing import Final
import httpx
import pytest
import respx
from opentelemetry.sdk.trace import ReadableSpan, TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor, SpanExportResult
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
import litellm
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
_OPENAI_URL: Final = "https://api.openai.com/v1/chat/completions"
_EXPECTED_SPAN_NAMES: Final = ("litellm_request", "raw_gen_ai_request")
_USER: Final = "OTEL_USER"
_USAGE: Final = {"prompt_tokens": 8, "completion_tokens": 2, "total_tokens": 10}
_COMPLETION: Final = {
"id": "chatcmpl-otel",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4.1-mini-2025-04-14",
"service_tier": "default",
"system_fingerprint": "fp_otel",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}],
"usage": _USAGE,
}
_STREAM: Final = (
"".join(
f"data: {json.dumps(chunk)}\n\n"
for chunk in (
{
"id": "chatcmpl-otel",
"object": "chat.completion.chunk",
"created": 1700000000,
"model": "gpt-4.1-mini-2025-04-14",
"choices": [{"index": 0, "delta": {"role": "assistant", "content": "hello"}, "finish_reason": None}],
},
{
"id": "chatcmpl-otel",
"object": "chat.completion.chunk",
"created": 1700000000,
"model": "gpt-4.1-mini-2025-04-14",
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
"usage": _USAGE,
},
)
)
+ "data: [DONE]\n\n"
)
_LITELLM_REQUEST_ATTRIBUTES: Final = (
"gen_ai.request.model",
"gen_ai.system",
"gen_ai.request.temperature",
"llm.is_streaming",
"llm.user",
"gen_ai.response.id",
"gen_ai.response.model",
"gen_ai.usage.total_tokens",
"gen_ai.usage.output_tokens",
"gen_ai.usage.input_tokens",
)
_RAW_STREAMING_ATTRIBUTES: Final = (
"llm.openai.messages",
"llm.openai.temperature",
"llm.openai.user",
"llm.openai.extra_body",
"llm.openai.model",
)
_RAW_NON_STREAMING_ATTRIBUTES: Final = (
*_RAW_STREAMING_ATTRIBUTES,
"llm.openai.id",
"llm.openai.choices",
"llm.openai.created",
"llm.openai.object",
"llm.openai.service_tier",
"llm.openai.system_fingerprint",
"llm.openai.usage",
)
def _is_our_request(span: ReadableSpan) -> bool:
return span.name == "litellm_request" and (span.attributes or {}).get("llm.user") == _USER
def _trace_id(span: ReadableSpan) -> int:
assert span.context is not None
return span.context.trace_id
class _SignallingExporter(InMemorySpanExporter):
def __init__(self, loop: asyncio.AbstractEventLoop) -> None:
super().__init__()
self.loop: Final = loop
self.request_span_exported: Final = asyncio.Event()
def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult:
result: Final = super().export(spans)
if any(_is_our_request(span) for span in spans):
self.loop.call_soon_threadsafe(self.request_span_exported.set)
return result
@pytest.mark.asyncio
@pytest.mark.parametrize("streaming", [True, False])
async def test_otel_callback_emits_the_request_and_raw_provider_spans(
streaming: bool, monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.delenv("OTEL_SEMCONV_STABILITY_OPT_IN", raising=False)
exporter: Final = _SignallingExporter(asyncio.get_running_loop())
tracer_provider: Final = TracerProvider()
tracer_provider.add_span_processor(SimpleSpanProcessor(exporter))
monkeypatch.setattr(
litellm,
"callbacks",
[OpenTelemetry(config=OpenTelemetryConfig(exporter=exporter), tracer_provider=tracer_provider)],
)
respx_mock.post(_OPENAI_URL).mock(
return_value=httpx.Response(200, text=_STREAM, headers={"content-type": "text/event-stream"})
if streaming
else httpx.Response(200, json=_COMPLETION)
)
response: Final = await litellm.acompletion(
model="gpt-4.1-mini",
messages=[{"role": "user", "content": "hi"}],
temperature=0.1,
user=_USER,
stream=streaming,
api_key="sk-unit-test",
)
if streaming:
assert [chunk async for chunk in response]
await asyncio.wait_for(exporter.request_span_exported.wait(), timeout=10)
finished: Final = exporter.get_finished_spans()
request_span: Final = next(span for span in finished if _is_our_request(span))
ours: Final = tuple(span for span in finished if _trace_id(span) == _trace_id(request_span))
assert tuple(sorted(span.name for span in ours)) == _EXPECTED_SPAN_NAMES
spans: Final = {span.name: span for span in ours}
request_attributes: Final = spans["litellm_request"].attributes or {}
assert all(request_attributes.get(name) is not None for name in _LITELLM_REQUEST_ATTRIBUTES)
assert request_attributes["gen_ai.request.model"] == "gpt-4.1-mini"
assert request_attributes["gen_ai.system"] == "openai"
assert request_attributes["gen_ai.request.temperature"] == 0.1
assert request_attributes["llm.is_streaming"] == str(streaming)
assert request_attributes["llm.user"] == _USER
assert request_attributes["gen_ai.response.id"] == "chatcmpl-otel"
assert request_attributes["gen_ai.usage.input_tokens"] == _USAGE["prompt_tokens"]
assert request_attributes["gen_ai.usage.output_tokens"] == _USAGE["completion_tokens"]
assert request_attributes["gen_ai.usage.total_tokens"] == _USAGE["total_tokens"]
raw_attributes: Final = spans["raw_gen_ai_request"].attributes or {}
expected_raw: Final = _RAW_STREAMING_ATTRIBUTES if streaming else _RAW_NON_STREAMING_ATTRIBUTES
assert all(raw_attributes.get(name) is not None for name in expected_raw)

View file

@ -0,0 +1,289 @@
import itertools
import json
from typing import Final, TypedDict
from typing_extensions import ReadOnly
import httpx
import pytest
import respx
from pydantic import TypeAdapter
import litellm
import litellm.proxy.proxy_server as proxy_server
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import VectorStorePreCallHook
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.vector_stores.vector_store_registry import LiteLLM_ManagedVectorStore, VectorStoreRegistry
_KB_ID: Final = "T37J8R4WTM"
_KB_URL: Final = f"https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases/{_KB_ID}/retrieve"
_ANTHROPIC_URL: Final = "https://api.anthropic.com/v1/messages"
_OPENAI_URL: Final = "https://api.openai.com/v1/chat/completions"
_KB_TEXT: Final = "LiteLLM is a library that simplifies LLM API access"
_PREFIX: Final = VectorStorePreCallHook.CONTENT_PREFIX_STRING
_BODY: Final = TypeAdapter(dict[str, object])
_ANTHROPIC_MESSAGE: Final = {
"id": "msg_kb",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "LiteLLM simplifies LLM access."}],
"model": "claude-haiku-4-5-20251001",
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 100, "output_tokens": 50},
}
_OPENAI_COMPLETION: Final = {
"id": "chatcmpl-kb",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-5-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12},
}
_ANTHROPIC_STREAM: Final = (
'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_kb_stream","type":"message",'
'"role":"assistant","content":[],"model":"claude-haiku-4-5-20251001","stop_reason":null,"stop_sequence":null,'
'"usage":{"input_tokens":10,"output_tokens":1}}}\n\n'
'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}\n\n'
'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"LiteLLM"}}\n\n'
'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n'
'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},'
'"usage":{"output_tokens":2}}\n\n'
'event: message_stop\ndata: {"type":"message_stop"}\n\n'
)
class _ChatMessage(TypedDict):
role: ReadOnly[str]
content: ReadOnly[str]
@pytest.fixture(autouse=True)
def _knowledge_base(monkeypatch: pytest.MonkeyPatch, fake_provider_credentials: None) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setenv("AWS_REGION", "us-west-2")
monkeypatch.setenv("AWS_REGION_NAME", "us-west-2")
monkeypatch.setattr(
litellm,
"vector_store_registry",
VectorStoreRegistry(
vector_stores=[LiteLLM_ManagedVectorStore(vector_store_id=_KB_ID, custom_llm_provider="bedrock")]
),
raising=False,
)
def _kb_route(respx_mock: respx.MockRouter) -> respx.Route:
return respx_mock.post(_KB_URL).mock(
return_value=httpx.Response(
200,
json={"retrievalResults": [{"content": {"text": _KB_TEXT, "type": "TEXT"}, "score": 0.9, "metadata": {}}]},
)
)
def _sent_body(route: respx.Route) -> dict[str, object]:
return _BODY.validate_json(route.calls.last.request.content)
@pytest.mark.asyncio
async def test_completion_with_vector_store_ids_prepends_the_kb_context_block(respx_mock: respx.MockRouter) -> None:
_kb_route(respx_mock)
anthropic: Final = respx_mock.post(_ANTHROPIC_URL).mock(return_value=httpx.Response(200, json=_ANTHROPIC_MESSAGE))
await litellm.acompletion(
model="anthropic/claude-haiku-4-5-20251001",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=[_KB_ID],
)
messages: Final = TypeAdapter(tuple[dict[str, object], ...]).validate_python(_sent_body(anthropic)["messages"])
content: Final = TypeAdapter(tuple[dict[str, str], ...]).validate_python(messages[0]["content"])
assert anthropic.call_count == 1
assert [block["type"] for block in content] == ["text", "text"]
assert content[0]["text"] == f"{_PREFIX}{_KB_TEXT}\n\n"
assert content[1]["text"] == "what is litellm?"
@pytest.mark.asyncio
async def test_streaming_completion_carries_the_search_results_on_a_chunk_delta(respx_mock: respx.MockRouter) -> None:
_kb_route(respx_mock)
respx_mock.post(_ANTHROPIC_URL).mock(
return_value=httpx.Response(200, text=_ANTHROPIC_STREAM, headers={"content-type": "text/event-stream"})
)
response: Final = await litellm.acompletion(
model="anthropic/claude-haiku-4-5-20251001",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=[_KB_ID],
stream=True,
)
chunks: Final = tuple([chunk async for chunk in response])
choices: Final = tuple(itertools.chain.from_iterable(chunk.choices for chunk in chunks))
annotated: Final = tuple(
choice.delta.provider_specific_fields["search_results"]
for choice in choices
if choice.delta.provider_specific_fields and "search_results" in choice.delta.provider_specific_fields
)
assert len(chunks) > 0
assert len(annotated) >= 1
assert annotated[0][0]["object"] == "vector_store.search_results.page"
assert annotated[0][0]["data"][0]["content"][0]["text"] == _KB_TEXT
@pytest.mark.asyncio
async def test_file_search_filters_reach_the_kb_as_a_bedrock_equals_filter(respx_mock: respx.MockRouter) -> None:
kb: Final = _kb_route(respx_mock)
respx_mock.post(_ANTHROPIC_URL).mock(return_value=httpx.Response(200, json=_ANTHROPIC_MESSAGE))
response: Final = await litellm.acompletion(
model="anthropic/claude-haiku-4-5-20251001",
messages=[{"role": "user", "content": "what is litellm?"}],
max_tokens=10,
tools=[
{
"type": "file_search",
"vector_store_ids": [_KB_ID],
"filters": {"key": "user_id", "value": "fake-user-id", "operator": "eq"},
}
],
)
retrieval: Final = TypeAdapter(dict[str, dict[str, dict[str, object]]]).validate_python(
_sent_body(kb)["retrievalConfiguration"]
)
assert retrieval["vectorSearchConfiguration"]["filter"] == {"equals": {"key": "user_id", "value": "fake-user-id"}}
assert response.choices[0].message.content == "LiteLLM simplifies LLM access."
def _openai_messages(route: respx.Route) -> tuple[_ChatMessage, ...]:
return TypeAdapter(tuple[_ChatMessage, ...]).validate_python(_sent_body(route)["messages"])
@pytest.mark.asyncio
async def test_openai_request_with_vector_store_ids_leads_with_a_kb_context_user_message(
respx_mock: respx.MockRouter,
) -> None:
_kb_route(respx_mock)
openai: Final = respx_mock.post(_OPENAI_URL).mock(return_value=httpx.Response(200, json=_OPENAI_COMPLETION))
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=[_KB_ID],
)
assert _openai_messages(openai) == (
_ChatMessage(role="user", content=f"{_PREFIX}{_KB_TEXT}\n\n"),
_ChatMessage(role="user", content="what is litellm?"),
)
@pytest.mark.asyncio
async def test_a_managed_file_search_tool_is_resolved_locally_and_not_sent_upstream(
respx_mock: respx.MockRouter,
) -> None:
_kb_route(respx_mock)
openai: Final = respx_mock.post(_OPENAI_URL).mock(return_value=httpx.Response(200, json=_OPENAI_COMPLETION))
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "what is litellm?"}],
tools=[{"type": "file_search", "vector_store_ids": [_KB_ID]}],
)
assert _openai_messages(openai)[0] == _ChatMessage(role="user", content=f"{_PREFIX}{_KB_TEXT}\n\n")
assert "tools" not in _sent_body(openai)
@pytest.mark.asyncio
async def test_an_unknown_vector_store_tool_is_forwarded_while_the_known_one_is_resolved(
respx_mock: respx.MockRouter,
) -> None:
_kb_route(respx_mock)
openai: Final = respx_mock.post(_OPENAI_URL).mock(return_value=httpx.Response(200, json=_OPENAI_COMPLETION))
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "what is litellm?"}],
tools=[
{"type": "file_search", "vector_store_ids": [_KB_ID]},
{"type": "file_search", "vector_store_ids": ["unknownVS"]},
],
)
assert _openai_messages(openai)[0] == _ChatMessage(role="user", content=f"{_PREFIX}{_KB_TEXT}\n\n")
assert _sent_body(openai)["tools"] == [{"type": "file_search", "vector_store_ids": ["unknownVS"]}]
def _authorized_key() -> UserAPIKeyAuth:
return UserAPIKeyAuth(api_key="sk-kb-proxy", user_id="kb-proxy-user")
class _SearchResultsPage(TypedDict):
object: ReadOnly[str]
search_query: ReadOnly[str]
data: ReadOnly[list[dict[str, object]]]
class _ProviderFields(TypedDict):
search_results: ReadOnly[list[_SearchResultsPage]]
class _ProxyMessage(TypedDict):
role: ReadOnly[str]
content: ReadOnly[str]
provider_specific_fields: ReadOnly[_ProviderFields]
class _ProxyChoice(TypedDict):
message: ReadOnly[_ProxyMessage]
class _ProxyCompletion(TypedDict):
choices: ReadOnly[list[_ProxyChoice]]
@pytest.mark.asyncio
async def test_proxy_http_response_keeps_the_kb_search_results_in_provider_specific_fields(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
kb: Final = _kb_route(respx_mock)
upstream: Final = respx_mock.post(_OPENAI_URL).mock(return_value=httpx.Response(200, json=_OPENAI_COMPLETION))
monkeypatch.setattr(
proxy_server,
"llm_router",
litellm.Router(
model_list=[
{"model_name": "gpt-5-mini", "litellm_params": {"model": "openai/gpt-5-mini", "api_key": "sk-fixture"}}
],
num_retries=0,
),
)
monkeypatch.setitem(proxy_server.app.dependency_overrides, user_api_key_auth, _authorized_key)
async with httpx.AsyncClient(
transport=httpx.ASGITransport(proxy_server.app), base_url="http://kb-proxy.test"
) as client:
result: Final = await client.post(
"/v1/chat/completions",
json={
"model": "gpt-5-mini",
"messages": [{"role": "user", "content": "what is litellm?"}],
"vector_store_ids": [_KB_ID],
},
)
assert result.status_code == 200, result.text
assert kb.call_count == 1
assert upstream.call_count == 1
message: Final = TypeAdapter(_ProxyCompletion).validate_json(result.content)["choices"][0]["message"]
assert message["content"] == "ok"
pages: Final = message["provider_specific_fields"]["search_results"]
assert [page["object"] for page in pages] == ["vector_store.search_results.page"]
assert pages[0]["search_query"] == "what is litellm?"
assert pages[0]["data"]
assert _KB_TEXT in json.dumps(pages[0]["data"])

View file

@ -0,0 +1,210 @@
import asyncio
import json
from typing import Final, Literal, TypedDict
from typing_extensions import ReadOnly
import httpx
import pytest
import respx
from pydantic import TypeAdapter
import litellm
from litellm.integrations.custom_logger import CustomLogger
_MODEL: Final = "gpt-sized-search-unit"
_INPUT_COST: Final = 1e-06
_OUTPUT_COST: Final = 4e-06
_PER_QUERY: Final = {
"search_context_size_low": 0.011,
"search_context_size_medium": 0.022,
"search_context_size_high": 0.033,
}
_PROMPT_TOKENS: Final = 100
_COMPLETION_TOKENS: Final = 20
_USAGE: Final = {
"prompt_tokens": _PROMPT_TOKENS,
"completion_tokens": _COMPLETION_TOKENS,
"total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS,
}
_CHAT_RESPONSE: Final = {
"id": "chatcmpl-search",
"object": "chat.completion",
"created": 1700000000,
"model": _MODEL,
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "A positive story",
"annotations": [
{
"type": "url_citation",
"url_citation": {
"start_index": 0,
"end_index": 5,
"title": "news",
"url": "https://news.example/a",
},
}
],
},
"finish_reason": "stop",
}
],
"usage": _USAGE,
}
_RESPONSES_BODY: Final = {
"id": "resp_search",
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": _MODEL,
"output": [
{"type": "web_search_call", "id": "ws_search", "status": "completed"},
{
"type": "message",
"id": "msg_search",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "A positive story", "annotations": []}],
},
],
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [],
"usage": {
"input_tokens": _PROMPT_TOKENS,
"output_tokens": _COMPLETION_TOKENS,
"total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS,
},
}
_RESPONSES_STREAM: Final = "".join(
f"event: {event['type']}\ndata: {json.dumps(event)}\n\n"
for event in (
{"type": "response.created", "response": {**_RESPONSES_BODY, "status": "in_progress", "output": []}},
{"type": "response.completed", "response": _RESPONSES_BODY},
)
)
_ContextSize = Literal["search_context_size_low", "search_context_size_medium", "search_context_size_high"]
class _LoggedCost(TypedDict):
response_cost: ReadOnly[float]
prompt_tokens: ReadOnly[int]
completion_tokens: ReadOnly[int]
class _CostRecorder(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.payloads: Final[list[_LoggedCost]] = []
self.logged: Final = asyncio.Event()
async def async_log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
self.payloads.append(TypeAdapter(_LoggedCost).validate_python(kwargs["standard_logging_object"]))
self.logged.set()
@pytest.fixture
def recorder(monkeypatch: pytest.MonkeyPatch) -> _CostRecorder:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
entry: Final = {
"input_cost_per_token": _INPUT_COST,
"output_cost_per_token": _OUTPUT_COST,
"litellm_provider": "openai",
"mode": "chat",
"max_tokens": 4096,
"max_input_tokens": 4096,
"max_output_tokens": 4096,
"supports_web_search": True,
"search_context_cost_per_query": _PER_QUERY,
}
monkeypatch.setitem(litellm.model_cost, _MODEL, entry)
monkeypatch.setitem(litellm.model_cost, f"openai/{_MODEL}", entry)
cost_recorder: Final = _CostRecorder()
monkeypatch.setattr(litellm, "callbacks", [cost_recorder])
return cost_recorder
async def _logged_cost(recorder: _CostRecorder) -> _LoggedCost:
await asyncio.wait_for(recorder.logged.wait(), timeout=10)
return recorder.payloads[-1]
def _expected_cost(payload: _LoggedCost, size: _ContextSize) -> float:
return payload["prompt_tokens"] * _INPUT_COST + payload["completion_tokens"] * _OUTPUT_COST + _PER_QUERY[size]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("web_search_options", "size"),
[
(None, "search_context_size_medium"),
({"search_context_size": "low"}, "search_context_size_low"),
({"search_context_size": "high"}, "search_context_size_high"),
],
)
async def test_chat_web_search_logged_cost_adds_the_per_query_cost_for_the_context_size(
web_search_options: dict[str, str] | None,
size: _ContextSize,
recorder: _CostRecorder,
respx_mock: respx.MockRouter,
) -> None:
respx_mock.post("https://api.openai.com/v1/chat/completions").mock(
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
)
options: Final = {"web_search_options": web_search_options} if web_search_options is not None else {}
await litellm.acompletion(
model=f"openai/{_MODEL}",
messages=[{"role": "user", "content": "What was a positive news story from today?"}],
api_key="sk-unit-test",
**options,
)
payload: Final = await _logged_cost(recorder)
assert (payload["prompt_tokens"], payload["completion_tokens"]) == (_PROMPT_TOKENS, _COMPLETION_TOKENS)
assert payload["response_cost"] == pytest.approx(_expected_cost(payload, size), abs=1e-12)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("tools", "size", "stream"),
[
([{"type": "web_search_preview", "search_context_size": "low"}], "search_context_size_low", True),
([{"type": "web_search_preview", "search_context_size": "low"}], "search_context_size_low", False),
([{"type": "web_search_preview"}], "search_context_size_medium", True),
([{"type": "web_search_preview"}], "search_context_size_medium", False),
],
)
async def test_responses_web_search_logged_cost_adds_the_per_query_cost_for_the_context_size(
tools: list[dict[str, str]],
size: _ContextSize,
stream: bool,
recorder: _CostRecorder,
respx_mock: respx.MockRouter,
) -> None:
respx_mock.post("https://api.openai.com/v1/responses").mock(
return_value=httpx.Response(200, text=_RESPONSES_STREAM, headers={"content-type": "text/event-stream"})
if stream
else httpx.Response(200, json=_RESPONSES_BODY)
)
response: Final = await litellm.aresponses(
model=f"openai/{_MODEL}",
input=[{"role": "user", "content": "What was a positive news story from today?"}],
tools=tools,
stream=stream,
api_key="sk-unit-test",
)
if stream:
assert [event async for event in response]
payload: Final = await _logged_cost(recorder)
assert (payload["prompt_tokens"], payload["completion_tokens"]) == (_PROMPT_TOKENS, _COMPLETION_TOKENS)
assert payload["response_cost"] == pytest.approx(_expected_cost(payload, size), abs=1e-12)

View file

@ -0,0 +1,88 @@
import asyncio
from typing import Final, Literal, TypedDict
from typing_extensions import ReadOnly
import httpx
import pytest
import respx
from pydantic import TypeAdapter
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.router import Router
_MODERATIONS_URL: Final = "https://api.openai.com/v1/moderations"
_MODEL_GROUP: Final = "internal-moderation-model"
_INPUT: Final = "Hello, how are you?"
_CATEGORIES: Final = ("harassment", "hate", "self-harm", "sexual", "violence")
_MODERATION_RESPONSE: Final = {
"id": "modr-logging",
"model": "omni-moderation-latest",
"results": [
{
"flagged": False,
"categories": {name: False for name in _CATEGORIES},
"category_scores": {name: 0.001 for name in _CATEGORIES},
}
],
}
class _LoggedModeration(TypedDict):
call_type: ReadOnly[str]
status: ReadOnly[str]
custom_llm_provider: ReadOnly[str | None]
messages: ReadOnly[object]
response: ReadOnly[object]
model_group: ReadOnly[str | None]
class _ModerationRecorder(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.payloads: Final[list[_LoggedModeration]] = []
self.logged: Final = asyncio.Event()
async def async_log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
self.payloads.append(TypeAdapter(_LoggedModeration).validate_python(kwargs["standard_logging_object"]))
self.logged.set()
@pytest.mark.asyncio
@pytest.mark.parametrize("caller", ["default-model", "named-model", "router-group"])
async def test_moderation_call_is_logged_as_an_amoderation_standard_payload(
caller: Literal["default-model", "named-model", "router-group"],
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setenv("OPENAI_API_KEY", "sk-unit-test")
recorder: Final = _ModerationRecorder()
monkeypatch.setattr(litellm, "callbacks", [recorder])
respx_mock.post(_MODERATIONS_URL).mock(return_value=httpx.Response(200, json=_MODERATION_RESPONSE))
router: Final = Router(
model_list=[{"model_name": _MODEL_GROUP, "litellm_params": {"model": "openai/omni-moderation-latest"}}]
)
response: Final = (
await router.amoderation(input=_INPUT, model=_MODEL_GROUP)
if caller == "router-group"
else await litellm.amoderation(
input=_INPUT, model=None if caller == "default-model" else "omni-moderation-latest"
)
)
await asyncio.wait_for(recorder.logged.wait(), timeout=10)
payload: Final = recorder.payloads[-1]
assert payload["call_type"] == litellm.utils.CallTypes.amoderation.value
assert payload["status"] == "success"
assert payload["custom_llm_provider"] == litellm.LlmProviders.OPENAI.value
assert TypeAdapter(tuple[dict[str, str], ...]).validate_python(payload["messages"])[0]["content"] == _INPUT
assert dict(TypeAdapter(dict[str, object]).validate_python(payload["response"])) == response.model_dump()
if caller == "router-group":
assert payload["model_group"] == _MODEL_GROUP
else:
assert not payload["model_group"]

View file

@ -0,0 +1,138 @@
import asyncio
import json
from typing import Final, TypedDict
from typing_extensions import ReadOnly
import httpx
import pytest
import respx
from pydantic import TypeAdapter
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.redact_messages import REDACTED_BY_LITELLM
from litellm.types.utils import Usage
_OPENAI_URL: Final = "https://api.openai.com/v1/chat/completions"
_BODY: Final = TypeAdapter(dict[str, object])
_PROMPT_TOKENS: Final = 607
_COMPLETION_TOKENS: Final = 23
def _sse(chunks: tuple[dict[str, object], ...]) -> str:
return "".join(f"data: {json.dumps(chunk)}\n\n" for chunk in chunks) + "data: [DONE]\n\n"
def _chunk(delta: dict[str, str], finish_reason: str | None) -> dict[str, object]:
return {
"id": "chatcmpl-usage",
"object": "chat.completion.chunk",
"created": 1700000000,
"model": "gpt-5.5",
"choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}],
}
_STREAM: Final = _sse(
(
_chunk({"role": "assistant", "content": "I am"}, None),
_chunk({"content": " well"}, "stop"),
{
"id": "chatcmpl-usage",
"object": "chat.completion.chunk",
"created": 1700000000,
"model": "gpt-5.5",
"choices": [],
"usage": {
"prompt_tokens": _PROMPT_TOKENS,
"completion_tokens": _COMPLETION_TOKENS,
"total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS,
},
},
)
)
class _LoggedUsage(TypedDict):
prompt_tokens: ReadOnly[int]
completion_tokens: ReadOnly[int]
total_tokens: ReadOnly[int]
messages: ReadOnly[object]
class _UsageRecorder(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.payloads: Final[list[_LoggedUsage]] = []
self.logged: Final = asyncio.Event()
async def async_log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
self.payloads.append(TypeAdapter(_LoggedUsage).validate_python(kwargs["standard_logging_object"]))
self.logged.set()
async def _stream_and_record(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter, include_usage: bool
) -> tuple[Usage, _LoggedUsage, dict[str, object]]:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
recorder: Final = _UsageRecorder()
monkeypatch.setattr(litellm, "callbacks", [recorder])
route: Final = respx_mock.post(_OPENAI_URL).mock(
return_value=httpx.Response(200, text=_STREAM, headers={"content-type": "text/event-stream"})
)
stream_options: Final = {"stream_options": {"include_usage": True}} if include_usage else {}
response: Final = await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "Hello, how are you?" * 100}],
stream=True,
api_key="sk-unit-test",
**stream_options,
)
usages: Final = tuple([chunk.usage async for chunk in response if getattr(chunk, "usage", None) is not None])
await asyncio.wait_for(recorder.logged.wait(), timeout=10)
return usages[-1], recorder.payloads[-1], _BODY.validate_json(route.calls.last.request.content)
def _assert_logged_usage_matches(client_usage: Usage, payload: _LoggedUsage) -> None:
assert client_usage.prompt_tokens == _PROMPT_TOKENS
assert client_usage.completion_tokens == _COMPLETION_TOKENS
assert (payload["prompt_tokens"], payload["completion_tokens"], payload["total_tokens"]) == (
client_usage.prompt_tokens,
client_usage.completion_tokens,
client_usage.total_tokens,
)
@pytest.mark.asyncio
async def test_logged_stream_usage_equals_the_final_chunk_usage_with_include_usage(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
client_usage, payload, body = await _stream_and_record(monkeypatch, respx_mock, include_usage=True)
assert body["stream_options"] == {"include_usage": True}
_assert_logged_usage_matches(client_usage, payload)
@pytest.mark.asyncio
async def test_logged_stream_usage_equals_the_usage_chunk_without_stream_options(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
client_usage, payload, body = await _stream_and_record(monkeypatch, respx_mock, include_usage=False)
assert body["stream_options"] == {"include_usage": True}
_assert_logged_usage_matches(client_usage, payload)
@pytest.mark.asyncio
async def test_logged_stream_usage_survives_message_redaction(
monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter
) -> None:
monkeypatch.setattr(litellm, "turn_off_message_logging", True)
client_usage, payload, _ = await _stream_and_record(monkeypatch, respx_mock, include_usage=False)
_assert_logged_usage_matches(client_usage, payload)
assert payload["messages"] == [{"role": "user", "content": REDACTED_BY_LITELLM}]

View file

@ -0,0 +1,260 @@
import asyncio
import importlib
from collections.abc import Sequence
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import Final, NotRequired, TypedDict
from typing_extensions import ReadOnly
from unittest.mock import AsyncMock, patch
import httpx
import pytest
from opentelemetry.sdk.trace import ReadableSpan, TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor, SpanExportResult
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from opentelemetry.trace import StatusCode
from prisma.errors import ClientNotConnectedError
from pydantic import TypeAdapter
import litellm
from litellm._service_logger import ServiceTypes
from litellm.integrations.datadog.datadog import DataDogLogger
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
from litellm.proxy.db.log_db_metrics import log_db_metrics
from litellm.proxy.db.prisma_client import _PrismaDrainTracker, _TrackedPrismaEngine
from litellm.proxy.proxy_server import proxy_logging_obj
from litellm.integrations.prometheus_services import PrometheusServicesLogger
from prometheus_client import REGISTRY
class _ServiceEvent(TypedDict):
service: ReadOnly[str]
call_type: ReadOnly[str]
duration: ReadOnly[float]
is_error: ReadOnly[bool]
error: ReadOnly[str | None]
event_metadata: ReadOnly[dict[str, str] | None]
table_name: NotRequired[ReadOnly[str]]
class _ServiceSpanExporter(InMemorySpanExporter):
def __init__(self) -> None:
super().__init__()
self.service_span_exported: Final = asyncio.Event()
def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult:
result: Final = super().export(spans)
self.service_span_exported.set()
return result
class _Rig:
def __init__(self, exporter: _ServiceSpanExporter, provider: TracerProvider, datadog: DataDogLogger) -> None:
self.exporter: Final = exporter
self.provider: Final = provider
self.datadog: Final = datadog
def service_spans(self) -> tuple[ReadableSpan, ...]:
return tuple(span for span in self.exporter.get_finished_spans() if span.name != "request")
def events(self) -> tuple[_ServiceEvent, ...]:
adapter: Final = TypeAdapter(_ServiceEvent)
return tuple(adapter.validate_json(entry["message"]) for entry in self.datadog.log_queue)
def _discard_periodic_flush(coroutine: object) -> None:
close: Final = getattr(coroutine, "close")
close()
@pytest.fixture
def rig(monkeypatch: pytest.MonkeyPatch) -> _Rig:
monkeypatch.setenv("DD_API_KEY", "test_api_key")
monkeypatch.setenv("DD_SITE", "test.datadoghq.com")
monkeypatch.setattr(litellm, "datadog_params", None)
exporter: Final = _ServiceSpanExporter()
provider: Final = TracerProvider()
provider.add_span_processor(SimpleSpanProcessor(exporter))
otel: Final = OpenTelemetry(config=OpenTelemetryConfig(exporter=exporter), tracer_provider=provider)
with patch("asyncio.create_task", side_effect=_discard_periodic_flush):
datadog: Final = DataDogLogger()
monkeypatch.setattr(litellm, "service_callback", [otel, datadog, "prometheus_system"])
monkeypatch.setattr(proxy_logging_obj.service_logging_obj, "dd_logger", datadog, raising=False)
monkeypatch.setattr(
proxy_logging_obj.service_logging_obj, "prometheusServicesLogger", PrometheusServicesLogger(), raising=False
)
return _Rig(exporter, provider, datadog)
async def _run_prisma_query() -> None:
engine: Final = _TrackedPrismaEngine(SimpleNamespace(query=AsyncMock(return_value={})), _PrismaDrainTracker())
await engine.query("{}", tx_id=None)
@log_db_metrics
async def read_spend_rows(**kwargs: object) -> str:
await _run_prisma_query()
return "success"
def _logged_db_latency() -> tuple[float, float]:
labels: Final = {ServiceTypes.DB.value: ServiceTypes.DB.value}
total: Final = REGISTRY.get_sample_value("litellm_postgres_latency_sum", labels)
count: Final = REGISTRY.get_sample_value("litellm_postgres_latency_count", labels)
return (total or 0.0, count or 0.0)
def _ns(moment: datetime) -> int:
return int(moment.timestamp() * 1e9)
_DB_CALL_START: Final = datetime(2026, 1, 1, 12, 0, 0)
_DB_CALL_DURATION: Final = timedelta(milliseconds=250)
class _ScriptedClock:
def __init__(self, *moments: datetime) -> None:
self._moments: Final = iter(moments)
def now(self) -> datetime:
return next(self._moments)
@pytest.fixture
def scripted_db_clock(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(
importlib.import_module("litellm.proxy.db.log_db_metrics"),
"datetime",
_ScriptedClock(_DB_CALL_START, _DB_CALL_START + _DB_CALL_DURATION),
)
@pytest.mark.asyncio
async def test_a_db_success_is_reported_on_the_parent_span_with_its_duration_and_times(
rig: _Rig, scripted_db_clock: None
) -> None:
parent: Final = rig.provider.get_tracer("db-test").start_span("request")
latency_before: Final = _logged_db_latency()
result: Final = await read_spend_rows(parent_otel_span=parent)
await asyncio.wait_for(rig.exporter.service_span_exported.wait(), timeout=10)
assert result == "success"
spans: Final = rig.service_spans()
assert len(spans) == 1
span: Final = spans[0]
assert span.parent is not None
assert span.parent.span_id == parent.get_span_context().span_id
assert span.attributes is not None
assert span.attributes["service"] == ServiceTypes.DB.value
assert span.attributes["call_type"] == "read_spend_rows"
assert span.status.status_code == StatusCode.OK
assert span.start_time == _ns(_DB_CALL_START)
assert span.end_time == _ns(_DB_CALL_START + _DB_CALL_DURATION)
latency_after: Final = _logged_db_latency()
assert latency_after[1] - latency_before[1] == 1
assert latency_after[0] - latency_before[0] == pytest.approx(_DB_CALL_DURATION.total_seconds())
assert rig.events() == ()
@pytest.mark.asyncio
async def test_db_event_metadata_names_only_the_table_and_never_the_raw_kwargs(rig: _Rig) -> None:
parent: Final = rig.provider.get_tracer("db-test").start_span("request")
await read_spend_rows(
parent_otel_span=parent,
table_name="LiteLLM_SpendLogs",
token="sk-secret-should-not-leak",
prisma_client=object(),
)
await asyncio.wait_for(rig.exporter.service_span_exported.wait(), timeout=10)
span_attributes: Final = rig.service_spans()[0].attributes
assert span_attributes is not None
assert span_attributes["table_name"] == "LiteLLM_SpendLogs"
assert not {"token", "prisma_client", "parent_otel_span"} & set(span_attributes)
assert all("sk-secret-should-not-leak" not in str(value) for value in span_attributes.values())
@pytest.mark.asyncio
async def test_the_logged_db_duration_is_the_span_wall_clock_of_the_wrapped_call(
rig: _Rig, scripted_db_clock: None
) -> None:
parent: Final = rig.provider.get_tracer("db-test").start_span("request")
latency_before: Final = _logged_db_latency()
await read_spend_rows(parent_otel_span=parent)
await asyncio.wait_for(rig.exporter.service_span_exported.wait(), timeout=10)
span: Final = rig.service_spans()[0]
assert span.start_time is not None and span.end_time is not None
latency_after: Final = _logged_db_latency()
logged_duration: Final = latency_after[0] - latency_before[0]
assert latency_after[1] - latency_before[1] == 1
assert logged_duration == pytest.approx((span.end_time - span.start_time) / 1e9, rel=1e-3, abs=2e-6)
assert logged_duration == pytest.approx(_DB_CALL_DURATION.total_seconds())
@log_db_metrics
async def disconnected_read(**kwargs: object) -> str:
raise ClientNotConnectedError()
@pytest.mark.asyncio
async def test_a_prisma_error_is_reported_as_a_db_failure_and_reraised(rig: _Rig) -> None:
parent: Final = rig.provider.get_tracer("db-test").start_span("request")
with pytest.raises(ClientNotConnectedError, match="Client is not connected to the query engine"):
await disconnected_read(parent_otel_span=parent)
spans: Final = rig.service_spans()
assert len(spans) == 1
assert spans[0].parent is not None
assert spans[0].parent.span_id == parent.get_span_context().span_id
assert spans[0].status.status_code == StatusCode.ERROR
assert spans[0].attributes is not None
assert spans[0].attributes["call_type"] == "disconnected_read"
assert spans[0].attributes["service"] == ServiceTypes.DB.value
assert "Client is not connected" in str(spans[0].attributes["error"])
events: Final = rig.events()
assert len(events) == 1
assert events[0]["is_error"] is True
assert events[0]["call_type"] == "disconnected_read"
assert isinstance(events[0]["duration"], float)
assert "Client is not connected" in (events[0]["error"] or "")
@pytest.mark.asyncio
@pytest.mark.parametrize(
("error", "is_db_error"),
[
(ValueError("Generic error"), False),
(KeyError("Missing key"), False),
(TypeError("Type error"), False),
(httpx.ConnectError("Failed to connect"), True),
(httpx.TimeoutException("Request timed out"), True),
(ClientNotConnectedError(), True),
],
)
async def test_only_db_errors_are_reported_as_db_failures(rig: _Rig, error: Exception, is_db_error: bool) -> None:
parent: Final = rig.provider.get_tracer("db-test").start_span("request")
@log_db_metrics
async def failing_read(**kwargs: object) -> str:
raise error
with pytest.raises(type(error)):
await failing_read(parent_otel_span=parent)
spans: Final = rig.service_spans()
events: Final = rig.events()
if is_db_error:
assert [span.status.status_code for span in spans] == [StatusCode.ERROR]
assert [(event["service"], event["call_type"], event["is_error"]) for event in events] == [
(ServiceTypes.DB.value, "failing_read", True)
]
assert isinstance(events[0]["duration"], float)
else:
assert spans == ()
assert events == ()

View file

@ -0,0 +1,509 @@
import asyncio
import inspect
import json
from collections import Counter
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime
from typing import Final, Literal, NamedTuple, TypedDict
from typing_extensions import ReadOnly
import httpx
import pytest
import respx
from pydantic import TypeAdapter
import litellm
from litellm.caching.caching import Cache
from litellm.integrations.custom_logger import CustomLogger
from litellm.router import Router
_PRIMARY: Final = "https://hooks-primary.openai.azure.com"
_FALLBACK: Final = "https://hooks-fallback.openai.azure.com"
_API_VERSION: Final = "2024-10-21"
_MESSAGES: Final = [{"role": "user", "content": "Hi - i'm openai"}]
_USAGE: Final = {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}
_COMPLETION: Final = {
"id": "chatcmpl-hooks",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4.1-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hello"}, "finish_reason": "stop"}],
"usage": _USAGE,
}
_EMBEDDING_VECTOR: Final = [0.1, 0.2, 0.3]
_EMBEDDING: Final = {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": _EMBEDDING_VECTOR}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 3, "total_tokens": 3},
}
_STREAM: Final = (
"".join(
f"data: {json.dumps(chunk)}\n\n"
for chunk in (
{
"id": "chatcmpl-hooks",
"object": "chat.completion.chunk",
"created": 1700000000,
"model": "gpt-4.1-mini",
"choices": [{"index": 0, "delta": {"role": "assistant", "content": "hel"}, "finish_reason": None}],
},
{
"id": "chatcmpl-hooks",
"object": "chat.completion.chunk",
"created": 1700000000,
"model": "gpt-4.1-mini",
"choices": [{"index": 0, "delta": {"content": "lo"}, "finish_reason": "stop"}],
},
)
)
+ "data: [DONE]\n\n"
)
_AUTH_ERROR: Final = httpx.Response(
401,
json={
"error": {"message": "Incorrect API key provided", "type": "invalid_request_error", "code": "invalid_api_key"}
},
)
_OUR_MODEL_GROUPS: Final = frozenset({"hooks-group", "primary-group", "fallback-group"})
_State = Literal[
"sync_pre_api_call",
"post_api_call",
"async_stream",
"sync_success",
"async_success",
"sync_failure",
"async_failure",
]
class _HookEvent(NamedTuple):
state: _State
model: object
kwargs: Mapping[str, object]
response: object
def _router_context_problems(kwargs: Mapping[str, object]) -> tuple[str, ...]:
litellm_params: Final = kwargs.get("litellm_params")
if not isinstance(litellm_params, dict):
return ("litellm_params",)
metadata: Final = litellm_params.get("metadata")
model_info: Final = litellm_params.get("model_info")
checks: Final = {
"metadata": isinstance(metadata, dict),
"model_group": isinstance(metadata, dict) and isinstance(metadata.get("model_group"), str),
"deployment": isinstance(metadata, dict) and isinstance(metadata.get("deployment"), str),
"model_info": isinstance(model_info, dict),
"model_info id": isinstance(model_info, dict) and isinstance(model_info.get("id"), str),
"proxy_server_request": isinstance(litellm_params.get("proxy_server_request"), (str, type(None))),
"preset_cache_key": isinstance(litellm_params.get("preset_cache_key"), (str, type(None))),
"stream_response": isinstance(litellm_params.get("stream_response"), dict),
}
return tuple(name for name, ok in checks.items() if not ok)
def _request_problems(kwargs: Mapping[str, object]) -> tuple[str, ...]:
checks: Final = {
"model": isinstance(kwargs.get("model"), str),
"messages": isinstance(kwargs.get("messages"), list),
"optional_params": isinstance(kwargs.get("optional_params"), dict),
"start_time": isinstance(kwargs.get("start_time"), (datetime, type(None))),
"stream": isinstance(kwargs.get("stream"), bool),
"user": isinstance(kwargs.get("user"), (str, type(None))),
}
return (*(name for name, ok in checks.items() if not ok), *_router_context_problems(kwargs))
def _call_detail_problems(kwargs: Mapping[str, object]) -> tuple[str, ...]:
original_response: Final = kwargs.get("original_response")
checks: Final = {
"input": isinstance(kwargs.get("input"), (list, dict, str)),
"api_key": isinstance(kwargs.get("api_key"), (str, type(None))),
"original_response": isinstance(original_response, (str, litellm.CustomStreamWrapper, type(None)))
or inspect.iscoroutine(original_response)
or inspect.isasyncgen(original_response),
"additional_args": isinstance(kwargs.get("additional_args"), (dict, type(None))),
"log_event_type": isinstance(kwargs.get("log_event_type"), str),
}
return tuple(name for name, ok in checks.items() if not ok)
def _is_from_this_test(kwargs: Mapping[str, object]) -> bool:
litellm_params: Final = kwargs.get("litellm_params")
metadata: Final = litellm_params.get("metadata") if isinstance(litellm_params, dict) else None
return isinstance(metadata, dict) and metadata.get("model_group") in _OUR_MODEL_GROUPS
class _HookRecorder(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.events: Final[list[_HookEvent]] = []
self.errors: Final[list[str]] = []
self.loop: asyncio.AbstractEventLoop | None = None
self.waiters: Final[list[tuple[Callable[[Sequence[_State]], bool], asyncio.Event]]] = []
@property
def states(self) -> list[_State]:
return [event.state for event in self.events]
def _record(
self, state: _State, model: object, kwargs: Mapping[str, object], response: object, problems: Sequence[str]
) -> None:
if not _is_from_this_test(kwargs):
return
self.errors.extend(f"{state}: {problem}" for problem in problems)
self.events.append(_HookEvent(state, model, kwargs, response))
if self.loop is not None:
self.loop.call_soon_threadsafe(self._notify)
def _notify(self) -> None:
for predicate, event in self.waiters:
if predicate(tuple(self.states)):
event.set()
async def until(self, predicate: Callable[[Sequence[_State]], bool]) -> tuple[_State, ...]:
self.loop = asyncio.get_running_loop()
if not predicate(tuple(self.states)):
event: Final = asyncio.Event()
self.waiters.append((predicate, event))
await asyncio.wait_for(event.wait(), timeout=10)
return tuple(self.states)
def log_pre_api_call(self, model: object, messages: object, kwargs: Mapping[str, object]) -> None:
problems: Final = (
*(("model",) if not isinstance(model, str) else ()),
*(("messages",) if not isinstance(messages, list) else ()),
*_request_problems(kwargs),
)
self._record("sync_pre_api_call", model, kwargs, messages, problems)
def log_post_api_call(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
problems: Final = (
*(("start_time",) if not isinstance(start_time, datetime) else ()),
*(("end_time",) if end_time is not None else ()),
*(("response_obj",) if response_obj is not None else ()),
*_request_problems(kwargs),
*_call_detail_problems(kwargs),
)
self._record("post_api_call", kwargs.get("model"), kwargs, response_obj, problems)
def log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
self._record("sync_success", kwargs.get("model"), kwargs, response_obj, ())
def log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
self._record("sync_failure", kwargs.get("model"), kwargs, response_obj, ())
async def async_log_stream_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
self._record("async_stream", kwargs.get("model"), kwargs, response_obj, ())
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
problems: Final = (
*(("times",) if not (isinstance(start_time, datetime) and isinstance(end_time, datetime)) else ()),
*(
("response_obj",)
if not isinstance(response_obj, (litellm.ModelResponse, litellm.EmbeddingResponse))
else ()
),
*(("cache_hit",) if not isinstance(kwargs.get("cache_hit"), (bool, type(None))) else ()),
*_request_problems(kwargs),
*_call_detail_problems(kwargs),
)
self._record("async_success", kwargs.get("model"), kwargs, response_obj, problems)
async def async_log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
problems: Final = (
*(("times",) if not (isinstance(start_time, datetime) and isinstance(end_time, datetime)) else ()),
*(("response_obj",) if response_obj is not None else ()),
*(("exception",) if not isinstance(kwargs.get("exception"), Exception) else ()),
*_request_problems(kwargs),
*_call_detail_problems(kwargs),
)
self._record("async_failure", kwargs.get("model"), kwargs, response_obj, problems)
def _settled(terminal: int, posts: int) -> Callable[[Sequence[_State]], bool]:
def satisfied(states: Sequence[_State]) -> bool:
terminals: Final = sum(state in ("async_success", "async_failure") for state in states)
return terminals >= terminal and states.count("post_api_call") >= posts
return satisfied
@pytest.fixture
def recorder(monkeypatch: pytest.MonkeyPatch) -> _HookRecorder:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
hook_recorder: Final = _HookRecorder()
monkeypatch.setattr(litellm, "callbacks", [hook_recorder])
return hook_recorder
def _router(model: str, api_base: str = _PRIMARY) -> Router:
return Router(
model_list=[
{
"model_name": "hooks-group",
"litellm_params": {
"model": model,
"api_key": "sk-unit-test",
"api_base": api_base,
"api_version": _API_VERSION,
},
"model_info": {"base_model": model},
}
],
num_retries=0,
)
class _RouterMetadata(TypedDict):
model_group: ReadOnly[str]
deployment: ReadOnly[str]
class _RouterModelInfo(TypedDict):
id: ReadOnly[str]
class _RouterParams(TypedDict):
metadata: ReadOnly[_RouterMetadata]
model_info: ReadOnly[_RouterModelInfo]
def _router_params(event: _HookEvent) -> _RouterParams:
return TypeAdapter(_RouterParams).validate_python(event.kwargs["litellm_params"])
def _model_group(event: _HookEvent) -> str:
return _router_params(event)["metadata"]["model_group"]
def _model_id(event: _HookEvent) -> str:
return _router_params(event)["model_info"]["id"]
def _of_state(recorder: _HookRecorder, state: _State) -> list[_HookEvent]:
return [event for event in recorder.events if event.state == state]
def _completion(event: _HookEvent) -> litellm.ModelResponse:
assert isinstance(event.response, litellm.ModelResponse)
return event.response
@pytest.mark.asyncio
async def test_router_chat_success_streaming_and_failure_fire_the_hooks_in_order(
recorder: _HookRecorder, respx_mock: respx.MockRouter
) -> None:
route: Final = respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/gpt-4.1-mini/chat/completions")
router: Final = _router("azure/gpt-4.1-mini")
route.mock(return_value=httpx.Response(200, json=_COMPLETION))
await router.acompletion(model="hooks-group", messages=_MESSAGES)
await recorder.until(_settled(1, posts=1))
assert recorder.states == ["sync_pre_api_call", "post_api_call", "async_success"]
pre, post, success = recorder.events
assert pre.model == "gpt-4.1-mini"
assert pre.response == _MESSAGES
assert {_model_group(event) for event in recorder.events} == {"hooks-group"}
assert len({_model_id(event) for event in recorder.events}) == 1
assert post.kwargs["messages"] == _MESSAGES
assert success.kwargs["stream"] is False
assert _completion(success).choices[0].message.content == "hello"
assert _completion(success).usage.total_tokens == _USAGE["total_tokens"]
route.mock(return_value=httpx.Response(200, text=_STREAM, headers={"content-type": "text/event-stream"}))
stream: Final = await router.acompletion(model="hooks-group", messages=_MESSAGES, stream=True)
assert "".join([chunk.choices[0].delta.content or "" async for chunk in stream]) == "hello"
await recorder.until(_settled(2, posts=2))
streamed: Final = recorder.events[3:]
assert sorted(event.state for event in streamed[:2]) == ["post_api_call", "sync_pre_api_call"]
assert [event.state for event in streamed[2:]] == ["async_success"]
assert all(event.kwargs["stream"] is True for event in streamed)
assert _completion(streamed[2]).choices[0].message.content == "hello"
route.mock(return_value=_AUTH_ERROR)
with pytest.raises(litellm.AuthenticationError):
await router.acompletion(model="hooks-group", messages=_MESSAGES)
await recorder.until(_settled(3, posts=3))
failed: Final = recorder.events[6:]
assert [event.state for event in failed] == ["sync_pre_api_call", "post_api_call", "async_failure"]
assert failed[2].response is None
assert isinstance(failed[2].kwargs["exception"], litellm.AuthenticationError)
assert recorder.errors == []
@pytest.mark.asyncio
async def test_router_embedding_success_and_failure_fire_the_hooks_in_order(
recorder: _HookRecorder, respx_mock: respx.MockRouter
) -> None:
route: Final = respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/text-embedding-3-small/embeddings")
router: Final = _router("azure/text-embedding-3-small")
route.mock(return_value=httpx.Response(200, json=_EMBEDDING))
await router.aembedding(model="hooks-group", input=["hello"])
await recorder.until(_settled(1, posts=1))
assert recorder.states == ["sync_pre_api_call", "post_api_call", "async_success"]
assert {event.model for event in recorder.events} == {"text-embedding-3-small"}
assert {_model_group(event) for event in recorder.events} == {"hooks-group"}
embedding: Final = recorder.events[2].response
assert isinstance(embedding, litellm.EmbeddingResponse)
assert embedding.model_dump()["data"][0]["embedding"] == _EMBEDDING_VECTOR
assert embedding.usage.prompt_tokens == 3
route.mock(return_value=_AUTH_ERROR)
with pytest.raises(litellm.AuthenticationError):
await router.aembedding(model="hooks-group", input=["hello"])
await recorder.until(_settled(2, posts=2))
assert recorder.states[3:] == ["sync_pre_api_call", "post_api_call", "async_failure"]
assert recorder.events[5].response is None
assert isinstance(recorder.events[5].kwargs["exception"], litellm.AuthenticationError)
assert recorder.errors == []
@pytest.mark.asyncio
async def test_router_fallback_fires_failure_then_success_hooks(
recorder: _HookRecorder, respx_mock: respx.MockRouter
) -> None:
respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/gpt-4.1-mini/chat/completions").mock(
return_value=_AUTH_ERROR
)
fallback: Final = respx_mock.post(
url__startswith=f"{_FALLBACK}/openai/deployments/gpt-4.1-mini/chat/completions"
).mock(return_value=httpx.Response(200, json=_COMPLETION))
router: Final = Router(
model_list=[
{
"model_name": "primary-group",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": "my-bad-key",
"api_base": _PRIMARY,
"api_version": _API_VERSION,
},
},
{
"model_name": "fallback-group",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": "sk-unit-test",
"api_base": _FALLBACK,
"api_version": _API_VERSION,
},
},
],
fallbacks=[{"primary-group": ["fallback-group"]}],
num_retries=0,
)
await router.acompletion(model="primary-group", messages=_MESSAGES)
await recorder.until(_settled(2, posts=2))
assert fallback.call_count == 1
assert recorder.states == [
"sync_pre_api_call",
"post_api_call",
"async_failure",
"sync_pre_api_call",
"post_api_call",
"async_success",
]
assert [_model_group(event) for event in recorder.events] == ["primary-group"] * 3 + ["fallback-group"] * 3
assert _model_id(recorder.events[0]) != _model_id(recorder.events[3])
assert isinstance(recorder.events[2].kwargs["exception"], litellm.AuthenticationError)
assert _completion(recorder.events[5]).choices[0].message.content == "hello"
assert recorder.errors == []
@pytest.mark.asyncio
async def test_router_completion_cache_hit_fires_a_second_success_hook(
recorder: _HookRecorder, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "cache", Cache())
route: Final = respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/gpt-4.1-mini/chat/completions").mock(
return_value=httpx.Response(200, json=_COMPLETION)
)
router: Final = _router("azure/gpt-4.1-mini")
await router.acompletion(model="hooks-group", messages=_MESSAGES, caching=True)
await recorder.until(_settled(1, posts=1))
await router.acompletion(model="hooks-group", messages=_MESSAGES, caching=True)
await recorder.until(_settled(2, posts=1))
assert route.call_count == 1
assert recorder.states == ["sync_pre_api_call", "post_api_call", "async_success", "async_success"]
first, second = _of_state(recorder, "async_success")
assert first.kwargs.get("cache_hit") is not True
assert second.kwargs.get("cache_hit") is True
assert _completion(second).choices[0].message.content == "hello"
assert recorder.errors == []
@pytest.mark.asyncio
async def test_router_streaming_cache_hit_still_fires_the_success_hook(
recorder: _HookRecorder, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "cache", Cache())
route: Final = respx_mock.post(url__startswith=f"{_PRIMARY}/openai/deployments/gpt-4.1-mini/chat/completions").mock(
return_value=httpx.Response(200, text=_STREAM, headers={"content-type": "text/event-stream"})
)
router: Final = _router("azure/gpt-4.1-mini")
first: Final = await router.acompletion(model="hooks-group", messages=_MESSAGES, stream=True, caching=True)
first_text: Final = "".join([chunk.choices[0].delta.content or "" async for chunk in first])
await recorder.until(_settled(1, posts=1))
states_after_first: Final = len(recorder.states)
second: Final = await router.acompletion(model="hooks-group", messages=_MESSAGES, stream=True, caching=True)
second_text: Final = "".join([chunk.choices[0].delta.content or "" async for chunk in second])
await recorder.until(_settled(2, posts=1))
assert route.call_count == 1
assert first_text == second_text == "hello"
assert sorted(recorder.states[:2]) == ["post_api_call", "sync_pre_api_call"]
assert recorder.states[2:states_after_first] == ["async_success"]
assert recorder.states[states_after_first:] == ["async_success"]
first_success, second_success = _of_state(recorder, "async_success")
assert first_success.kwargs.get("cache_hit") is not True
assert second_success.kwargs.get("cache_hit") is True
assert _completion(second_success).choices[0].message.content == "hello"
assert recorder.errors == []
@pytest.mark.asyncio
async def test_router_embedding_cache_hit_fires_a_second_success_hook(
recorder: _HookRecorder, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "cache", Cache())
route: Final = respx_mock.post(
url__startswith=f"{_PRIMARY}/openai/deployments/text-embedding-3-small/embeddings"
).mock(return_value=httpx.Response(200, json=_EMBEDDING))
router: Final = _router("azure/text-embedding-3-small")
await router.aembedding(model="hooks-group", input=["hello"], caching=True)
await recorder.until(_settled(1, posts=1))
await router.aembedding(model="hooks-group", input=["hello"], caching=True)
await recorder.until(_settled(2, posts=1))
assert route.call_count == 1
assert recorder.states == ["sync_pre_api_call", "post_api_call", "async_success", "async_success"]
first, second = _of_state(recorder, "async_success")
assert first.kwargs.get("cache_hit") is not True
assert second.kwargs.get("cache_hit") is True
assert isinstance(second.response, litellm.EmbeddingResponse)
assert second.response.model_dump()["data"][0]["embedding"] == _EMBEDDING_VECTOR
assert recorder.errors == []