fix(proxy): stamp the client alias on a copy of each streamed chunk so pricing sees the deployment model (#44341)

* fix(proxy): stamp the client alias on a copy of each streamed chunk so pricing sees the deployment model

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(spend): streamed alias matching a capability rule bills the deployment price

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(spend): assert every streamed chunk carries the client alias

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(logging): log the client alias on the priced streamed response, the same as non-streamed

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: kerry <kerry@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-03 06:44:52 +00:00 • committed by GitHub
parent a76ba8c01e
commit 7b432d78d2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 320 additions and 11 deletions

View file

@ -679,6 +679,9 @@ class Logging(LiteLLMLoggingBaseClass):
self.truncated_messages_for_logging: str | list | dict | None = None # mutable-ok: logged messages shape
## TIME TO FIRST TOKEN LOGGING ##
self.completion_start_time: datetime.datetime | None = None
# The model the proxy shows the client on streamed chunks. The logged streamed response carries it
# once that response is priced, the same way a non-streamed response is logged
self.client_facing_stream_model: str | None = None
self.zero_cost_warned: bool = False
self._llm_caching_handler: LLMCachingHandler | None = None
@ -2482,6 +2485,15 @@ class Logging(LiteLLMLoggingBaseClass):
setattr(result, "usage", transformed_usage)
return result
def _with_client_facing_stream_model(
self,
response: ModelResponse | TextCompletionResponse | ResponsesAPIResponse | InteractionsAPIResponse,
) -> ModelResponse | TextCompletionResponse | ResponsesAPIResponse | InteractionsAPIResponse:
model: Final = self.client_facing_stream_model
if model is None or getattr(response, "model", None) in (None, model):
return response
return response.model_copy(update={"model": model})
def _success_handler_helper_fn(
self,
result=None,
@ -2797,9 +2809,11 @@ class Logging(LiteLLMLoggingBaseClass):
result=complete_streaming_response
)
self._merge_hidden_params_from_response_into_metadata(complete_streaming_response)
logged_streaming_response: Final = self._with_client_facing_stream_model(complete_streaming_response)
self.model_call_details["complete_streaming_response"] = logged_streaming_response
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
logged_streaming_response, start_time, end_time
)
standard_logging_payload: Final[StandardLoggingPayload | None] = self.model_call_details.get(
"standard_logging_object"
@ -3338,10 +3352,13 @@ class Logging(LiteLLMLoggingBaseClass):
await self._prepare_baseline_cache_estimate(complete_streaming_response)
logged_streaming_response: Final = self._with_client_facing_stream_model(complete_streaming_response)
self.model_call_details["async_complete_streaming_response"] = logged_streaming_response
## STANDARDIZED LOGGING PAYLOAD
try:
self.model_call_details["standard_logging_object"] = self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
logged_streaming_response, start_time, end_time
)
except Exception: # noqa: BLE001 # payload build must never block later callbacks (slot release)
verbose_logger.exception(

View file

@ -9467,12 +9467,17 @@ def _restamp_streaming_chunk_model(
)
model_mismatch_logged = True
# The streaming wrapper keeps these same chunk objects to assemble the response it
# prices, so stamp a copy for the client and leave the provider's model for pricing.
# The logging object stamps the same model on the assembled response after pricing it.
logging_obj: Final = request_data.get("litellm_logging_obj")
if isinstance(logging_obj, LiteLLMLoggingObj):
logging_obj.client_facing_stream_model = target_model
if isinstance(chunk, dict):
chunk["model"] = target_model
return chunk, model_mismatch_logged
return {**chunk, "model": target_model}, model_mismatch_logged
try:
chunk.model = target_model
return chunk.model_copy(update={"model": target_model}), model_mismatch_logged
except Exception as e:
verbose_proxy_logger.error(
"litellm_call_id=%s: failed to override chunk.model=%r on chunk_type=%s. error=%s",

View file

@ -0,0 +1,253 @@
"""A streamed alias never replaces the deployment's model for pricing (LIT-9065).
The proxy shows the client's alias on every streamed chunk, but the chunks kept for end-of-stream cost calculation
keep the deployment's model. "claude-opus-4.8-<digits>" is no cost-map key and only matches the claude capability
rules, whose model info carries no prices, so a stream through that alias must bill exactly what the plain alias
"integration-<hex>" bills at the same deployment rates, and the client must still see the alias it asked for.
Logging callbacks see that alias as the response model on streamed requests, the same as on non-streamed ones
"""
import json
from collections.abc import Callable, Iterator, Mapping
from hashlib import sha256
from pathlib import Path
from typing import Final
from uuid import uuid4
import pytest
import yaml
from integration._support.client import (
Gateway,
Scenario,
eventually,
gateway_from_environment,
object_value,
string_value,
)
from integration._support.database import read_rows
from integration._support.otlp_sink import owned_sinks, recorded_spans
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue
def _sse_event(name: str, payload: dict[str, JsonValue]) -> bytes:
return f"event: {name}\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n".encode()
def _anthropic_reply(request: Request) -> Reply:
assert request.target.endswith("/v1/messages"), request.target
body: Final = json.loads(request.body)
assert body["model"] == "claude-opus-4-8", body
if body.get("stream") is not True:
return Reply(
body=json.dumps(
{
"id": f"msg_{uuid4().hex[:12]}",
"type": "message",
"role": "assistant",
"model": "claude-opus-4-8",
"content": [{"type": "text", "text": "hi"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 30, "output_tokens": 40},
}
).encode()
)
return Reply(
content_type="text/event-stream",
chunks=(
_sse_event(
"message_start",
{
"type": "message_start",
"message": {
"id": f"msg_{uuid4().hex[:12]}",
"type": "message",
"role": "assistant",
"model": "claude-opus-4-8",
"content": [],
"stop_reason": None,
"stop_sequence": None,
"usage": {"input_tokens": 30, "output_tokens": 1},
},
},
),
_sse_event(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
),
_sse_event(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hi"}},
),
_sse_event("content_block_stop", {"type": "content_block_stop", "index": 0}),
_sse_event(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 40},
},
),
_sse_event("message_stop", {"type": "message_stop"}),
),
)
def _deployment(
scenario: Scenario,
model_name: str,
litellm_params: dict[str, JsonValue],
model_info: dict[str, JsonValue] | None = None,
) -> str:
created: Final = scenario.gateway.post(
"/model/new", {"model_name": model_name, "litellm_params": litellm_params, "model_info": model_info or {}}
)
identity: Final = string_value(object_value(created["model_info"])["id"])
scenario.cleanups.callback(scenario.delete_model, identity)
return model_name
def _streamed_spend(gateway: Gateway, scenario: Scenario, model: str, content: str) -> dict[str, JsonValue]:
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": content}],
"stream": True,
"stream_options": {"include_usage": True},
},
key=key,
)
assert response.status_code == 200, response.text
chunks: Final = tuple(
json.loads(line.removeprefix("data: "))
for line in response.text.splitlines()
if line.startswith("data: ") and line != "data: [DONE]"
)
assert chunks and {chunk["model"] for chunk in chunks} == {model}, response.text
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE api_key=%s',
(sha256(key.encode()).hexdigest(),),
),
lambda values: len(values) == 1,
seconds=70,
)
return rows[0]
def _listed_deployments(gateway: Gateway, model_name: str) -> tuple[dict[str, JsonValue], ...]:
entries: Final = gateway.get("/model/info")["data"]
assert isinstance(entries, list)
return tuple(object_value(entry) for entry in entries if object_value(entry)["model_name"] == model_name)
def _deployment_pricing(gateway: Gateway, model_name: str) -> dict[str, JsonValue]:
listed: Final = eventually(lambda: _listed_deployments(gateway, model_name), lambda found: len(found) == 1)
return object_value(listed[0]["model_info"])
_BACKENDS: Final = (
pytest.param(
lambda _: {"model": "vertex_ai/claude-opus-4-8@default", "mock_response": "hi"},
id="vertex-mock-response",
),
pytest.param(
lambda wire_url: {
"model": "anthropic/claude-opus-4-8",
"api_key": "integration-provider-key",
"api_base": wire_url,
},
id="anthropic-upstream",
),
)
@pytest.mark.parametrize("litellm_params", _BACKENDS)
@pytest.mark.timeout(180)
def test_streamed_alias_matching_a_capability_rule_bills_the_deployment_price(
gateway: Gateway, litellm_params: Callable[[str], dict[str, JsonValue]]
) -> None:
with wire_server(_anthropic_reply) as wire, gateway.scenario() as scenario:
content: Final = f"alias billing {uuid4().hex}"
plain_alias: Final = f"integration-{uuid4().hex}"
rule_alias: Final = f"claude-opus-4.8-{uuid4().int % 10**8:08d}"
exact_row: Final = _streamed_spend(
gateway, scenario, _deployment(scenario, plain_alias, litellm_params(wire.url)), content
)
alias_row: Final = _streamed_spend(
gateway, scenario, _deployment(scenario, rule_alias, litellm_params(wire.url)), content
)
for model_name, row in ((plain_alias, exact_row), (rule_alias, alias_row)):
pricing: Final = _deployment_pricing(gateway, model_name)
input_rate: Final = float(str(pricing["input_cost_per_token"]))
output_rate: Final = float(str(pricing["output_cost_per_token"]))
uplift: Final = float(str(pricing["regional_endpoint_uplift_multiplier"] or 1))
assert input_rate > 0 and output_rate > 0, pricing
assert float(str(row["spend"])) == pytest.approx(
uplift
* (float(str(row["prompt_tokens"])) * input_rate + float(str(row["completion_tokens"])) * output_rate)
), (model_name, row, pricing)
@pytest.fixture(scope="module")
def otel_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[tuple[Gateway, str]]:
directory: Final = tmp_path_factory.mktemp("stream-alias-otel")
with owned_sinks(directory / "sinks") as sinks, gateway_from_environment() as base:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"] = {**config["litellm_settings"], "callbacks": ["otel"]}
config["callback_settings"] = {
"otel": {"exporter": "http/json", "endpoint": sinks.operator, "use_simple_processor": True}
}
path: Final = directory / "otel.yaml"
path.write_text(yaml.safe_dump(config))
overrides: Final = {"OTEL_EXPORTER": "http/json", "OTEL_ENDPOINT": sinks.operator}
with owned_proxy(base, directory, overrides, config=path) as candidate:
yield candidate, sinks.operator
def _logged_response_models(sink: str, call_ids: Mapping[str, str]) -> dict[str, JsonValue]:
_, spans = recorded_spans(sink)
return {
label: span["attributes"]["gen_ai.response.model"]
for span in spans
for label, call_id in call_ids.items()
if span["attributes"].get("litellm.call_id") == call_id and "gen_ai.response.model" in span["attributes"]
}
@pytest.mark.parametrize("litellm_params", _BACKENDS)
@pytest.mark.timeout(240)
def test_logged_response_model_is_the_client_alias_whether_or_not_the_request_streams(
otel_proxy: tuple[Gateway, str], litellm_params: Callable[[str], dict[str, JsonValue]]
) -> None:
candidate, sink = otel_proxy
with wire_server(_anthropic_reply) as wire, candidate.scenario() as scenario:
alias: Final = f"claude-opus-4.8-{uuid4().int % 10**8:08d}"
key: Final = scenario.key(models=[_deployment(scenario, alias, litellm_params(wire.url))])
call_ids: Final[dict[str, str]] = {}
for label, stream_fields in (
("non-streamed", {}),
("streamed", {"stream": True, "stream_options": {"include_usage": True}}),
):
response = candidate.request(
"POST",
"/v1/chat/completions",
{"model": alias, "messages": [{"role": "user", "content": f"logged alias {uuid4().hex}"}]}
| stream_fields,
key=key,
)
assert response.status_code == 200, response.text
call_ids[label] = response.headers["x-litellm-call-id"]
logged: Final = eventually(
lambda: _logged_response_models(sink, call_ids),
lambda found: len(found) == 2,
seconds=60,
return_last_on_timeout=True,
)
assert logged == {"non-streamed": alias, "streamed": alias}, call_ids

