mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
a76ba8c01e
commit
7b432d78d2
5 changed files with 320 additions and 11 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
253
tests/integration/spend/test_stream_alias_billing.py
Normal file
253
tests/integration/spend/test_stream_alias_billing.py
Normal 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
|
||||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue