fix(realtime): preserve translation deployment pricing and input usage

This commit is contained in:
Emerson Gomes 2026-09-26 09:24:14 -05:00
parent 152b1e15d2
commit a66dd7dda0
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
5 changed files with 188 additions and 30 deletions

View file

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

View file

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

View file

View file

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

View file

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