mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(realtime): preserve provider duration and model capabilities
This commit is contained in:
parent
db3e6f06a0
commit
39ac30122e
11 changed files with 149 additions and 10 deletions
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
43
tests/unit/cookbook/test_gpt_realtime_translate.py
Normal file
43
tests/unit/cookbook/test_gpt_realtime_translate.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue