mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
8bae65e0b6
commit
f09408dd37
32 changed files with 3103 additions and 4842 deletions
|
|
@ -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() == ()
|
||||
|
|
|
|||
171
tests/integration/authorization/test_model_access_allow_lists.py
Normal file
171
tests/integration/authorization/test_model_access_allow_lists.py
Normal 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")
|
||||
|
|
@ -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
|
||||
153
tests/integration/observability/test_guardrail_attachment.py
Normal file
153
tests/integration/observability/test_guardrail_attachment.py
Normal 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() == ()
|
||||
|
|
@ -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)]
|
||||
172
tests/integration/spend/test_budget_limit_envelopes.py
Normal file
172
tests/integration/spend/test_budget_limit_envelopes.py
Normal 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
|
||||
|
|
@ -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()
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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)}")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
@ -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
|
||||
|
|
@ -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"]
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
)
|
||||
|
|
@ -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"
|
||||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
@ -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
|
||||
|
|
@ -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]
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
157
tests/unit/integrations/test_opentelemetry_request_spans.py
Normal file
157
tests/unit/integrations/test_opentelemetry_request_spans.py
Normal 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)
|
||||
|
|
@ -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"])
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"]
|
||||
138
tests/unit/litellm_core_utils/test_stream_usage_logging.py
Normal file
138
tests/unit/litellm_core_utils/test_stream_usage_logging.py
Normal 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}]
|
||||
260
tests/unit/proxy/db/test_log_db_metrics_service_spans.py
Normal file
260
tests/unit/proxy/db/test_log_db_metrics_service_spans.py
Normal 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 == ()
|
||||
509
tests/unit/test_router/test_router_callback_hook_sequence.py
Normal file
509
tests/unit/test_router/test_router_callback_hook_sequence.py
Normal 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 == []
|
||||
Loading…
Add table
Reference in a new issue