fix(realtime): retain missing translation input usage at close

This commit is contained in:
Emerson Gomes 2026-09-26 09:36:33 -05:00
parent c0b814dbe8
commit 998bde0963
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
2 changed files with 69 additions and 52 deletions

View file

@ -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:

View file

@ -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