fix(databricks): surface prompt-cache token counts in streaming usage

chunk_parser built ModelResponseStream without passing usage, so the
cache_read_input_tokens and cache_creation_input_tokens that Databricks
returns for Anthropic models never reached the cost calculator. Every
streamed request was billed at the full input rate even when served
from cache.

ModelResponseStream already coerces a usage dict into Usage, which maps
those keys into prompt_tokens_details, so passing the chunk's usage
through is sufficient.
This commit is contained in:
pokepoke81 2026-08-14 10:45:21 -04:00
parent 423b791ee0
commit ce66cbce0e
2 changed files with 77 additions and 0 deletions

View file

@ -733,6 +733,7 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator):
created=chunk["created"],
model=chunk["model"],
choices=translated_choices,
usage=chunk.get("usage"),
)
except KeyError as e:
raise DatabricksException(

View file

@ -423,3 +423,79 @@ def test_databricks_config_probes_capabilities_under_databricks_namespace():
without this override they probed the ``anthropic`` cost-map namespace and
ignored the exact ``databricks/databricks-claude-*`` entries."""
assert DatabricksConfig().custom_llm_provider == "databricks"
def _streaming_chunk(usage=None, choices=None):
base = {
"id": "chatcmpl-test",
"created": 1234567890,
"model": "databricks-claude-sonnet-5",
"choices": [{"delta": {"content": "hi"}}] if choices is None else choices,
}
return base if usage is None else {**base, "usage": usage}
@pytest.mark.parametrize(
"cache_read, cache_creation, expected_cached, expected_written",
[
(12002, 0, 12002, 0),
(0, 12002, 0, 12002),
],
ids=["warm_cache_read", "cold_cache_write"],
)
def test_chunk_parser_surfaces_prompt_cache_usage(cache_read, cache_creation, expected_cached, expected_written):
"""Databricks returns Anthropic prompt-cache counts in the streaming usage object,
but chunk_parser dropped usage entirely, so cache-aware pricing never reached the
cost calculator and every streamed request was billed at the full input rate."""
iterator = DatabricksChatResponseIterator(streaming_response=None, sync_stream=True)
result = iterator.chunk_parser(
_streaming_chunk(
usage={
"prompt_tokens": 12011,
"completion_tokens": 8,
"total_tokens": 12019,
"cache_read_input_tokens": cache_read,
"cache_creation_input_tokens": cache_creation,
}
)
)
assert result.usage is not None
assert result.usage.prompt_tokens == 12011
assert result.usage.completion_tokens == 8
assert result.usage.prompt_tokens_details is not None
assert result.usage.prompt_tokens_details.cached_tokens == expected_cached
assert result.usage._cache_creation_input_tokens == expected_written
def test_chunk_parser_surfaces_usage_only_final_chunk():
"""stream_options={"include_usage": True} emits a trailing chunk whose choices
list is empty; usage must still reach the caller."""
iterator = DatabricksChatResponseIterator(streaming_response=None, sync_stream=True)
result = iterator.chunk_parser(
_streaming_chunk(
usage={
"prompt_tokens": 100,
"completion_tokens": 5,
"total_tokens": 105,
"cache_read_input_tokens": 90,
},
choices=[],
)
)
assert result.choices == []
assert result.usage is not None
assert result.usage.prompt_tokens_details.cached_tokens == 90
def test_chunk_parser_without_usage_still_parses_content():
iterator = DatabricksChatResponseIterator(streaming_response=None, sync_stream=True)
result = iterator.chunk_parser(_streaming_chunk())
assert result.id == "chatcmpl-test"
assert result.model == "databricks-claude-sonnet-5"
assert result.choices[0]["delta"]["content"] == "hi"