From a66dd7dda03533d0c4843ea86e3da4e048d520fa Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Sat, 26 Sep 2026 09:24:14 -0500 Subject: [PATCH] fix(realtime): preserve translation deployment pricing and input usage --- litellm/cost_calculator.py | 36 +++++---- .../litellm_core_utils/realtime_streaming.py | 77 +++++++++++++++---- tests/unit/cookbook/__init__.py | 0 .../test_realtime_streaming.py | 56 ++++++++++++++ tests/unit/test_cost_calculator.py | 49 ++++++++++++ 5 files changed, 188 insertions(+), 30 deletions(-) create mode 100644 tests/unit/cookbook/__init__.py diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 01a84ac9c9a..81d913a24b9 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -3014,6 +3014,7 @@ def handle_realtime_stream_cost_calculation( results=results, custom_llm_provider=custom_llm_provider, litellm_model_name=litellm_model_name, + potential_model_names=potential_model_names, ) total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost + translation_cost @@ -3043,6 +3044,7 @@ def handle_realtime_translation_cost_calculation( results: OpenAIRealtimeStreamList, custom_llm_provider: str, litellm_model_name: str, + potential_model_names: Sequence[str | None] = (), ) -> float: usage_events: Final = tuple( usage @@ -3062,22 +3064,28 @@ def handle_realtime_translation_cost_calculation( ) if input_seconds <= 0 and output_seconds <= 0: return 0.0 - try: - model_info: Final = litellm.get_model_info( - model=litellm_model_name, - custom_llm_provider=custom_llm_provider, - ) - except Exception: # noqa: BLE001 # unknown model metadata should yield zero translation cost - return 0.0 - input_cost_per_second: Final = model_info.get("input_cost_per_second") - output_cost_per_second: Final = model_info.get("output_cost_per_second") - input_cost: Final = ( - input_seconds * input_cost_per_second if isinstance(input_cost_per_second, (int, float)) else 0.0 + model_infos: Final = tuple( + _get_model_info_or_none(model, custom_llm_provider) + for model in (*potential_model_names, litellm_model_name) + if model is not None ) - output_cost: Final = ( - output_seconds * output_cost_per_second if isinstance(output_cost_per_second, (int, float)) else 0.0 + input_cost_per_second: Final = next( + ( + rate + for info in model_infos + if (rate := _declared_transcription_rate(info, ("input_cost_per_second",))) is not None + ), + 0.0, ) - return input_cost + output_cost + output_cost_per_second: Final = next( + ( + rate + for info in model_infos + if (rate := _declared_transcription_rate(info, ("output_cost_per_second",))) is not None + ), + 0.0, + ) + return input_seconds * input_cost_per_second + output_seconds * output_cost_per_second def handle_realtime_transcription_cost_calculation( diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 455967f7cfe..95d6b1eb7ba 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -154,6 +154,8 @@ class RealTimeStreaming: self.session_tools: list[dict] = [] self.tool_calls: list[dict] = [] self._is_translation_session = translation_session + self._translation_input_seconds = 0.0 + self._translation_input_bytes_per_second = 48000.0 self._translation_output_audio_bytes = 0 self._translation_output_bytes_per_second = 48000.0 self._translation_usage_finalized = False @@ -496,27 +498,61 @@ class RealTimeStreaming: return self._translation_output_audio_bytes += len(decoded) + @staticmethod + def _translation_audio_bytes_per_second(audio_format: object) -> float | None: + if audio_format == "pcm16": + return 48000.0 + if audio_format in ("g711_ulaw", "g711_alaw"): + return 8000.0 + if not isinstance(audio_format, Mapping): + return None + rate: Final = normalized_audio_duration_seconds(audio_format.get("rate")) + if rate is None or rate <= 0: + return None + match audio_format.get("type"): + case "audio/pcm": + return rate * 2 + case "audio/pcmu" | "audio/pcma": + return rate + case _: + return None + + def _capture_translation_input_audio(self, message: str) -> None: + if not self._is_translation_session: + return + try: + event: Final = _decode_json_object(message) + except (json.JSONDecodeError, TypeError): + return + self._capture_translation_output_format(event) + if event.get("type") not in ( + "input_audio_buffer.append", + "session.input_audio_buffer.append", + ) or not isinstance(audio := event.get("audio"), str): + return + try: + decoded: Final = base64.b64decode(audio, validate=True) + except (ValueError, TypeError): + return + self._translation_input_seconds += len(decoded) / self._translation_input_bytes_per_second + def _capture_translation_output_format(self, event_obj: Mapping[str, object]) -> None: session: Final = event_obj.get("session") if not isinstance(session, dict): return audio: Final = session.get("audio") + audio_input: Final = audio.get("input") if isinstance(audio, dict) else None + input_format: Final = ( + audio_input.get("format") if isinstance(audio_input, dict) else session.get("input_audio_format") + ) + input_rate: Final = self._translation_audio_bytes_per_second(input_format) + if input_rate is not None: + self._translation_input_bytes_per_second = input_rate output: Final = audio.get("output") if isinstance(audio, dict) else None audio_format: Final = output.get("format") if isinstance(output, dict) else None - if isinstance(audio_format, str): - if audio_format in ("g711_ulaw", "g711_alaw"): - self._translation_output_bytes_per_second = 8000.0 - return - if not isinstance(audio_format, dict): - return - format_type: Final = audio_format.get("type") - rate: Final = audio_format.get("rate") - if not isinstance(rate, (int, float)) or rate <= 0: - return - if format_type == "audio/pcm": - self._translation_output_bytes_per_second = float(rate) * 2 - elif format_type in ("audio/pcmu", "audio/pcma"): - self._translation_output_bytes_per_second = float(rate) + output_rate: Final = self._translation_audio_bytes_per_second(audio_format) + if output_rate is not None: + self._translation_output_bytes_per_second = output_rate def _finalize_translation_usage(self) -> None: if self._translation_usage_finalized: @@ -530,10 +566,16 @@ class RealTimeStreaming: ): self._translation_usage_finalized = True return - if self._translation_output_audio_bytes == 0: + if self._translation_output_audio_bytes == 0 and self._translation_input_seconds == 0: return output_seconds: Final = self._translation_output_audio_bytes / self._translation_output_bytes_per_second - synthetic_usage: Final = OpenAIRealtimeTranslationDurationUsage(type="duration", output_seconds=output_seconds) + synthetic_usage: Final = ( + OpenAIRealtimeTranslationDurationUsage( + type="duration", input_seconds=self._translation_input_seconds, output_seconds=output_seconds + ) + if self._translation_input_seconds > 0 + else OpenAIRealtimeTranslationDurationUsage(type="duration", output_seconds=output_seconds) + ) self.messages.append(OpenAIRealtimeTranslationClosedEvent(type="session.closed", usage=synthetic_usage)) self._translation_usage_finalized = True @@ -587,8 +629,11 @@ class RealTimeStreaming: if is_content_message: self._content_sent_after_setup = True sent = True + if sent: + self._capture_translation_input_audio(message) return sent await self.backend_ws.send(message) + self._capture_translation_input_audio(message) return True async def _apply_nested_transcription_model_policy(self, message: str) -> str: diff --git a/tests/unit/cookbook/__init__.py b/tests/unit/cookbook/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index 6dd582e2cec..d0ee76bb3c8 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -2933,6 +2933,62 @@ def test_translation_audio_duration_is_finalized_once(event_type: str): assert closed_events[0]["usage"] == {"type": "duration", "output_seconds": 1.0} +@pytest.mark.asyncio +@pytest.mark.parametrize("output_bytes", (0, 48000)) +@pytest.mark.parametrize("event_type", ("input_audio_buffer.append", "session.input_audio_buffer.append")) +@pytest.mark.parametrize( + "audio_format,bytes_per_second", + [("pcm16", 48000), ("g711_ulaw", 8000), ({"type": "audio/pcm", "rate": 16000}, 32000)], +) +async def test_translation_disconnect_bills_sent_input_audio( + output_bytes: int, event_type: str, audio_format: str | Mapping[str, object], bytes_per_second: int +) -> None: + import base64 + + backend: Final = MagicMock() + backend.send = AsyncMock() + streaming: Final = RealTimeStreaming( + websocket=_ga_client_ws(), + backend_ws=backend, + logging_obj=MagicMock(), + model="gpt-realtime-translate", + translation_session=True, + ) + await streaming._send_to_backend( + json.dumps({"type": "session.update", "session": {"audio": {"input": {"format": audio_format}}}}) + ) + await streaming._send_to_backend( + json.dumps({"type": event_type, "audio": base64.b64encode(bytes(2 * bytes_per_second)).decode()}) + ) + streaming._capture_translation_output_audio( + {"type": "session.output_audio.delta", "delta": base64.b64encode(bytes(output_bytes)).decode()} + ) + streaming._finalize_translation_usage() + streaming._finalize_translation_usage() + + assert streaming.messages == [ + { + "type": "session.closed", + "usage": {"type": "duration", "input_seconds": 2.0, "output_seconds": output_bytes / 48000}, + } + ] + + +@pytest.mark.asyncio +async def test_translation_failed_audio_send_is_not_billed() -> None: + backend: Final = MagicMock() + backend.send = AsyncMock(side_effect=RuntimeError("send failed")) + streaming: Final = RealTimeStreaming( + websocket=_ga_client_ws(), backend_ws=backend, logging_obj=MagicMock(), translation_session=True + ) + + with pytest.raises(RuntimeError, match="send failed"): + await streaming._send_to_backend(json.dumps({"type": "input_audio_buffer.append", "audio": "AAAA"})) + streaming._finalize_translation_usage() + + assert streaming.messages == [] + + def test_translation_audio_duration_uses_session_output_format(): import base64 diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 71c2ce3c742..bfb89f98347 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -4967,6 +4967,55 @@ def test_realtime_cached_multimodal_token_cost(_local_model_cost_map, provider: assert actual == pytest.approx(expected) +@pytest.mark.parametrize("input_override,output_override", [(None, None), (0.25, 0.75), (0.0, 0.0), (0.0, None)]) +def test_realtime_translation_uses_deployment_rates_before_base_rates( + _local_model_cost_map: None, + monkeypatch: pytest.MonkeyPatch, + input_override: float | None, + output_override: float | None, +) -> None: + monkeypatch.setitem( + litellm.model_cost, + "translation-base", + { + "litellm_provider": "azure", + "mode": "realtime", + "input_cost_per_second": 0.5, + "output_cost_per_second": 1.0, + }, + ) + monkeypatch.setitem( + litellm.model_cost, + "translation-deployment", + { + "litellm_provider": "azure", + "mode": "realtime", + **{ + key: rate + for key, rate in (("input_cost_per_second", input_override), ("output_cost_per_second", output_override)) + if rate is not None + }, + }, + ) + litellm.get_model_info.cache_clear() + events: Final[OpenAIRealtimeStreamList] = [ + {"type": "session.closed", "usage": {"type": "duration", "input_seconds": 3.0, "output_seconds": 2.0}} + ] + cost: Final = handle_realtime_stream_cost_calculation( + results=events, + combined_usage_object=Usage(), + custom_llm_provider="azure", + litellm_model_name="unmapped-provider-deployment", + custom_pricing_model="translation-deployment", + base_pricing_model="translation-base", + ) + + assert cost == pytest.approx( + 3 * (0.5 if input_override is None else input_override) + + 2 * (1.0 if output_override is None else output_override) + ) + + def test_realtime_translation_duration_cost(_local_model_cost_map): from litellm.cost_calculator import handle_realtime_translation_cost_calculation