From 9bd89290cb0f394009607b753b374bcab1b458cd Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Fri, 24 Jul 2026 16:44:00 -0700 Subject: [PATCH] fix(guardrails): derive tokens_saved when Headroom compression service omits it The savings readers (extract_compression_saved_tokens, feeding compression_saved_tokens on the daily spend tables) key exclusively on tokens_saved in the guardrail_response stats, but the Headroom guardrail builds those stats as a filtered pass-through of the compression service response and the live service omits tokens_saved. Every compressed request recorded 0 saved tokens on the Cost Optimization dashboard. Derive tokens_saved = tokens_before - tokens_after when the key is absent and both operands are numeric; a service-sent value still wins. The two sibling writers (compresr, native compression interception) already derive it the same way. Co-Authored-By: Claude Fable 5 --- .../guardrail_hooks/headroom/headroom.py | 13 ++ .../guardrail_hooks/test_headroom.py | 113 ++++++++++++++++++ 2 files changed, 126 insertions(+) diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 2d67c22f0aa..ef7b52cdd4e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -410,6 +410,19 @@ class HeadroomGuardrail(CustomGuardrail): ) if key in body } + tokens_before = stats.get("tokens_before") + tokens_after = stats.get("tokens_after") + if ( + "tokens_saved" not in stats + and isinstance(tokens_before, (int, float)) + and not isinstance(tokens_before, bool) + and isinstance(tokens_after, (int, float)) + and not isinstance(tokens_after, bool) + ): + # Spend tracking (extract_compression_saved_tokens) reads only + # tokens_saved, which the live compression service omits; derive it + # so savings are counted, but let a service-sent value win. + stats["tokens_saved"] = tokens_before - tokens_after return filtered, True, stats async def _call_retrieve(self, hash_value: str, query: str | None = None) -> str: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py index 7f412c008ca..85f04d82360 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_headroom.py @@ -11,6 +11,9 @@ Tests cover: - /v1/compress non-2xx surfaces as httpx.HTTPStatusError (raise_for_status), not a status_code check on the returned response -- both are handled - unreachable_fallback="fail_open" forwards the request uncompressed instead of raising +- tokens_saved is derived from tokens_before/tokens_after when the compression + service omits it, passed through verbatim when present, and skipped (without + breaking compression) when the token counts are not numeric - CCR: headroom_retrieve tool injected when compressed messages contain hashes - CCR: async_should_run_agentic_loop returns True when response has headroom_retrieve tool calls - CCR: async_build_agentic_loop_plan calls retrieve endpoint and builds follow-up messages @@ -32,6 +35,9 @@ from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import ( has_headroom_retrieve_tool, HEADROOM_RETRIEVE_TOOL_NAME, ) +from litellm.proxy.spend_tracking.compression_savings import ( + extract_compression_saved_tokens, +) from litellm.types.utils import GenericGuardrailAPIInputs FAKE_API_BASE = "https://headroom.example.com" @@ -139,6 +145,113 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages( assert result.get("structured_messages") == COMPRESSED_MESSAGES +def _recorded_guardrail_response(request_data: dict) -> dict: + entries = request_data["metadata"]["standard_logging_guardrail_information"] + assert len(entries) == 1 + return entries[0]["guardrail_response"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_derives_tokens_saved_when_service_omits_it( + guardrail: HeadroomGuardrail, +): + inputs = GenericGuardrailAPIInputs( + texts=["A" * 5000], + structured_messages=ORIGINAL_MESSAGES, + ) + # _make_compress_response omits tokens_saved, matching the live service. + mock_response = _make_compress_response(COMPRESSED_MESSAGES) + request_data: dict = {"model": "gpt-4o"} + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + stats = _recorded_guardrail_response(request_data) + assert stats["tokens_saved"] == 900 + + # Spend tracking reads the entry under the spend-log metadata key. + entry = request_data["metadata"]["standard_logging_guardrail_information"][0] + assert extract_compression_saved_tokens({"guardrail_information": [entry]}) == 900 + + +@pytest.mark.asyncio +async def test_apply_guardrail_passes_through_service_sent_tokens_saved( + guardrail: HeadroomGuardrail, +): + inputs = GenericGuardrailAPIInputs( + texts=["A" * 5000], + structured_messages=ORIGINAL_MESSAGES, + ) + mock_response = _make_compress_response(COMPRESSED_MESSAGES) + # Deliberately different from tokens_before - tokens_after (900): the + # service-sent value must win over the derived one. + mock_response.json.return_value["tokens_saved"] = 123 + request_data: dict = {"model": "gpt-4o"} + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + assert _recorded_guardrail_response(request_data)["tokens_saved"] == 123 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "tokens_before, tokens_after", + [ + ("1000", "100"), + (True, False), + (None, None), + ], +) +async def test_apply_guardrail_skips_derivation_for_non_numeric_token_counts( + guardrail: HeadroomGuardrail, + tokens_before, + tokens_after, +): + inputs = GenericGuardrailAPIInputs( + texts=["A" * 5000], + structured_messages=ORIGINAL_MESSAGES, + ) + mock_response = _make_compress_response(COMPRESSED_MESSAGES) + mock_response.json.return_value["tokens_before"] = tokens_before + mock_response.json.return_value["tokens_after"] = tokens_after + request_data: dict = {"model": "gpt-4o"} + + with patch.object( + guardrail.async_handler, + "post", + new_callable=AsyncMock, + return_value=mock_response, + ): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + assert "tokens_saved" not in _recorded_guardrail_response(request_data) + # Compression itself is unaffected by the skipped derivation. + assert result.get("structured_messages") == COMPRESSED_MESSAGES + + @pytest.mark.asyncio async def test_apply_guardrail_injects_retrieve_tool_when_hashes_present( guardrail: HeadroomGuardrail,