mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(realtime): retain missing translation input usage at close
This commit is contained in:
parent
c0b814dbe8
commit
998bde0963
2 changed files with 69 additions and 52 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue