From f09408dd37d47de075284a18db6a50fbd69e601a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 8 Oct 2026 15:55:57 -0700 Subject: [PATCH] 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 --- .../test_cli_sso_login_ui_disabled.py | 58 +- .../test_model_access_allow_lists.py | 171 ++++ .../test_config_declared_model_behaviours.py | 162 ++++ .../test_guardrail_attachment.py | 153 +++ .../test_prometheus_request_metrics.py | 274 ++++++ .../spend/test_budget_limit_envelopes.py | 172 ++++ tests/logging_callback_tests/test_alerting.py | 338 ------- .../test_bedrock_knowledgebase_hook.py | 612 ------------ .../test_built_in_tools_cost_tracking.py | 160 --- .../test_custom_callback_router.py | 754 --------------- tests/logging_callback_tests/test_datadog.py | 226 ----- .../test_log_db_redis_services.py | 212 +--- .../test_moderations_api_logging.py | 100 -- .../test_otel_logging.py | 133 --- .../test_token_counting.py | 157 --- tests/otel_tests/test_e2e_budgeting.py | 557 ----------- tests/otel_tests/test_e2e_model_access.py | 304 ------ tests/otel_tests/test_guardrails.py | 191 ---- .../otel_tests/test_key_logging_callbacks.py | 70 -- tests/otel_tests/test_model_info.py | 29 - tests/otel_tests/test_moderations.py | 22 - tests/otel_tests/test_prometheus.py | 911 ------------------ tests/otel_tests/test_team_tag_routing.py | 65 -- .../test_slack_alerting_delivery.py | 344 +++++++ .../unit/integrations/datadog/test_datadog.py | 119 ++- .../test_opentelemetry_request_spans.py | 157 +++ .../test_bedrock_kb_context_offline.py | 289 ++++++ .../test_web_search_logged_cost.py | 210 ++++ .../test_moderation_standard_logging.py | 88 ++ .../test_stream_usage_logging.py | 138 +++ .../db/test_log_db_metrics_service_spans.py | 260 +++++ .../test_router_callback_hook_sequence.py | 509 ++++++++++ 32 files changed, 3103 insertions(+), 4842 deletions(-) create mode 100644 tests/integration/authorization/test_model_access_allow_lists.py create mode 100644 tests/integration/configuration/test_config_declared_model_behaviours.py create mode 100644 tests/integration/observability/test_guardrail_attachment.py create mode 100644 tests/integration/observability/test_prometheus_request_metrics.py create mode 100644 tests/integration/spend/test_budget_limit_envelopes.py delete mode 100644 tests/logging_callback_tests/test_alerting.py delete mode 100644 tests/logging_callback_tests/test_built_in_tools_cost_tracking.py delete mode 100644 tests/logging_callback_tests/test_custom_callback_router.py delete mode 100644 tests/logging_callback_tests/test_datadog.py delete mode 100644 tests/logging_callback_tests/test_moderations_api_logging.py delete mode 100644 tests/logging_callback_tests/test_otel_logging.py delete mode 100644 tests/logging_callback_tests/test_token_counting.py delete mode 100644 tests/otel_tests/test_e2e_budgeting.py delete mode 100644 tests/otel_tests/test_e2e_model_access.py delete mode 100644 tests/otel_tests/test_key_logging_callbacks.py delete mode 100644 tests/otel_tests/test_model_info.py delete mode 100644 tests/otel_tests/test_prometheus.py delete mode 100644 tests/otel_tests/test_team_tag_routing.py create mode 100644 tests/unit/integrations/SlackAlerting/test_slack_alerting_delivery.py create mode 100644 tests/unit/integrations/test_opentelemetry_request_spans.py create mode 100644 tests/unit/integrations/vector_store_integrations/test_bedrock_kb_context_offline.py create mode 100644 tests/unit/litellm_core_utils/llm_cost_calc/test_web_search_logged_cost.py create mode 100644 tests/unit/litellm_core_utils/test_moderation_standard_logging.py create mode 100644 tests/unit/litellm_core_utils/test_stream_usage_logging.py create mode 100644 tests/unit/proxy/db/test_log_db_metrics_service_spans.py create mode 100644 tests/unit/test_router/test_router_callback_hook_sequence.py diff --git a/tests/integration/authorization/test_cli_sso_login_ui_disabled.py b/tests/integration/authorization/test_cli_sso_login_ui_disabled.py index 126fe27bef2..a5e9b2752ea 100644 --- a/tests/integration/authorization/test_cli_sso_login_ui_disabled.py +++ b/tests/integration/authorization/test_cli_sso_login_ui_disabled.py @@ -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() == () diff --git a/tests/integration/authorization/test_model_access_allow_lists.py b/tests/integration/authorization/test_model_access_allow_lists.py new file mode 100644 index 00000000000..cdefc8c8a3e --- /dev/null +++ b/tests/integration/authorization/test_model_access_allow_lists.py @@ -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") diff --git a/tests/integration/configuration/test_config_declared_model_behaviours.py b/tests/integration/configuration/test_config_declared_model_behaviours.py new file mode 100644 index 00000000000..6b1f25c35a3 --- /dev/null +++ b/tests/integration/configuration/test_config_declared_model_behaviours.py @@ -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 diff --git a/tests/integration/observability/test_guardrail_attachment.py b/tests/integration/observability/test_guardrail_attachment.py new file mode 100644 index 00000000000..f291b3877f3 --- /dev/null +++ b/tests/integration/observability/test_guardrail_attachment.py @@ -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() == () diff --git a/tests/integration/observability/test_prometheus_request_metrics.py b/tests/integration/observability/test_prometheus_request_metrics.py new file mode 100644 index 00000000000..042e4350640 --- /dev/null +++ b/tests/integration/observability/test_prometheus_request_metrics.py @@ -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)] diff --git a/tests/integration/spend/test_budget_limit_envelopes.py b/tests/integration/spend/test_budget_limit_envelopes.py new file mode 100644 index 00000000000..6ac2a11e657 --- /dev/null +++ b/tests/integration/spend/test_budget_limit_envelopes.py @@ -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 diff --git a/tests/logging_callback_tests/test_alerting.py b/tests/logging_callback_tests/test_alerting.py deleted file mode 100644 index 691bb58b998..00000000000 --- a/tests/logging_callback_tests/test_alerting.py +++ /dev/null @@ -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() diff --git a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py index bf6294baa00..8c17ee132d8 100644 --- a/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py +++ b/tests/logging_callback_tests/test_bedrock_knowledgebase_hook.py @@ -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") diff --git a/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py b/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py deleted file mode 100644 index 2d3933c3dad..00000000000 --- a/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py +++ /dev/null @@ -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) diff --git a/tests/logging_callback_tests/test_custom_callback_router.py b/tests/logging_callback_tests/test_custom_callback_router.py deleted file mode 100644 index 4a12f8d536d..00000000000 --- a/tests/logging_callback_tests/test_custom_callback_router.py +++ /dev/null @@ -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 diff --git a/tests/logging_callback_tests/test_datadog.py b/tests/logging_callback_tests/test_datadog.py deleted file mode 100644 index 35b46a58fd4..00000000000 --- a/tests/logging_callback_tests/test_datadog.py +++ /dev/null @@ -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)}") diff --git a/tests/logging_callback_tests/test_log_db_redis_services.py b/tests/logging_callback_tests/test_log_db_redis_services.py index ba7b333e097..aad404a5525 100644 --- a/tests/logging_callback_tests/test_log_db_redis_services.py +++ b/tests/logging_callback_tests/test_log_db_redis_services.py @@ -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 diff --git a/tests/logging_callback_tests/test_moderations_api_logging.py b/tests/logging_callback_tests/test_moderations_api_logging.py deleted file mode 100644 index a2a356d3665..00000000000 --- a/tests/logging_callback_tests/test_moderations_api_logging.py +++ /dev/null @@ -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 diff --git a/tests/logging_callback_tests/test_otel_logging.py b/tests/logging_callback_tests/test_otel_logging.py deleted file mode 100644 index 6274916ddb2..00000000000 --- a/tests/logging_callback_tests/test_otel_logging.py +++ /dev/null @@ -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" diff --git a/tests/logging_callback_tests/test_token_counting.py b/tests/logging_callback_tests/test_token_counting.py deleted file mode 100644 index 513d2242fdf..00000000000 --- a/tests/logging_callback_tests/test_token_counting.py +++ /dev/null @@ -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 - - diff --git a/tests/otel_tests/test_e2e_budgeting.py b/tests/otel_tests/test_e2e_budgeting.py deleted file mode 100644 index f180403e115..00000000000 --- a/tests/otel_tests/test_e2e_budgeting.py +++ /dev/null @@ -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 diff --git a/tests/otel_tests/test_e2e_model_access.py b/tests/otel_tests/test_e2e_model_access.py deleted file mode 100644 index bdd6b597f68..00000000000 --- a/tests/otel_tests/test_e2e_model_access.py +++ /dev/null @@ -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"] diff --git a/tests/otel_tests/test_guardrails.py b/tests/otel_tests/test_guardrails.py index 36f2d019a34..ae82bb8908c 100644 --- a/tests/otel_tests/test_guardrails.py +++ b/tests/otel_tests/test_guardrails.py @@ -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" diff --git a/tests/otel_tests/test_key_logging_callbacks.py b/tests/otel_tests/test_key_logging_callbacks.py deleted file mode 100644 index 1736831eb69..00000000000 --- a/tests/otel_tests/test_key_logging_callbacks.py +++ /dev/null @@ -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"] - ) diff --git a/tests/otel_tests/test_model_info.py b/tests/otel_tests/test_model_info.py deleted file mode 100644 index 66a81eeee51..00000000000 --- a/tests/otel_tests/test_model_info.py +++ /dev/null @@ -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" diff --git a/tests/otel_tests/test_moderations.py b/tests/otel_tests/test_moderations.py index a9c73c93500..1a67827aa61 100644 --- a/tests/otel_tests/test_moderations.py +++ b/tests/otel_tests/test_moderations.py @@ -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(): """ diff --git a/tests/otel_tests/test_prometheus.py b/tests/otel_tests/test_prometheus.py deleted file mode 100644 index 7224334105d..00000000000 --- a/tests/otel_tests/test_prometheus.py +++ /dev/null @@ -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}" diff --git a/tests/otel_tests/test_team_tag_routing.py b/tests/otel_tests/test_team_tag_routing.py deleted file mode 100644 index b818aa0cea4..00000000000 --- a/tests/otel_tests/test_team_tag_routing.py +++ /dev/null @@ -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 diff --git a/tests/unit/integrations/SlackAlerting/test_slack_alerting_delivery.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting_delivery.py new file mode 100644 index 00000000000..c8fb5d30376 --- /dev/null +++ b/tests/unit/integrations/SlackAlerting/test_slack_alerting_delivery.py @@ -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] diff --git a/tests/unit/integrations/datadog/test_datadog.py b/tests/unit/integrations/datadog/test_datadog.py index e86d83ba467..eec0022c1db 100644 --- a/tests/unit/integrations/datadog/test_datadog.py +++ b/tests/unit/integrations/datadog/test_datadog.py @@ -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) diff --git a/tests/unit/integrations/test_opentelemetry_request_spans.py b/tests/unit/integrations/test_opentelemetry_request_spans.py new file mode 100644 index 00000000000..9ec40c00cb4 --- /dev/null +++ b/tests/unit/integrations/test_opentelemetry_request_spans.py @@ -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) diff --git a/tests/unit/integrations/vector_store_integrations/test_bedrock_kb_context_offline.py b/tests/unit/integrations/vector_store_integrations/test_bedrock_kb_context_offline.py new file mode 100644 index 00000000000..20d3a310bdb --- /dev/null +++ b/tests/unit/integrations/vector_store_integrations/test_bedrock_kb_context_offline.py @@ -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"]) diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_web_search_logged_cost.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_web_search_logged_cost.py new file mode 100644 index 00000000000..1448318b07d --- /dev/null +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_web_search_logged_cost.py @@ -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) diff --git a/tests/unit/litellm_core_utils/test_moderation_standard_logging.py b/tests/unit/litellm_core_utils/test_moderation_standard_logging.py new file mode 100644 index 00000000000..1f8c722785f --- /dev/null +++ b/tests/unit/litellm_core_utils/test_moderation_standard_logging.py @@ -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"] diff --git a/tests/unit/litellm_core_utils/test_stream_usage_logging.py b/tests/unit/litellm_core_utils/test_stream_usage_logging.py new file mode 100644 index 00000000000..a108af09dd1 --- /dev/null +++ b/tests/unit/litellm_core_utils/test_stream_usage_logging.py @@ -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}] diff --git a/tests/unit/proxy/db/test_log_db_metrics_service_spans.py b/tests/unit/proxy/db/test_log_db_metrics_service_spans.py new file mode 100644 index 00000000000..a7fd830b948 --- /dev/null +++ b/tests/unit/proxy/db/test_log_db_metrics_service_spans.py @@ -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 == () diff --git a/tests/unit/test_router/test_router_callback_hook_sequence.py b/tests/unit/test_router/test_router_callback_hook_sequence.py new file mode 100644 index 00000000000..a7c41964950 --- /dev/null +++ b/tests/unit/test_router/test_router_callback_hook_sequence.py @@ -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 == []