mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(realtime): preserve translation deployment pricing and input usage
This commit is contained in:
parent
152b1e15d2
commit
a66dd7dda0
5 changed files with 188 additions and 30 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
0
tests/unit/cookbook/__init__.py
Normal file
0
tests/unit/cookbook/__init__.py
Normal 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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue