fix(realtime): preserve provider duration and model capabilities

This commit is contained in:
Emerson Gomes 2026-09-22 18:59:55 -05:00
parent db3e6f06a0
commit 39ac30122e
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
11 changed files with 149 additions and 10 deletions

View file

@ -116,6 +116,11 @@ ARRAY_KEYS: dict[str, JsonSchema] = {
"description": "OpenAI-style API routes this model can be called through, e.g. /v1/chat/completions.",
"items": STRING,
},
"supported_transcription_response_formats": {
"type": "array",
"description": "Response formats accepted by the model for file transcription.",
"items": STRING,
},
"supported_modalities": {
"type": "array",
"description": "Input modalities the model accepts.",

View file

@ -176,7 +176,11 @@ async def receive_translation(
output.setframerate(SAMPLE_RATE)
write_stdout("Translation: ", end="", flush=True)
while True:
timeout = OUTPUT_IDLE_TIMEOUT_SECONDS if sender_finished.is_set() else INITIAL_RESPONSE_TIMEOUT_SECONDS
timeout = (
OUTPUT_IDLE_TIMEOUT_SECONDS
if sender_finished.is_set() and audio_received.is_set()
else INITIAL_RESPONSE_TIMEOUT_SECONDS
)
try:
raw_event = await asyncio.wait_for(connection.recv(), timeout=timeout)
except TimeoutError:

View file

@ -437,6 +437,21 @@ class RealTimeStreaming:
def _capture_translation_output_audio(self, event_obj: Mapping[str, object]) -> None:
if not self._is_translation_session:
return
if event_obj.get("type") == "session.closed":
usage: Final = event_obj.get("usage")
output_seconds: Final = usage.get("output_seconds") if isinstance(usage, dict) else None
if isinstance(output_seconds, (int, float)):
if not self._should_store_message(event_obj):
self.messages.append(
OpenAIRealtimeTranslationClosedEvent(
type="session.closed",
usage=OpenAIRealtimeTranslationDurationUsage(
type="duration", output_seconds=output_seconds
),
)
)
self._translation_usage_finalized = True
return
self._capture_translation_output_format(event_obj)
if event_obj.get("type") not in (
"session.output_audio.delta",
@ -1117,6 +1132,8 @@ class RealTimeStreaming:
for event in events:
if self._should_drop_event_from_client(event):
continue
if isinstance(event, dict):
self._capture_translation_output_audio(event)
is_session_created_event = isinstance(event, dict) and event.get("type") == "session.created"
if is_session_created_event:
if self._uses_deferred_backend_setup() and not self._backend_setup_complete:

View file

@ -23,7 +23,7 @@ from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Mappin
from concurrent import futures
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from copy import deepcopy
from functools import partial
from functools import lru_cache, partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, Union, cast, get_args
from urllib.parse import urlsplit
@ -39,7 +39,7 @@ import httpx
import openai
from openai import AsyncStream, Stream
from openai.types.audio import TranscriptionStreamEvent
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter
from typing_extensions import overload
import litellm
@ -7871,6 +7871,23 @@ async def atranscription(
)
@lru_cache(maxsize=1)
def _bundled_transcription_response_formats() -> Mapping[str, tuple[str, ...]]:
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
catalog: Final = TypeAdapter(dict[str, dict[str, object]]).validate_python(
GetModelCostMap.load_local_model_cost_map()
)
formats_adapter: Final = TypeAdapter(tuple[str, ...])
return MappingProxyType(
{
model: formats_adapter.validate_python(entry["supported_transcription_response_formats"])
for model, entry in catalog.items()
if "supported_transcription_response_formats" in entry
}
)
def _validate_gpt_transcription_request(
model: str,
custom_llm_provider: str,
@ -7879,25 +7896,30 @@ def _validate_gpt_transcription_request(
response_format: str | None,
api_version: str | None,
) -> str | None:
model_info: Final = get_model_info(model=model) if model in litellm.model_cost else {}
supported_endpoints: Final = model_info.get("supported_endpoints")
supported_formats: Final = model_info.get("supported_transcription_response_formats") or (
_bundled_transcription_response_formats().get(model)
)
if language is not None and languages is not None:
raise litellm.UnsupportedParamsError(
message="language and languages cannot be used together",
model=model,
llm_provider=custom_llm_provider,
)
if model == "gpt-live-transcribe":
if supported_endpoints is not None and "/v1/audio/transcriptions" not in supported_endpoints:
raise litellm.UnsupportedParamsError(
message="gpt-live-transcribe is available through the Realtime API, not file transcription",
message=f"{model} is available through the Realtime API, not file transcription",
model=model,
llm_provider=custom_llm_provider,
)
if model == "gpt-transcribe" and response_format not in (None, "json"):
if supported_formats is not None and response_format is not None and response_format not in supported_formats:
raise litellm.UnsupportedParamsError(
message="gpt-transcribe only supports response_format='json'",
message=f"{model} only supports response_format={', '.join(repr(fmt) for fmt in supported_formats)}",
model=model,
llm_provider=custom_llm_provider,
)
if custom_llm_provider == "azure" and model == "gpt-transcribe":
if custom_llm_provider == "azure" and supported_formats is not None:
if api_version in (None, "v1", "latest", "preview"):
return litellm.AZURE_DEFAULT_API_VERSION
return api_version

View file

@ -214,6 +214,9 @@
"/v1/realtime",
"/v1/realtime/transcription_sessions"
],
"supported_transcription_response_formats": [
"json"
],
"supported_modalities": [
"audio",
"text"
@ -57555,6 +57558,9 @@
"/v1/audio/transcriptions",
"/v1/realtime/transcription_sessions"
],
"supported_transcription_response_formats": [
"json"
],
"supported_modalities": [
"text",
"audio"

View file

@ -396,6 +396,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
]
]
supported_endpoints: list[str] | None
supported_transcription_response_formats: ReadOnly[Sequence[str] | None]
use_openai_responses_path: bool | None
tpm: int | None
rpm: int | None

View file

@ -6235,6 +6235,9 @@ def _get_model_info_helper(
litellm_provider=_model_info.get("litellm_provider", custom_llm_provider),
mode=_model_info.get("mode"),
supported_endpoints=_model_info.get("supported_endpoints", None),
supported_transcription_response_formats=_model_info.get(
"supported_transcription_response_formats", None
),
supports_system_messages=_model_info.get("supports_system_messages", None),
supports_response_schema=_model_info.get("supports_response_schema", None),
supports_vision=_model_info.get("supports_vision", None),

View file

@ -214,6 +214,9 @@
"/v1/realtime",
"/v1/realtime/transcription_sessions"
],
"supported_transcription_response_formats": [
"json"
],
"supported_modalities": [
"audio",
"text"
@ -57555,6 +57558,9 @@
"/v1/audio/transcriptions",
"/v1/realtime/transcription_sessions"
],
"supported_transcription_response_formats": [
"json"
],
"supported_modalities": [
"text",
"audio"

View file

@ -864,6 +864,13 @@
"type": "string"
}
},
"supported_transcription_response_formats": {
"type": "array",
"description": "Response formats accepted by the model for file transcription.",
"items": {
"type": "string"
}
},
"supports_adaptive_thinking": {
"type": "boolean"
},

View file

@ -1073,7 +1073,7 @@ async def test_translation_session_update_rejects_disallowed_nested_transcriptio
translation_session=True,
)
with pytest.raises(Exception, match="Tried to access gpt-live-transcribe"):
with pytest.raises(Exception, match=r"gpt-live-transcribe.*not available"):
await streaming._send_to_backend(
json.dumps(
{
@ -2910,7 +2910,9 @@ def test_store_message_skips_pydantic_for_unlogged_audio_delta():
assert streaming.messages == []
@pytest.mark.parametrize("event_type", ["session.output_audio.delta", "response.output_audio.delta", "response.audio.delta"])
@pytest.mark.parametrize(
"event_type", ["session.output_audio.delta", "response.output_audio.delta", "response.audio.delta"]
)
def test_translation_audio_duration_is_finalized_once(event_type: str):
import base64
@ -2974,6 +2976,29 @@ def test_translation_does_not_duplicate_provider_duration_usage():
assert len(closed_events) == 1
def test_translation_prefers_provider_duration_over_audio_byte_estimate():
import base64
streaming = RealTimeStreaming(
websocket=MagicMock(),
backend_ws=MagicMock(),
logging_obj=MagicMock(),
model="gpt-realtime-translate",
translation_session=True,
)
streaming._capture_translation_output_audio(
{"type": "session.output_audio.delta", "delta": base64.b64encode(bytes(48000)).decode()}
)
streaming._capture_translation_output_audio(
{"type": "session.closed", "usage": {"type": "duration", "output_seconds": 0.5}}
)
streaming._finalize_translation_usage()
closed_events = [event for event in streaming.messages if event.get("type") == "session.closed"]
assert len(closed_events) == 1
assert closed_events[0]["usage"] == {"type": "duration", "output_seconds": 0.5}
@pytest.mark.asyncio
async def test_audio_delta_frame_parsed_at_most_once():
client_ws = _beta_client_ws()

View file

@ -0,0 +1,43 @@
import asyncio
import base64
import json
import wave
from pathlib import Path
from types import SimpleNamespace
from typing import Final, cast
import pytest
from websockets.asyncio.client import ClientConnection
from cookbook import gpt_realtime_translate as translate
@pytest.mark.asyncio
async def test_short_upload_waits_for_first_translated_audio(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None:
monkeypatch.setattr(translate, "OUTPUT_IDLE_TIMEOUT_SECONDS", 0.01)
monkeypatch.setattr(translate, "INITIAL_RESPONSE_TIMEOUT_SECONDS", 0.1)
audio: Final = bytes(480)
events: Final = iter(
(
{"type": "session.output_audio.delta", "delta": base64.b64encode(audio).decode()},
{"type": "error", "error": {"message": "session closed"}},
)
)
async def recv() -> str:
event: Final = next(events)
if event["type"] == "session.output_audio.delta":
await asyncio.sleep(0.03)
return json.dumps(event)
sender_finished: Final = asyncio.Event()
sender_finished.set()
output: Final = tmp_path / "translation.wav"
result: Final = await translate.receive_translation(
cast(ClientConnection, SimpleNamespace(recv=recv)), output, sender_finished
)
assert result == 'Realtime API error: {"message": "session closed"}'
with wave.open(str(output), "rb") as rendered:
assert rendered.readframes(240) == audio