diff --git a/headroom/proxy/handlers/gemini.py b/headroom/proxy/handlers/gemini.py index 0d176d129..199a79fd1 100644 --- a/headroom/proxy/handlers/gemini.py +++ b/headroom/proxy/handlers/gemini.py @@ -22,6 +22,7 @@ from headroom.proxy.auth_mode import classify_client from headroom.proxy.compression_decision import CompressionDecision from headroom.proxy.helpers import COMPRESSION_TIMEOUT_SECONDS, extract_tags from headroom.proxy.outcome import RequestOutcome +from headroom.proxy.token_counting import gemini_output_tokens logger = logging.getLogger("headroom.proxy") @@ -462,7 +463,9 @@ class GeminiHandlerMixin: # output_tokens) would then raise TypeError on the non-error # path. Mirrors the streaming _usage_int guard. total_input_tokens = _usage_int(usage.get("promptTokenCount")) - output_tokens = _usage_int(usage.get("candidatesTokenCount")) + output_tokens = gemini_output_tokens( + usage + ) # includes thinking tokens (2.5-family) cache_read_tokens = _usage_int(usage.get("cachedContentTokenCount")) except (json.JSONDecodeError, ValueError, KeyError, TypeError, AttributeError): pass @@ -714,7 +717,9 @@ class GeminiHandlerMixin: if usage.get("promptTokenCount") is None else usage["promptTokenCount"] ) - output_tokens = _usage_int(usage.get("candidatesTokenCount")) + output_tokens = gemini_output_tokens( + usage + ) # includes thinking tokens (2.5-family) # Gemini returns cachedContentTokenCount for context-cached tokens # These are charged at 10-25% of the input price depending on model cache_read_tokens = _usage_int(usage.get("cachedContentTokenCount")) diff --git a/headroom/proxy/handlers/openai.py b/headroom/proxy/handlers/openai.py index f380ab4c9..62b46c03e 100644 --- a/headroom/proxy/handlers/openai.py +++ b/headroom/proxy/handlers/openai.py @@ -75,6 +75,7 @@ from headroom.proxy.passthrough import ( custom_base_passthrough_telemetry as _custom_base_passthrough_telemetry, ) from headroom.proxy.project_context import classify_project, set_current_project +from headroom.proxy.token_counting import gemini_output_tokens logger = logging.getLogger("headroom.proxy") @@ -331,7 +332,7 @@ def _passthrough_usage_from_json(payload: Any) -> dict[str, int]: if isinstance(usage_meta, dict): return { "input_tokens": _usage_int(usage_meta.get("promptTokenCount")), - "output_tokens": _usage_int(usage_meta.get("candidatesTokenCount")), + "output_tokens": gemini_output_tokens(usage_meta), "cache_read_input_tokens": _usage_int(usage_meta.get("cachedContentTokenCount")), } diff --git a/headroom/proxy/handlers/streaming.py b/headroom/proxy/handlers/streaming.py index 3b864f9dd..7ce93a9fb 100644 --- a/headroom/proxy/handlers/streaming.py +++ b/headroom/proxy/handlers/streaming.py @@ -18,6 +18,7 @@ from headroom.proxy.helpers import ( jitter_delay_ms, retry_after_ms, ) +from headroom.proxy.token_counting import gemini_output_tokens if TYPE_CHECKING: from fastapi.responses import Response, StreamingResponse @@ -224,7 +225,7 @@ class StreamingMixin: usage_meta = data.get("usageMetadata") if usage_meta: usage["input_tokens"] = usage_meta.get("promptTokenCount", 0) - usage["output_tokens"] = usage_meta.get("candidatesTokenCount", 0) + usage["output_tokens"] = gemini_output_tokens(usage_meta) # Gemini also has cachedContentTokenCount for context caching usage["cache_read_input_tokens"] = usage_meta.get( "cachedContentTokenCount", 0 @@ -342,7 +343,7 @@ class StreamingMixin: usage_meta = data.get("usageMetadata") if usage_meta: usage_found["input_tokens"] = usage_meta.get("promptTokenCount", 0) - usage_found["output_tokens"] = usage_meta.get("candidatesTokenCount", 0) + usage_found["output_tokens"] = gemini_output_tokens(usage_meta) usage_found["cache_read_input_tokens"] = usage_meta.get( "cachedContentTokenCount", 0 ) diff --git a/headroom/proxy/token_counting.py b/headroom/proxy/token_counting.py index fb9ca107b..0e5c10017 100644 --- a/headroom/proxy/token_counting.py +++ b/headroom/proxy/token_counting.py @@ -77,3 +77,33 @@ async def count_texts_offloaded(owner: Any, model: Any, texts: Any) -> tuple[Any return await _count_offloaded( owner, model, lambda counter: sum(counter.count_text(text) for text in text_list) ) + + +def gemini_output_tokens(usage_meta: dict[str, Any]) -> int: + """Output-token count for a Gemini ``usageMetadata``, including thinking tokens. + + Gemini reports ``candidatesTokenCount`` sometimes inclusive of the + ``thoughtsTokenCount`` (2.5-family reasoning) and sometimes exclusive of it. + When ``promptTokenCount + candidatesTokenCount != totalTokenCount`` the + thinking tokens are a separate bucket and must be added, or the output cost + (billed at the output rate) is undercounted. Mirrors litellm's + ``is_candidate_token_count_inclusive`` rule. Robust to missing/null fields. + """ + + def _int(value: Any) -> int: + try: + return max(int(value), 0) + except (TypeError, ValueError): + return 0 + + candidates = _int(usage_meta.get("candidatesTokenCount")) + thoughts = _int(usage_meta.get("thoughtsTokenCount")) + if thoughts <= 0: + return candidates + prompt = _int(usage_meta.get("promptTokenCount")) + total = _int(usage_meta.get("totalTokenCount")) + # Inclusive iff prompt + candidates already equals total; otherwise the + # thinking tokens are a separate bucket that belongs in the output count. + if prompt + candidates == total: + return candidates + return candidates + thoughts diff --git a/tests/test_proxy_handler_helpers.py b/tests/test_proxy_handler_helpers.py index bb86076aa..7845396b0 100644 --- a/tests/test_proxy_handler_helpers.py +++ b/tests/test_proxy_handler_helpers.py @@ -381,6 +381,51 @@ def test_passthrough_usage_normalizes_vertex_usage_metadata() -> None: } +def test_gemini_output_tokens_includes_thinking_when_exclusive() -> None: + """Gemini 2.5 thinking: when prompt + candidates != total, thoughtsTokenCount + is a separate output bucket and must be added, or output cost undercounts.""" + from headroom.proxy.token_counting import gemini_output_tokens + + exclusive = { + "promptTokenCount": 1000, + "candidatesTokenCount": 200, + "thoughtsTokenCount": 500, + "totalTokenCount": 1700, + } + assert gemini_output_tokens(exclusive) == 700 # 200 visible + 500 thinking + + # Inclusive: candidatesTokenCount already covers thoughts (prompt+cand==total). + inclusive = { + "promptTokenCount": 1000, + "candidatesTokenCount": 700, + "thoughtsTokenCount": 500, + "totalTokenCount": 1700, + } + assert gemini_output_tokens(inclusive) == 700 + + # No thinking tokens: just the candidates count (common non-2.5 case). + assert gemini_output_tokens({"candidatesTokenCount": 42, "totalTokenCount": 100}) == 42 + # Robust to empty / missing fields. + assert gemini_output_tokens({}) == 0 + + +def test_passthrough_usage_counts_gemini_thinking_tokens() -> None: + """_passthrough_usage_from_json must include thinking tokens in output_tokens.""" + usage = _passthrough_usage_from_json( + { + "usageMetadata": { + "promptTokenCount": 1000, + "candidatesTokenCount": 200, + "thoughtsTokenCount": 500, + "totalTokenCount": 1700, + "cachedContentTokenCount": 100, + } + } + ) + assert usage["output_tokens"] == 700 + assert usage["input_tokens"] == 1000 + + def test_vertex_passthrough_records_usage_metadata_for_dashboard() -> None: handler = object.__new__(HeadroomProxy) handler.http_client = _VertexUsageClient()