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 <noreply@anthropic.com>
This commit is contained in:
Tin Chi Lo 2026-07-24 16:44:00 -07:00
parent 3eaf7b1c0a
commit 9bd89290cb
2 changed files with 126 additions and 0 deletions

View file

@ -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:

View file

@ -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,