diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index a43574b1a04..05393049f58 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3592d682d72..9af080f1b50 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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", diff --git a/tests/integration/spend/test_stream_alias_billing.py b/tests/integration/spend/test_stream_alias_billing.py new file mode 100644 index 00000000000..c9f8dba615a --- /dev/null +++ b/tests/integration/spend/test_stream_alias_billing.py @@ -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-" 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-" 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 diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index e07ffe00d4c..4b25da2ff79 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -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"], diff --git a/tests/unit/proxy/proxy_server/test_streaming_helpers.py b/tests/unit/proxy/proxy_server/test_streaming_helpers.py index 92de00a4a3f..69fa195e9d6 100644 --- a/tests/unit/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/unit/proxy/proxy_server/test_streaming_helpers.py @@ -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():