headroom/headroom/proxy/savings_attribution.py
Tejas Chopra 3145242645
Unify savings attribution across stats, perf, metrics, and dashboard (#2976)
## Summary

Adds a small provider-neutral savings attribution seam. Named sources
can attach realized or projected token/USD deltas to a request without
changing headline arithmetic or introducing private-package inventory
into OSS.

Also fixes the Anthropic buffered lifecycle so normal successful
responses run response hooks, applies stream-safety filtering, includes
tool savings in per-model perf totals, and surfaces the same breakdown
in request logs, `/stats`, `headroom perf`, Prometheus, OTEL, and the
dashboard.

## Why

Request-local savings were split between canonical token deltas,
process-global extension counters, and tool-only tags. This made correct
headline totals possible while losing attribution in perf, recent
requests, metrics, and the dashboard. Normal Anthropic responses also
skipped response hooks unless CCR ran.

## Validation

- 74 focused tests passed: turn hooks, OpenAI hook lifecycle, outcome
funnel, perf formats, and tool-search repair
- Ruff passes on all changed Python files
- Existing compression-observability suite: 11 passed; 2 tokenizer-cache
tests require network access to fetch the tiktoken vocabulary

## Compatibility

No named private packages or private inventory are encoded in OSS.
Existing hooks remain source-compatible because all new TurnContext
fields are optional.

---------

Co-authored-by: Tejas Chopra <tejas@Tejass-MacBook-Pro.local>
2026-08-13 17:13:23 -07:00

110 lines
3.4 KiB
Python

"""Bounded attribution for savings that do not have a built-in metric."""
from __future__ import annotations
import base64
import json
import re
from collections.abc import MutableMapping
from typing import Any
SAVINGS_ATTRIBUTION_TAG = "_headroom_savings_attribution"
_NAME_RE = re.compile(r"[^a-z0-9_.-]+")
MAX_SOURCES = 32
_SCOPE_KEY = "headroom_savings_attribution"
def _source_name(value: object) -> str:
name = _NAME_RE.sub("_", str(value or "other").strip().lower()).strip("_.-")
return (name or "other")[:64]
def _ledger(tags: MutableMapping[str, Any]) -> list[dict[str, Any]]:
current = tags.get(SAVINGS_ATTRIBUTION_TAG)
if isinstance(current, list):
return current
current = []
tags[SAVINGS_ATTRIBUTION_TAG] = current
return current
def bind_scope(tags: MutableMapping[str, Any], scope: MutableMapping[str, Any]) -> None:
"""Share one ledger between ASGI middleware and the request handler."""
state = scope.setdefault("state", {})
ledger = state.get(_SCOPE_KEY)
if not isinstance(ledger, list):
ledger = []
state[_SCOPE_KEY] = ledger
tags[SAVINGS_ATTRIBUTION_TAG] = ledger
def record_scope_savings(scope: MutableMapping[str, Any], source: object, **values: Any) -> None:
state = scope.setdefault("state", {})
ledger = state.get(_SCOPE_KEY)
if not isinstance(ledger, list):
ledger = []
state[_SCOPE_KEY] = ledger
record_savings({SAVINGS_ATTRIBUTION_TAG: ledger}, source, **values)
def record_savings(
tags: MutableMapping[str, Any],
source: object,
*,
tokens: int = 0,
usd: float = 0.0,
realized: bool = True,
estimated: bool = False,
details: dict[str, Any] | None = None,
) -> None:
"""Attribute savings to a source; this never changes headline totals."""
ledger = _ledger(tags)
if len(ledger) >= MAX_SOURCES:
return
item: dict[str, Any] = {
"source": _source_name(source),
"realized": bool(realized),
"estimated": bool(estimated),
"tokens": max(0, int(tokens or 0)),
"usd": round(float(usd or 0.0), 12),
}
if details:
item["details"] = {
_source_name(key): value
for key, value in list(details.items())[:12]
if isinstance(value, (str, int, float, bool)) or value is None
}
ledger.append(item)
def from_tags(tags: MutableMapping[str, Any] | None) -> list[dict[str, Any]]:
raw = (tags or {}).get(SAVINGS_ATTRIBUTION_TAG)
if not isinstance(raw, list):
return []
return [dict(item) for item in raw[:MAX_SOURCES] if isinstance(item, dict)]
def public_tags(tags: MutableMapping[str, Any] | None) -> dict[str, Any]:
return {key: value for key, value in (tags or {}).items() if key != SAVINGS_ATTRIBUTION_TAG}
def encode(items: list[dict[str, Any]]) -> str:
if not items:
return "none"
payload = json.dumps(items, separators=(",", ":"), sort_keys=True).encode()
return base64.urlsafe_b64encode(payload).decode().rstrip("=")
def decode(value: str) -> list[dict[str, Any]]:
if not value or value == "none":
return []
try:
padded = value + "=" * (-len(value) % 4)
decoded = json.loads(base64.urlsafe_b64decode(padded).decode())
except Exception:
return []
return (
[dict(item) for item in decoded if isinstance(item, dict)]
if isinstance(decoded, list)
else []
)