From 998bde096306eb2ac2fd7a96e046958bba9516a3 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Sat, 26 Sep 2026 09:36:33 -0500 Subject: [PATCH] fix(realtime): retain missing translation input usage at close --- .../litellm_core_utils/realtime_streaming.py | 73 ++++++++----------- .../test_realtime_streaming.py | 48 +++++++++--- 2 files changed, 69 insertions(+), 52 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 95d6b1eb7ba..4d2d4dde999 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -155,7 +155,6 @@ class RealTimeStreaming: 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 @@ -445,9 +444,14 @@ class RealTimeStreaming: output_seconds: Final = ( normalized_audio_duration_seconds(usage.get("output_seconds")) if isinstance(usage, dict) else None ) - input_seconds: Final = ( + reported_input_seconds: Final = ( normalized_audio_duration_seconds(usage.get("input_seconds")) if isinstance(usage, dict) else None ) + input_seconds: Final = ( + reported_input_seconds + if reported_input_seconds is not None + else self._translation_input_seconds or None + ) synthetic_output_seconds: Final = ( self._translation_output_audio_bytes / self._translation_output_bytes_per_second if self._translation_output_audio_bytes > 0 @@ -456,13 +460,19 @@ class RealTimeStreaming: resolved_output_seconds: Final = output_seconds if output_seconds is not None else synthetic_output_seconds if input_seconds is not None or resolved_output_seconds is not None: if self._should_store_message(event_obj): - if output_seconds is None and synthetic_output_seconds is not None: + supplemental_usage: Final = OpenAIRealtimeTranslationDurationUsage( + type="duration", + input_seconds=float(input_seconds or 0.0) if reported_input_seconds is None else 0.0, + output_seconds=float(synthetic_output_seconds or 0.0) if output_seconds is None else 0.0, + ) + if ( + supplemental_usage.get("input_seconds", 0.0) > 0 + or supplemental_usage.get("output_seconds", 0.0) > 0 + ): self.messages.append( OpenAIRealtimeTranslationClosedEvent( type="session.closed", - usage=OpenAIRealtimeTranslationDurationUsage( - type="duration", output_seconds=synthetic_output_seconds - ), + usage=supplemental_usage, ) ) else: @@ -498,25 +508,6 @@ 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 @@ -524,35 +515,35 @@ class RealTimeStreaming: 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): + if event.get("type") != "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 + self._translation_input_seconds += len(decoded) / 48000.0 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 - 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 + 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) def _finalize_translation_usage(self) -> None: if self._translation_usage_finalized: diff --git a/tests/unit/litellm_core_utils/test_realtime_streaming.py b/tests/unit/litellm_core_utils/test_realtime_streaming.py index d0ee76bb3c8..94b0d215bf4 100644 --- a/tests/unit/litellm_core_utils/test_realtime_streaming.py +++ b/tests/unit/litellm_core_utils/test_realtime_streaming.py @@ -2935,13 +2935,8 @@ def test_translation_audio_duration_is_finalized_once(event_type: str): @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 + output_bytes: int, ) -> None: import base64 @@ -2955,10 +2950,7 @@ async def test_translation_disconnect_bills_sent_input_audio( 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()}) + json.dumps({"type": "session.input_audio_buffer.append", "audio": base64.b64encode(bytes(96000)).decode()}) ) streaming._capture_translation_output_audio( {"type": "session.output_audio.delta", "delta": base64.b64encode(bytes(output_bytes)).decode()} @@ -2983,12 +2975,46 @@ async def test_translation_failed_audio_send_is_not_billed() -> None: ) with pytest.raises(RuntimeError, match="send failed"): - await streaming._send_to_backend(json.dumps({"type": "input_audio_buffer.append", "audio": "AAAA"})) + await streaming._send_to_backend(json.dumps({"type": "session.input_audio_buffer.append", "audio": "AAAA"})) streaming._finalize_translation_usage() assert streaming.messages == [] +@pytest.mark.asyncio +@pytest.mark.parametrize("retain_close", (False, True)) +@pytest.mark.parametrize("reported_input,expected_input", [(None, 2.0), (0.0, 0.0), (0.25, 0.25)]) +async def test_translation_terminal_usage_fills_only_missing_input_duration( + monkeypatch: pytest.MonkeyPatch, retain_close: bool, reported_input: float | None, expected_input: float +) -> None: + import base64 + + monkeypatch.setattr(litellm, "logged_real_time_event_types", "*" if retain_close else None) + backend: Final = MagicMock() + backend.send = AsyncMock() + streaming: Final = RealTimeStreaming( + websocket=_ga_client_ws(), backend_ws=backend, logging_obj=MagicMock(), translation_session=True + ) + await streaming._send_to_backend( + json.dumps({"type": "session.input_audio_buffer.append", "audio": base64.b64encode(bytes(96000)).decode()}) + ) + close_event: Final = { + "type": "session.closed", + "usage": { + "type": "duration", + "output_seconds": 0.5, + **({"input_seconds": reported_input} if reported_input is not None else {}), + }, + } + streaming._capture_translation_output_audio(close_event) + streaming.store_message(close_event) + streaming._finalize_translation_usage() + + usage: Final = tuple(event["usage"] for event in streaming.messages if event.get("type") == "session.closed") + assert sum(item.get("input_seconds") or 0.0 for item in usage) == expected_input + assert sum(item.get("output_seconds") or 0.0 for item in usage) == 0.5 + + def test_translation_audio_duration_uses_session_output_format(): import base64