fix(streaming): drop zero cost on Vertex stream chunks

Vertex chunks reach chunk_creator with usage.cost=0 before cost resolution.
Treat that non-positive cost as absent so token pricing runs for the stream.
This commit is contained in:
leilei3167 2026-09-29 12:36:38 +00:00
parent 8567010897
commit 366c5f38e7
2 changed files with 109 additions and 1 deletions

View file

@ -1658,7 +1658,27 @@ class CustomStreamWrapper:
self.tool_call = True
if hasattr(chunk, "usage") and chunk.usage is not None:
model_response.usage = chunk.usage
attached_usage = chunk.usage
if self.custom_llm_provider == "vertex_ai":
# The vertex_ai branch returns before cost propagation, so a
# leftover usage.cost=0 would otherwise be copied here and
# treated as a provider total. Resolve it the same way as
# stream assembly: non-positive is absent, token pricing runs.
raw_cost = (
attached_usage.get("cost")
if isinstance(attached_usage, dict)
else getattr(attached_usage, "cost", None)
)
if raw_cost is not None and CustomStreamWrapper._resolve_provider_reported_cost(raw_cost) is None:
if isinstance(attached_usage, dict):
attached_usage = {key: value for key, value in attached_usage.items() if key != "cost"}
elif hasattr(attached_usage, "model_copy"):
attached_usage = attached_usage.model_copy(update={"cost": None})
try:
del attached_usage.cost
except (AttributeError, TypeError):
pass
model_response.usage = attached_usage
## RETURN ARG
result: Final = self.return_processed_chunk_logic(

View file

@ -2079,6 +2079,94 @@ def test_stream_spend_prices_vertex_anthropic_without_cache_read():
assert stream_spend >= token_total - 1e-12
def test_vertex_chunk_creator_drops_zero_cost_and_prices_tokens():
"""Vertex stream chunks can carry usage.cost=0 with real token counts.
chunk_creator must not keep that 0 as a provider total. Token pricing on
the assembled stream still has to be a positive spend.
"""
logging_obj = Logging(
model="claude-opus-5",
messages=[{"role": "user", "content": "count to five"}],
stream=True,
call_type="completion",
start_time=time.time(),
litellm_call_id="vertex-zero-cost-chunk",
function_id="1245",
)
logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai"
logging_obj.optional_params = {}
wrapper = CustomStreamWrapper(
completion_stream=None,
model="claude-opus-5",
logging_obj=logging_obj,
custom_llm_provider="vertex_ai",
)
source_chunks = _vertex_anthropic_stream_chunks(cache_read=44616, prompt_tokens=70961, completion_tokens=4096)
zero_cost_usage = source_chunks[-1].usage.model_copy(update={"cost": 0})
zero_cost_chunk = ModelResponseStream(
id="chatcmpl-vertex-zero-cost",
created=1745513207,
model="claude-opus-5",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content="Hi", role="assistant"),
)
],
usage=zero_cost_usage,
)
processed = wrapper.chunk_creator(chunk=zero_cost_chunk)
assert processed is not None
assert processed.usage.prompt_tokens == 70961
assert processed.usage.completion_tokens == 4096
assert processed.usage.cache_read_input_tokens == 44616
assert getattr(processed.usage, "cost", None) in (None,)
priced_chunk = ModelResponseStream(
id="chatcmpl-vertex-priced",
created=1745513208,
model="claude-opus-5",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content="Hi", role="assistant"),
)
],
usage=zero_cost_usage.model_copy(update={"cost": 0.01}),
)
priced = wrapper.chunk_creator(chunk=priced_chunk)
assert priced is not None
assert priced.usage.cost == 0.01
complete_response = litellm.stream_chunk_builder(
chunks=[processed],
messages=[{"role": "user", "content": "count to five"}],
logging_obj=logging_obj,
)
assert complete_response is not None
token_prompt, token_completion = litellm.cost_per_token(
model="claude-opus-5",
custom_llm_provider="vertex_ai",
usage_object=complete_response.usage,
)
token_total = token_prompt + token_completion
assert token_total > 0
stream_spend = _stream_spend_via_logging(
complete_response,
model="claude-opus-5",
custom_llm_provider="vertex_ai",
)
assert stream_spend > 0
assert stream_spend >= token_total - 1e-12
def test_handle_special_delta_attributes(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):