headroom/tests/test_bedrock_streaming_input_tokens.py

Ignoring revisions in .git-blame-ignore-revs. Click here to bypass and see the normal blame view.

153 lines
5.6 KiB
Python
Raw Normal View History

fix(proxy): report real input tokens on streaming message_start (#1132) (#1305) ## Description LiteLLM/Bedrock streaming never surfaces prompt tokens mid-stream — it emits `message_start` with `usage.input_tokens=0` and only reports `output_tokens` (at the end, in `message_delta`). Anthropic clients such as Claude Code read `usage.input_tokens` from the **first** SSE event (`message_start`) to emit OTel/cost metrics, so every Headroom + Bedrock streaming request reported ~0 input tokens — underreporting token usage by ~99% in Athena/CloudWatch dashboards. Only `output_tokens` was tracked correctly. `StreamingMixin._stream_response_bedrock` now backfills `input_tokens` on `message_start` with the count Headroom actually sent upstream (`optimized_tokens`, already a parameter of that method) when the backend left it unset/zero. A non-zero value the backend genuinely reports is preserved untouched. Closes #1132 ## Type of Change - [x] Bug fix (non-breaking change that fixes an issue) - [ ] New feature (non-breaking change that adds functionality) - [ ] Breaking change (fix or feature that would cause existing functionality to change) - [ ] Documentation update - [ ] Performance improvement - [ ] Code refactoring (no functional changes) ## Changes Made - `headroom/proxy/handlers/streaming.py`: in `_stream_response_bedrock`, rewrite the `message_start` event's `usage.input_tokens` to `optimized_tokens` before it is serialized to the client, when the backend reported `0`/unset (and `optimized_tokens > 0`). Non-zero upstream values pass through unchanged. - `tests/test_bedrock_streaming_input_tokens.py`: new test that drives the Bedrock streaming route end-to-end with a LiteLLM-shaped backend (data-only `StreamEvent`s, `raw_sse=None`) and asserts the client-received `message_start` carries a real input-token count; plus a guard that a genuine non-zero upstream value is preserved. - `CHANGELOG.md`: Bug Fixes entry under Unreleased. ## Testing - [x] Unit tests pass (`pytest`) - [x] Linting passes (`ruff check .`) - [x] Type checking passes (`mypy headroom`) - [x] New tests added for new functionality - [x] Manual testing performed ### Test Output ```text $ uv run pytest tests/test_bedrock_streaming_input_tokens.py \ tests/test_backend_streaming_cache_metrics.py \ tests/test_proxy_streaming_resilience.py tests/test_streaming_usage_parser.py -q 39 passed, 1 warning in 46.59s $ uv run ruff check headroom/proxy/handlers/streaming.py tests/test_bedrock_streaming_input_tokens.py All checks passed! $ uv run ruff format --check ... # 2 files already formatted $ uv run mypy headroom/proxy/handlers/streaming.py Success: no issues found in 1 source file ``` ## TDD verification (RED → GREEN) The new test exercises the exact bug path (LiteLLM-shaped `message_start` with `input_tokens=0`, `raw_sse=None` → handler re-serializes `event.data`). **RED** — prod fix reverted (`git stash push -- headroom/proxy/handlers/streaming.py`): ```text FAILED tests/test_bedrock_streaming_input_tokens.py::test_bedrock_streaming_backfills_input_tokens_on_message_start E AssertionError: message_start.usage.input_tokens reached the client as 0; expected the upstream-sent token count (#1132). E assert 0 > 0 1 failed, 1 passed ``` (The 1 passing test on RED is the backwards-compat guard — it asserts a genuine non-zero upstream value is *preserved*, which holds with or without the fix.) **GREEN** — fix applied: ```text tests/test_bedrock_streaming_input_tokens.py .. [100%] 2 passed, 1 warning in 27.88s ``` ## Real Behavior Proof - Environment: Linux, Python 3.13.12, headroom-ai @ this branch, `uv run`. - Exact command / steps: drive the real `/v1/messages` streaming route through `create_app(ProxyConfig(backend="anyllm", anyllm_provider="anthropic", optimize=False))` with a LiteLLM-shaped backend whose `message_start` reports `usage.input_tokens=0` (exactly what `LiteLLMBackend.stream_message` emits), then parse the SSE the client receives. - Observed result: **before fix** the client's `message_start` event carries `usage.input_tokens=0`; **after fix** it carries the real upstream-sent token count (`> 0`), matching the issue's expected behavior. Captured verbatim in the RED→GREEN block above. - Not tested: a live AWS Bedrock account end-to-end (no Bedrock credentials available). The test reproduces the exact SSE shape `LiteLLMBackend.stream_message` produces — `message_start` with `input_tokens=0` and no `raw_sse` — which is the code path the issue identifies. Cache-token fields (`cache_read_input_tokens`/`cache_creation_input_tokens`) are out of scope: LiteLLM streaming does not surface them mid-stream and they cannot be reliably known at `message_start` time. ## Review Readiness - [x] I have performed a self-review - [x] This PR is ready for human review ## Checklist - [x] My code follows the project's style guidelines - [x] I have performed a self-review of my code - [x] I have commented my code, particularly in hard-to-understand areas - [ ] I have made corresponding changes to the documentation — N/A (no doc surface enumerates this behavior) - [x] My changes generate no new warnings - [x] I have added tests that prove my fix is effective - [x] New and existing unit tests pass locally with my changes - [x] I have updated the CHANGELOG.md ## Additional Notes - The fix lives in the proxy handler (`_stream_response_bedrock`), not the LiteLLM backend, because that is the layer that knows `optimized_tokens` — the authoritative count of input tokens Headroom sent upstream. Wiring it into the generic backend interface would be invasive and would duplicate tokenization. - Scope is intentionally limited to `input_tokens` (the headline metric from the issue). Cache-token fields are not inferable upfront and are left as-is. Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-23 19:53:15 +02:00
"""Bedrock/LiteLLM streaming must report real input tokens on message_start.
Regression coverage for issue #1132.
LiteLLM/Bedrock streaming never surfaces prompt tokens during a stream it
emits ``message_start`` with ``usage.input_tokens=0`` and only reports
``output_tokens`` (at the end, in ``message_delta``). Anthropic clients such as
Claude Code read ``usage.input_tokens`` from the ``message_start`` SSE event to
emit OTel/cost metrics, so every Headroom+Bedrock streaming request reported ~0
input tokens underreporting token usage by ~99%.
``StreamingMixin._stream_response_bedrock`` now backfills ``input_tokens`` on
``message_start`` with the count Headroom actually sent upstream
(``optimized_tokens``) when the backend left it unset/zero, while preserving any
non-zero value the backend genuinely reported.
"""
from __future__ import annotations
import json
from collections.abc import AsyncIterator
from unittest.mock import MagicMock, patch
import pytest
fastapi = pytest.importorskip("fastapi")
httpx = pytest.importorskip("httpx")
from fastapi.testclient import TestClient # noqa: E402
from headroom.backends.base import StreamEvent # noqa: E402
from headroom.proxy.server import ProxyConfig, create_app # noqa: E402
def _make_bedrock_backend(events: list[StreamEvent]) -> MagicMock:
"""Mock backend yielding Anthropic ``StreamEvent``s (no ``raw_sse``).
Mirrors ``LiteLLMBackend.stream_message``: it constructs each event from a
``data`` dict and never sets ``raw_sse``, so the handler re-serializes
``event.data`` the exact path that carries the #1132 bug.
"""
async def fake_stream(body: dict, headers: dict) -> AsyncIterator[StreamEvent]:
for evt in events:
yield evt
mock = MagicMock()
mock.name = "bedrock"
mock.stream_message = fake_stream
mock.map_model_id = MagicMock(return_value="claude-3-5-sonnet-20241022")
mock.supports_model = MagicMock(return_value=True)
return mock
def _bedrock_events(input_tokens: int) -> list[StreamEvent]:
"""Build a minimal Anthropic streaming sequence as LiteLLM emits it."""
message_start = {
"type": "message_start",
"message": {
"id": "msg_1",
"model": "claude-3-5-sonnet-20241022",
"role": "assistant",
"type": "message",
"content": [],
# LiteLLM hardcodes this to 0 — the bug under test.
"usage": {"input_tokens": input_tokens, "output_tokens": 0},
},
}
block_start = {
"type": "content_block_start",
"index": 0,
"content_block": {"type": "text", "text": ""},
}
block_delta = {
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": "hi"},
}
block_stop = {"type": "content_block_stop", "index": 0}
message_delta = {
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 50},
}
message_stop = {"type": "message_stop"}
# raw_sse=None mirrors LiteLLMBackend.stream_message (data-only events).
return [
StreamEvent(event_type=e["type"], data=e)
for e in [
message_start,
block_start,
block_delta,
block_stop,
message_delta,
message_stop,
]
]
def _message_start_input_tokens(sse_body: str) -> int:
"""Extract ``message.usage.input_tokens`` from the message_start SSE event."""
for block in sse_body.split("\n\n"):
if "message_start" not in block:
continue
for line in block.splitlines():
if line.startswith("data: "):
payload = json.loads(line[len("data: ") :])
return payload["message"]["usage"]["input_tokens"]
raise AssertionError(f"No message_start event with usage found in SSE:\n{sse_body[:500]}")
def _post_stream(backend: MagicMock) -> str:
config = ProxyConfig(
optimize=False,
cache_enabled=False,
rate_limit_enabled=False,
backend="anyllm",
anyllm_provider="anthropic",
)
with patch("headroom.proxy.server.AnyLLMBackend", return_value=backend):
app = create_app(config)
with TestClient(app) as client:
resp = client.post(
"/v1/messages",
json={
"model": "claude-3-5-sonnet-20241022",
"messages": [{"role": "user", "content": "hello there general"}],
"max_tokens": 64,
"stream": True,
},
headers={"x-api-key": "sk-ant-test", "anthropic-version": "2023-06-01"},
)
assert resp.status_code == 200, resp.text[:200]
return resp.text
def test_bedrock_streaming_backfills_input_tokens_on_message_start() -> None:
"""When the backend reports input_tokens=0, the client must see a real count."""
body = _post_stream(_make_bedrock_backend(_bedrock_events(input_tokens=0)))
client_input_tokens = _message_start_input_tokens(body)
assert client_input_tokens > 0, (
"message_start.usage.input_tokens reached the client as "
f"{client_input_tokens}; expected the upstream-sent token count (#1132)."
)
def test_bedrock_streaming_preserves_nonzero_upstream_input_tokens() -> None:
"""A genuine non-zero input_tokens from the backend must pass through untouched."""
upstream_input_tokens = 777
body = _post_stream(_make_bedrock_backend(_bedrock_events(input_tokens=upstream_input_tokens)))
assert _message_start_input_tokens(body) == upstream_input_tokens