View file

@ -321,6 +321,18 @@ async def test_mcp_direct_content_edit_invalidates_stale_structured_data(logging
assert "SECRET-1234" not in result.model_dump_json()
def test_with_client_facing_stream_model_stamps_a_copy_of_the_priced_response(logging_obj):
response = ModelResponse(model="claude-opus-4-6@default")
logging_obj.client_facing_stream_model = "claude-opus-4.6"
logged = logging_obj._with_client_facing_stream_model(response)
assert (logged.model, response.model) == ("claude-opus-4.6", "claude-opus-4-6@default")
def test_with_client_facing_stream_model_keeps_the_response_when_the_proxy_set_no_model(logging_obj):
response = ModelResponse(model="claude-opus-4-6@default")
assert logging_obj._with_client_facing_stream_model(response) is response
def test_get_combined_callback_list_preserves_insertion_order(logging_obj):
assert logging_obj.get_combined_callback_list(
dynamic_success_callbacks=["prometheus", "langfuse", "datadog", "otel", "s3"],

View file

@ -283,8 +283,9 @@ def test_restamp_streaming_chunk_model_overrides_model_on_basemodel():
"model": new_chunk.model,
"logged": logged,
"same_object": new_chunk is chunk,
"original_model": chunk.model,
}
assert snapshot == {"model": "gpt-4", "logged": True, "same_object": True}
assert snapshot == {"model": "gpt-4", "logged": True, "same_object": False, "original_model": "openai/internal-x"}
@pytest.mark.parametrize("return_raw_model_name", [False, True])
@ -310,8 +311,7 @@ def test_restamp_streaming_chunk_model_overrides_model_on_dict():
request_data={},
model_mismatch_logged=True,
)
assert new_chunk["model"] == "gpt-4"
assert logged is True
assert (new_chunk["model"], chunk["model"], logged) == ("gpt-4", "internal", True)
def test_restamp_streaming_chunk_model_uses_fallback_model_from_metadata():
@ -443,7 +443,7 @@ def test_restamp_streaming_chunk_model_fastest_response_preserves_model():
assert logged is False
def test_restamp_streaming_chunk_model_setattr_exception_logs_and_returns():
def test_restamp_streaming_chunk_model_restamps_a_frozen_chunk_through_a_copy():
from pydantic import ConfigDict
class FrozenChunk(_simple_chunk().__class__):
@ -462,8 +462,30 @@ def test_restamp_streaming_chunk_model_setattr_exception_logs_and_returns():
request_data={"litellm_call_id": "test-id"},
model_mismatch_logged=False,
)
assert new_chunk.model == "openai/internal-x"
assert logged is True
assert (new_chunk.model, chunk.model, logged) == ("gpt-4", "openai/internal-x", True)
def test_restamp_streaming_chunk_model_records_the_client_model_on_the_logging_object():
import time
from litellm.litellm_core_utils.litellm_logging import Logging
logging_obj = Logging(
model="openai/internal-x",
messages=[],
stream=True,
call_type="acompletion",
start_time=time.time(),
litellm_call_id="test-id",
function_id="test-id",
)
_restamp_streaming_chunk_model(
chunk=_simple_chunk(model="openai/internal-x"),
requested_model_from_client="gpt-4",
request_data={"litellm_call_id": "test-id", "litellm_logging_obj": logging_obj},
model_mismatch_logged=False,
)
assert logging_obj.client_facing_stream_model == "gpt-4"
def test_format_fallback_metadata_sse_event():