diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index e6f76b67c3c..2d67c22f0aa 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -11,6 +11,7 @@ from fastapi import HTTPException import litellm from httpx import Response as HttpxResponse +from litellm.proxy.spend_tracking.compression_savings import HEADROOM_GUARDRAIL_PROVIDER from typing_extensions import TypeGuard from litellm._logging import verbose_proxy_logger @@ -487,7 +488,7 @@ class HeadroomGuardrail(CustomGuardrail): guardrail_json_response=stats, request_data=request_data, guardrail_status="success", - guardrail_provider="headroom", + guardrail_provider=HEADROOM_GUARDRAIL_PROVIDER, start_time=start_time, end_time=end_time, duration=end_time - start_time, diff --git a/litellm/proxy/spend_tracking/compression_savings.py b/litellm/proxy/spend_tracking/compression_savings.py index f3b57758b48..34241715756 100644 --- a/litellm/proxy/spend_tracking/compression_savings.py +++ b/litellm/proxy/spend_tracking/compression_savings.py @@ -10,11 +10,11 @@ HEADROOM_GUARDRAIL_PROVIDER = "headroom" def _saved_tokens_or_zero(value: object) -> int: - if isinstance(value, bool) or not isinstance(value, int): + if isinstance(value, bool) or not isinstance(value, (int, float)): return 0 if value < 0: return 0 - return value + return int(value) def _tokens_saved_from_stats(stats: object) -> int: @@ -32,9 +32,10 @@ def _headroom_entry_saved_tokens(entry: object) -> int: def _headroom_saved_tokens(guardrail_information: object) -> int: - if not isinstance(guardrail_information, list): + entries = [guardrail_information] if isinstance(guardrail_information, Mapping) else guardrail_information + if not isinstance(entries, list): return 0 - return sum(_headroom_entry_saved_tokens(entry) for entry in guardrail_information) + return sum(_headroom_entry_saved_tokens(entry) for entry in entries) def extract_compression_saved_tokens(metadata: Mapping[str, object]) -> int: @@ -52,6 +53,8 @@ def extract_compression_saved_tokens(metadata: Mapping[str, object]) -> int: different stages (guardrail pre-call vs deployment pre-call), so when both fire on one request their measured savings are independent and additive; summing them never double-counts. Malformed or missing values contribute 0. + A bare dict ``guardrail_information`` is treated as a single entry, matching + the spend-log redactor's normalization of that legacy shape. """ return _tokens_saved_from_stats(metadata.get("compression_savings")) + _headroom_saved_tokens( metadata.get("guardrail_information") diff --git a/tests/test_litellm/proxy/spend_tracking/test_compression_savings.py b/tests/test_litellm/proxy/spend_tracking/test_compression_savings.py index dc8f071d77e..77e7b4b55c3 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_compression_savings.py +++ b/tests/test_litellm/proxy/spend_tracking/test_compression_savings.py @@ -69,7 +69,6 @@ def test_non_headroom_guardrail_entries_ignored(): {}, {"tokens_saved": None}, {"tokens_saved": "7000"}, - {"tokens_saved": 12.5}, {"tokens_saved": True}, {"tokens_saved": -5}, {"tokens_before": 100, "tokens_after": 50}, @@ -108,3 +107,25 @@ def test_valid_headroom_entry_survives_alongside_malformed_ones(): ] } assert extract_compression_saved_tokens(metadata) == 600 + + +def test_float_tokens_saved_counts_as_int(): + entry = {"guardrail_provider": "headroom", "guardrail_response": {"tokens_saved": 600.0}} + assert extract_compression_saved_tokens({"guardrail_information": [entry]}) == 600 + assert extract_compression_saved_tokens({"compression_savings": {"tokens_saved": 7000.0}}) == 7000 + assert extract_compression_saved_tokens({"compression_savings": {"tokens_saved": 12.5}}) == 12 + + +def test_bare_dict_guardrail_information_counts_as_single_entry(): + assert extract_compression_saved_tokens({"guardrail_information": HEADROOM_ENTRY}) == 600 + + +def test_headroom_writer_and_reader_share_provider_slug(): + from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import ( + HEADROOM_GUARDRAIL_PROVIDER as writer_slug, + ) + from litellm.proxy.spend_tracking.compression_savings import ( + HEADROOM_GUARDRAIL_PROVIDER as reader_slug, + ) + + assert writer_slug == reader_slug == "headroom"