feat(realtime): support latest OpenAI audio models

This commit is contained in:
Emerson Gomes 2026-09-15 12:12:59 -05:00
parent bc3b5b1d5b
commit bd43c233ef
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
50 changed files with 3744 additions and 496 deletions

View file

@ -0,0 +1,242 @@
#!/usr/bin/env python3
import argparse
import asyncio
import base64
import json
import os
import sys
import wave
from collections.abc import Iterator, Sequence
from dataclasses import dataclass
from pathlib import Path
from urllib.parse import urlencode, urlsplit, urlunsplit
import websockets
from websockets.asyncio.client import ClientConnection
SAMPLE_RATE = 24_000
CHANNELS = 1
SAMPLE_WIDTH = 2
CHUNK_DURATION_SECONDS = 0.1
CHUNK_BYTES = int(SAMPLE_RATE * CHANNELS * SAMPLE_WIDTH * CHUNK_DURATION_SECONDS)
OUTPUT_IDLE_TIMEOUT_SECONDS = 3.0
INITIAL_RESPONSE_TIMEOUT_SECONDS = 30.0
AUDIO_EVENT_TYPES = frozenset(
{
"session.output_audio.delta",
"response.audio.delta",
"response.output_audio.delta",
}
)
TRANSCRIPT_EVENT_TYPES = frozenset(
{
"session.output_transcript.delta",
"response.text.delta",
"response.output_audio_transcript.delta",
}
)
@dataclass(frozen=True, slots=True)
class Settings:
input_wav: Path
output_wav: Path
base_url: str
model: str
target_language: str
trailing_silence_seconds: float
api_key: str
def write_stdout(message: str = "", *, end: str = "\n", flush: bool = False) -> None:
sys.stdout.write(f"{message}{end}")
if flush:
sys.stdout.flush()
def write_stderr(message: str) -> None:
sys.stderr.write(f"{message}\n")
def parse_args(argv: Sequence[str] | None = None) -> Settings | str:
parser = argparse.ArgumentParser(
description="Stream a 24 kHz PCM16 WAV through gpt-realtime-translate and save the translated audio",
)
parser.add_argument("input_wav", type=Path)
parser.add_argument("--output", type=Path, default=Path("translated.wav"))
parser.add_argument("--base-url", default=os.getenv("LITELLM_BASE_URL", "http://localhost:4000"))
parser.add_argument("--model", default=os.getenv("REALTIME_TRANSLATE_MODEL", "gpt-realtime-translate"))
parser.add_argument("--target-language", default="fr")
parser.add_argument("--trailing-silence", type=float, default=1.5)
parsed = parser.parse_args(argv)
api_key = os.getenv("LITELLM_API_KEY") or os.getenv("OPENAI_API_KEY")
if not api_key:
return "Set LITELLM_API_KEY or OPENAI_API_KEY before running the script"
if parsed.trailing_silence < 0:
return "--trailing-silence must be zero or greater"
return Settings(
input_wav=parsed.input_wav,
output_wav=parsed.output,
base_url=parsed.base_url,
model=parsed.model,
target_language=parsed.target_language,
trailing_silence_seconds=parsed.trailing_silence,
api_key=api_key,
)
def translation_url(base_url: str, model: str) -> str | None:
parsed = urlsplit(base_url.rstrip("/"))
scheme = {"http": "ws", "https": "wss", "ws": "ws", "wss": "wss"}.get(parsed.scheme)
if not scheme or not parsed.netloc:
return None
base_path = parsed.path.rstrip("/")
realtime_path = (
f"{base_path}/realtime/translations" if base_path.endswith("/v1") else f"{base_path}/v1/realtime/translations"
)
return urlunsplit((scheme, parsed.netloc, realtime_path, urlencode({"model": model}), ""))
def read_pcm16_wav(path: Path) -> bytes | str:
try:
with wave.open(str(path), "rb") as source:
actual_format = (
source.getnchannels(),
source.getsampwidth(),
source.getframerate(),
source.getcomptype(),
)
expected_format = (CHANNELS, SAMPLE_WIDTH, SAMPLE_RATE, "NONE")
if actual_format != expected_format:
return (
f"{path} must be mono, 16-bit PCM, 24 kHz WAV; received "
f"channels={actual_format[0]}, sample_width={actual_format[1]}, "
f"sample_rate={actual_format[2]}, compression={actual_format[3]}"
)
return source.readframes(source.getnframes())
except (OSError, EOFError, wave.Error) as exc:
return f"Unable to read {path}: {exc}"
def audio_chunks(audio: bytes) -> Iterator[bytes]:
return (audio[offset : offset + CHUNK_BYTES] for offset in range(0, len(audio), CHUNK_BYTES))
def audio_message(audio: bytes) -> str:
return json.dumps(
{
"type": "session.input_audio_buffer.append",
"audio": base64.b64encode(audio).decode("ascii"),
}
)
async def configure_session(connection: ClientConnection, target_language: str) -> str | None:
await connection.send(
json.dumps(
{
"type": "session.update",
"session": {"audio": {"output": {"language": target_language}}},
}
)
)
while True:
raw_event = await asyncio.wait_for(connection.recv(), timeout=20)
event = json.loads(raw_event)
event_type = event.get("type")
if event_type == "session.created":
write_stdout(f"Session: {event.get('session', {}).get('id', 'created')}")
if event_type == "session.updated":
return None
if event_type == "error":
return f"Realtime API error: {json.dumps(event.get('error', event), ensure_ascii=False)}"
async def send_audio(
connection: ClientConnection, pcm: bytes, trailing_silence_seconds: float, finished: asyncio.Event
) -> None:
silence = bytes(round(trailing_silence_seconds * SAMPLE_RATE * CHANNELS * SAMPLE_WIDTH))
try:
for chunk in audio_chunks(pcm + silence):
await connection.send(audio_message(chunk))
await asyncio.sleep(len(chunk) / (SAMPLE_RATE * CHANNELS * SAMPLE_WIDTH))
finally:
finished.set()
async def receive_translation(
connection: ClientConnection, output_path: Path, sender_finished: asyncio.Event
) -> str | None:
audio_received = asyncio.Event()
try:
with wave.open(str(output_path), "wb") as output:
output.setnchannels(CHANNELS)
output.setsampwidth(SAMPLE_WIDTH)
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
try:
raw_event = await asyncio.wait_for(connection.recv(), timeout=timeout)
except TimeoutError:
if sender_finished.is_set() and audio_received.is_set():
write_stdout()
return None
return "The translation stream ended without translated audio"
event = json.loads(raw_event)
event_type = event.get("type")
if event_type in AUDIO_EVENT_TYPES:
output.writeframes(base64.b64decode(event.get("delta", ""), validate=True))
audio_received.set()
elif event_type in TRANSCRIPT_EVENT_TYPES:
write_stdout(event.get("delta", event.get("text", "")), end="", flush=True)
elif event_type == "error":
return f"Realtime API error: {json.dumps(event.get('error', event), ensure_ascii=False)}"
except (OSError, wave.Error) as exc:
return f"Unable to write {output_path}: {exc}"
async def translate(settings: Settings, pcm: bytes) -> str | None:
url = translation_url(settings.base_url, settings.model)
if not url:
return f"Invalid --base-url: {settings.base_url}"
sender_finished = asyncio.Event()
try:
async with websockets.connect(
url,
additional_headers={"Authorization": f"Bearer {settings.api_key}"},
proxy=None,
open_timeout=20,
close_timeout=5,
) as connection:
configuration_error = await configure_session(connection, settings.target_language)
if configuration_error:
return configuration_error
async with asyncio.TaskGroup() as tasks:
receiver = tasks.create_task(receive_translation(connection, settings.output_wav, sender_finished))
tasks.create_task(send_audio(connection, pcm, settings.trailing_silence_seconds, sender_finished))
return receiver.result()
except Exception as exc:
return f"Translation failed: {type(exc).__name__}: {exc}"
def main(argv: Sequence[str] | None = None) -> int:
settings = parse_args(argv)
if isinstance(settings, str):
write_stderr(settings)
return 2
pcm = read_pcm16_wav(settings.input_wav)
if isinstance(pcm, str):
write_stderr(pcm)
return 2
error = asyncio.run(translate(settings, pcm))
if error:
write_stderr(error)
return 1
write_stdout(f"Translated audio: {settings.output_wav}")
return 0
if __name__ == "__main__":
raise SystemExit(main())

View file

@ -29,7 +29,7 @@ def _dev_env_hot_reload_enabled() -> bool:
if os.getenv("LITELLM_MODE", "DEV") == "DEV":
_dotenv.load_dotenv(override=_dev_env_hot_reload_enabled())
from collections.abc import Mapping, Sequence
from collections.abc import Sequence
from typing import (
Any,
Callable,
@ -72,6 +72,7 @@ from litellm.constants import (
OPENAI_CHAT_COMPLETION_PARAMS as _openai_completion_params, # backwards compatibility
OPENAI_FINISH_REASONS,
OPENAI_FINISH_REASONS as _openai_finish_reasons, # backwards compatibility
OPENAI_REALTIME_AND_TRANSCRIPTION_MODELS,
openai_compatible_endpoints,
openai_compatible_providers,
openai_text_completion_compatible_providers,
@ -1007,6 +1008,7 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
_populate_provider_model_sets(model_cost)
open_ai_chat_completion_models.update(OPENAI_REALTIME_AND_TRANSCRIPTION_MODELS)
# known openai compatible endpoints - we'll eventually move this list to the model_prices_and_context_window.json dictionary
# this is maintained for Exception Mapping
@ -1466,7 +1468,9 @@ from .realtime_api.main import (
_arealtime,
acreate_realtime_client_secret,
acreate_realtime_transcription_session,
acreate_realtime_translation_client_secret,
arealtime_calls,
arealtime_translation_calls,
)
from .responses.main import _aresponses_websocket
from .fine_tuning.main import *
@ -1650,6 +1654,9 @@ if TYPE_CHECKING:
from .llms.vertex_ai.rerank.transformation import (
VertexAIRerankConfig as VertexAIRerankConfig,
)
from .llms.together_ai.chat.transformation import (
TogetherAIChatConfig as TogetherAIChatConfig,
)
from .llms.fireworks_ai.rerank.transformation import (
FireworksAIRerankConfig as FireworksAIRerankConfig,
)
@ -1695,9 +1702,6 @@ if TYPE_CHECKING:
BedrockMantleAnthropicMessagesConfig as BedrockMantleAnthropicMessagesConfig,
)
from .llms.together_ai.chat import TogetherAIConfig as TogetherAIConfig
from .llms.together_ai.chat.transformation import (
TogetherAIChatConfig as TogetherAIChatConfig,
)
from .llms.nlp_cloud.chat.handler import NLPCloudConfig as NLPCloudConfig
from .llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig as VertexGeminiConfig,
@ -1853,6 +1857,9 @@ if TYPE_CHECKING:
from .llms.xai.responses.transformation import (
XAIResponsesAPIConfig as XAIResponsesAPIConfig,
)
from .llms.vertex_ai.interactions.transformation import (
VertexAIInteractionsConfig as VertexAIInteractionsConfig,
)
from .llms.litellm_proxy.responses.transformation import (
LiteLLMProxyResponsesAPIConfig as LiteLLMProxyResponsesAPIConfig,
)
@ -1877,9 +1884,6 @@ if TYPE_CHECKING:
from .llms.gemini.interactions.transformation import (
GoogleAIStudioInteractionsConfig as GoogleAIStudioInteractionsConfig,
)
from .llms.vertex_ai.interactions.transformation import (
VertexAIInteractionsConfig as VertexAIInteractionsConfig,
)
from .llms.openai.chat.o_series_transformation import (
OpenAIOSeriesConfig as OpenAIOSeriesConfig,
OpenAIOSeriesConfig as OpenAIO1Config,

View file

@ -828,10 +828,44 @@ OPENAI_CHAT_COMPLETION_PARAMS: Final = [
OPENAI_TRANSCRIPTION_PARAMS: Final = [
"language",
"languages",
"keywords",
"response_format",
"stream",
"timestamp_granularities",
]
OPENAI_REALTIME_AND_TRANSCRIPTION_MODELS: Final = frozenset(
{
"gpt-realtime-2",
"gpt-realtime-2.1",
"gpt-realtime-2.1-mini",
"gpt-realtime-translate",
"gpt-realtime-whisper",
"gpt-transcribe",
"gpt-live-transcribe",
}
)
AZURE_GA_REALTIME_MODELS: Final = frozenset(
{
"gpt-realtime-2",
"gpt-realtime-2-2026-05-06",
"gpt-realtime-2.1",
"gpt-realtime-2.1-2026-07-07",
"gpt-realtime-2.1-mini",
"gpt-realtime-2.1-mini-2026-07-07",
"gpt-realtime-translate",
"gpt-realtime-translate-2026-05-06",
"gpt-realtime-translate-2026-05-07",
"gpt-realtime-whisper",
"gpt-realtime-whisper-2026-05-06",
"gpt-realtime-whisper-2026-05-07",
"gpt-transcribe",
"gpt-live-transcribe",
}
)
OPENAI_EMBEDDING_PARAMS: Final = ["dimensions", "encoding_format", "user"]
DEFAULT_EMBEDDING_PARAM_VALUES: Final = {

View file

@ -1006,6 +1006,25 @@ def get_usage_object(
return None
def _get_transcription_usage_duration(completion_response: object) -> float | None:
usage_object: Final = (
completion_response.get("usage")
if isinstance(completion_response, dict)
else getattr(completion_response, "usage", None)
)
usage_type: Final = (
usage_object.get("type") if isinstance(usage_object, dict) else getattr(usage_object, "type", None)
)
if usage_type != "duration":
return None
seconds: Final = (
usage_object.get("seconds") if isinstance(usage_object, dict) else getattr(usage_object, "seconds", None)
)
if isinstance(seconds, bool) or not isinstance(seconds, (int, float)) or seconds < 0:
return None
return float(seconds)
def _is_known_usage_objects(usage_obj):
"""Returns True if the usage obj is a known Usage type"""
return (
@ -1595,9 +1614,14 @@ def completion_cost(
# the response attribute (for verbose_json responses that
# naturally include duration from the provider).
_hidden = getattr(completion_response, "_hidden_params", {}) or {}
audio_transcription_file_duration = _hidden.get(
"audio_transcription_duration",
getattr(completion_response, "duration", 0.0),
provider_duration = _get_transcription_usage_duration(completion_response)
audio_transcription_file_duration = (
provider_duration
if provider_duration is not None
else _hidden.get(
"audio_transcription_duration",
getattr(completion_response, "duration", 0.0),
)
)
elif call_type in _RERANK_CALL_TYPES:
if completion_response is not None and isinstance(completion_response, RerankResponse):
@ -2848,6 +2872,7 @@ class ResponsesWebSocketTokenUsageProcessor(BaseTokenUsageProcessor):
_TRANSCRIPTION_COMPLETED_EVENT_TYPE: Final = "conversation.item.input_audio_transcription.completed"
_TRANSLATION_CLOSED_EVENT_TYPE: Final = "session.closed"
def _candidate_realtime_token_costs(
@ -2947,7 +2972,21 @@ def handle_realtime_stream_cost_calculation(
if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results)
else 0.0
)
total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost
translation_cost: Final = handle_realtime_translation_cost_calculation(
results=results,
custom_llm_provider=custom_llm_provider,
litellm_model_name=litellm_model_name,
)
total_cost: Final = input_cost_per_token + output_cost_per_token + transcription_cost + translation_cost
additional_costs: Final = { # mutable-ok: logging stores a mutable per-request cost breakdown
key: value
for key, value in (
("transcription_cost", transcription_cost),
("translation_cost", translation_cost),
)
if value > 0
}
_store_cost_breakdown_in_logging_obj(
litellm_logging_obj=litellm_logging_obj,
@ -2955,13 +2994,40 @@ def handle_realtime_stream_cost_calculation(
completion_tokens_cost_usd_dollar=output_cost_per_token,
cost_for_built_in_tools_cost_usd_dollar=0.0,
total_cost_usd_dollar=total_cost,
additional_costs={"transcription_cost": transcription_cost} if transcription_cost > 0 else None,
data_residency=data_residency,
additional_costs=additional_costs or None,
)
return total_cost
def handle_realtime_translation_cost_calculation(
results: OpenAIRealtimeStreamList,
custom_llm_provider: str,
litellm_model_name: str,
) -> float:
output_seconds = 0.0 # rebind-ok: duration is accumulated across translation close events
for result in results:
if result.get("type") != _TRANSLATION_CLOSED_EVENT_TYPE:
continue
usage = result.get("usage")
if isinstance(usage, dict) and isinstance(usage.get("output_seconds"), (int, float)):
output_seconds += float(usage["output_seconds"])
if 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
output_cost_per_second: Final = model_info.get("output_cost_per_second")
if not isinstance(output_cost_per_second, (int, float)):
return 0.0
return output_seconds * output_cost_per_second
def handle_realtime_transcription_cost_calculation(
results: OpenAIRealtimeStreamList,
custom_llm_provider: str,

View file

@ -0,0 +1,204 @@
from __future__ import annotations
import datetime
import traceback
from collections.abc import AsyncIterator, Iterator
from typing import Final, Protocol
from openai import AsyncStream, Stream
from openai.types.audio import (
TranscriptionStreamEvent,
TranscriptionTextDeltaEvent,
TranscriptionTextDoneEvent,
)
from litellm.types.utils import TranscriptionResponse
class TranscriptionStreamLogging(Protocol):
def success_handler(
self,
result: TranscriptionResponse,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
async def async_success_handler(
self,
result: TranscriptionResponse,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def handle_sync_success_callbacks_for_async_calls(
self,
result: TranscriptionResponse,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
def failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
async def async_failure_handler(
self,
exception: Exception,
traceback_exception: str,
start_time: datetime.datetime,
end_time: datetime.datetime,
) -> None: ...
class _TranscriptionEventCollector:
def __init__(self, duration: float | None) -> None:
self.duration = duration
self.text_deltas: list[str] = [] # mutable-ok: streaming deltas accumulate until the terminal event
self.done_event: TranscriptionTextDoneEvent | None = None
def add(self, event: TranscriptionStreamEvent) -> None:
if isinstance(event, TranscriptionTextDeltaEvent):
self.text_deltas.append(event.delta)
elif isinstance(event, TranscriptionTextDoneEvent):
self.done_event = event
def response(self) -> TranscriptionResponse:
done_event: Final = self.done_event
done_languages: Final = getattr(done_event, "languages", None) if done_event is not None else None
response: Final = TranscriptionResponse(
text=done_event.text if done_event is not None else "".join(self.text_deltas),
usage=done_event.usage.model_dump() if done_event is not None and done_event.usage is not None else None,
languages=(
[ # mutable-ok: the response model requires a concrete serialized language list
language.model_dump() for language in done_languages
]
if done_languages is not None
else None
),
)
if self.duration is not None:
response.set_audio_transcription_duration(self.duration)
return response
class LoggingTranscriptionStream(Stream[TranscriptionStreamEvent]):
def __init__(
self,
stream: Stream[TranscriptionStreamEvent],
logging_obj: TranscriptionStreamLogging,
start_time: datetime.datetime,
) -> None:
self.__dict__.update(stream.__dict__)
self._logging_obj = logging_obj
self._start_time = start_time
self._collector = _TranscriptionEventCollector(getattr(stream, "_litellm_audio_duration", None))
self._finalized = False
self._failed = False
source_iterator: Final = self._iterator
self._iterator = self._logging_iterator(source_iterator)
def _logging_iterator(
self, source_iterator: Iterator[TranscriptionStreamEvent]
) -> Iterator[TranscriptionStreamEvent]:
try:
for event in source_iterator:
self._collector.add(event)
yield event
except Exception as exception:
self._failed = True
end_time: Final = datetime.datetime.now() # noqa: DTZ005 # callback timestamps use the legacy naive contract
self._logging_obj.failure_handler(exception, traceback.format_exc(), self._start_time, end_time)
raise
finally:
self._finalize()
def _finalize(self) -> None:
if self._finalized or self._failed:
return
self._finalized = True
self._logging_obj.success_handler(
self._collector.response(),
self._start_time,
datetime.datetime.now(), # noqa: DTZ005 # callback timestamps use the legacy naive contract
)
def close(self) -> None:
try:
super().close()
finally:
self._finalize()
class LoggingAsyncTranscriptionStream(AsyncStream[TranscriptionStreamEvent]):
def __init__(
self,
stream: AsyncStream[TranscriptionStreamEvent],
logging_obj: TranscriptionStreamLogging,
start_time: datetime.datetime,
) -> None:
self.__dict__.update(stream.__dict__)
self._logging_obj = logging_obj
self._start_time = start_time
self._collector = _TranscriptionEventCollector(getattr(stream, "_litellm_audio_duration", None))
self._finalized = False
self._failed = False
source_iterator: Final = self._iterator
self._iterator = self._logging_iterator(source_iterator)
async def _logging_iterator(
self, source_iterator: AsyncIterator[TranscriptionStreamEvent]
) -> AsyncIterator[TranscriptionStreamEvent]:
try:
async for event in source_iterator:
self._collector.add(event)
yield event
except Exception as exception:
self._failed = True
end_time: Final = datetime.datetime.now() # noqa: DTZ005 # callback timestamps use the legacy naive contract
self._logging_obj.failure_handler(exception, traceback.format_exc(), self._start_time, end_time)
await self._logging_obj.async_failure_handler(
exception,
traceback.format_exc(),
self._start_time,
end_time,
)
raise
finally:
await self._finalize()
async def _finalize(self) -> None:
if self._finalized or self._failed:
return
self._finalized = True
response: Final = self._collector.response()
end_time: Final = datetime.datetime.now() # noqa: DTZ005 # callback timestamps use the legacy naive contract
self._logging_obj.handle_sync_success_callbacks_for_async_calls(
result=response,
start_time=self._start_time,
end_time=end_time,
)
await self._logging_obj.async_success_handler(
result=response,
start_time=self._start_time,
end_time=end_time,
)
async def close(self) -> None:
try:
await super().close()
finally:
await self._finalize()
def wrap_transcription_stream(
stream: Stream[TranscriptionStreamEvent] | AsyncStream[TranscriptionStreamEvent],
logging_obj: TranscriptionStreamLogging,
start_time: datetime.datetime,
) -> LoggingTranscriptionStream | LoggingAsyncTranscriptionStream:
if isinstance(stream, AsyncStream):
return LoggingAsyncTranscriptionStream(stream, logging_obj, start_time)
return LoggingTranscriptionStream(stream, logging_obj, start_time)

View file

@ -910,6 +910,11 @@ def calculate_cache_writing_cost(
class PromptTokensDetailsResult(TypedDict):
cache_hit_tokens: int
cache_hit_audio_tokens: ReadOnly[int]
cached_text_tokens: ReadOnly[int]
cached_audio_tokens: ReadOnly[int]
cached_image_tokens: ReadOnly[int]
has_cached_tokens_details: ReadOnly[bool]
cache_creation_tokens: int
cache_creation_token_details: CacheCreationTokenDetails | None
text_tokens: int
@ -996,6 +1001,10 @@ def parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult:
return PromptTokensDetailsResult(
cache_hit_tokens=cache_hit_tokens,
cache_hit_audio_tokens=cached_audio_tokens,
cached_text_tokens=cached_text_tokens,
cached_audio_tokens=cached_audio_tokens,
cached_image_tokens=cached_image_tokens,
has_cached_tokens_details=cached_tokens_details is not None,
cache_creation_tokens=cache_creation_tokens,
cache_creation_token_details=cache_creation_token_details,
text_tokens=text_tokens,
@ -1079,15 +1088,11 @@ def _calculate_input_cost(
prompt_cost = float(prompt_tokens_details["text_tokens"]) * prompt_base_cost
### CACHE READ COST - Now uses tiered pricing
cache_hit_audio_tokens: Final = prompt_tokens_details["cache_hit_audio_tokens"]
audio_cache_read_rate: Final = _get_cost_per_unit(
model_info,
_get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier),
None,
)
prompt_cost += float(prompt_tokens_details["cache_hit_tokens"] - cache_hit_audio_tokens) * cache_read_cost
prompt_cost += float(cache_hit_audio_tokens) * (
audio_cache_read_rate if audio_cache_read_rate is not None else cache_read_cost
prompt_cost += _calculate_cache_read_cost(
prompt_tokens_details=prompt_tokens_details,
model_info=model_info,
cache_read_cost=cache_read_cost,
service_tier=service_tier,
)
### AUDIO COST
@ -1167,6 +1172,38 @@ def _calculate_input_cost(
return prompt_cost
def _calculate_cache_read_cost(
prompt_tokens_details: PromptTokensDetailsResult,
model_info: ModelInfo,
cache_read_cost: float,
service_tier: str | None,
) -> float:
cached_text_tokens: Final = prompt_tokens_details["cached_text_tokens"]
cached_audio_tokens: Final = prompt_tokens_details["cached_audio_tokens"]
cached_image_tokens: Final = prompt_tokens_details["cached_image_tokens"]
classified_cached_tokens: Final = cached_text_tokens + cached_audio_tokens + cached_image_tokens
unclassified_cached_tokens: Final = max(prompt_tokens_details["cache_hit_tokens"] - classified_cached_tokens, 0)
total_cost = ( # rebind-ok: cached modality components accumulate into one cache-read cost
float(cached_text_tokens + unclassified_cached_tokens) * cache_read_cost
)
if cached_audio_tokens:
cached_audio_cost_key: Final = _get_service_tier_cost_key("cache_read_input_audio_token_cost", service_tier)
cached_audio_cost: Final = _get_cost_per_unit(model_info, cached_audio_cost_key, cache_read_cost)
total_cost += ( # rebind-ok: cached audio contributes to cache-read cost
float(cached_audio_tokens) * float(cached_audio_cost or 0.0)
)
if cached_image_tokens:
cached_image_cost_key: Final = _get_service_tier_cost_key("cache_read_input_image_token_cost", service_tier)
cached_image_cost: Final = _get_cost_per_unit(model_info, cached_image_cost_key, cache_read_cost)
total_cost += ( # rebind-ok: cached images contribute to cache-read cost
float(cached_image_tokens) * float(cached_image_cost or 0.0)
)
return total_cost
def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | None) -> float:
"""
Resolve the per-model regional-processing uplift multiplier for a given
@ -1328,6 +1365,10 @@ def generic_cost_per_token(
prompt_tokens_details = PromptTokensDetailsResult(
cache_hit_tokens=0,
cache_hit_audio_tokens=0,
cached_text_tokens=0,
cached_audio_tokens=0,
cached_image_tokens=0,
has_cached_tokens_details=False,
cache_creation_tokens=0,
cache_creation_token_details=None,
text_tokens=usage.prompt_tokens,
@ -1502,6 +1543,7 @@ class BilledTokenRates:
cache_creation_input_token_cost: float
cache_creation_input_token_cost_above_1hr: float
output_cost_per_reasoning_token: float
cache_read_input_image_token_cost: float | None = None
def scaled(self, multiplier: float) -> "BilledTokenRates":
if multiplier == 1.0:
@ -1514,6 +1556,11 @@ class BilledTokenRates:
cache_creation_input_token_cost=self.cache_creation_input_token_cost * multiplier,
cache_creation_input_token_cost_above_1hr=self.cache_creation_input_token_cost_above_1hr * multiplier,
output_cost_per_reasoning_token=self.output_cost_per_reasoning_token * multiplier,
cache_read_input_image_token_cost=(
self.cache_read_input_image_token_cost * multiplier
if self.cache_read_input_image_token_cost is not None
else None
),
)
@ -1617,6 +1664,11 @@ def _cost_map_billed_rates(
cache_creation_input_token_cost=cache_creation_cost_rate,
cache_creation_input_token_cost_above_1hr=cache_creation_cost_above_1hr_rate,
output_cost_per_reasoning_token=reasoning_rate,
cache_read_input_image_token_cost=_get_cost_per_unit(
model_info,
_get_service_tier_cost_key("cache_read_input_image_token_cost", service_tier),
None,
),
).scaled(multiplier)
@ -1692,6 +1744,12 @@ def get_token_type_cost_breakdown(
cache_read_tokens, cached_audio_tokens, cache_creation_tokens, cache_creation_token_details = _cache_token_counts(
usage
)
cached_image_tokens: Final = parse_prompt_tokens_details(usage)["cached_image_tokens"]
image_cache_read_rate: Final = (
rates.cache_read_input_image_token_cost
if rates.cache_read_input_image_token_cost is not None
else rates.cache_read_input_token_cost
)
cache_creation_cost: Final = (
float(cache_creation_tokens) * rates.cache_creation_input_token_cost
if custom_cost_per_token is not None
@ -1705,8 +1763,9 @@ def get_token_type_cost_breakdown(
return TokenTypeCostBreakdown(
reasoning_cost=float(_reasoning_token_count(usage)) * rates.output_cost_per_reasoning_token,
cache_read_cost=(
float(cache_read_tokens - cached_audio_tokens) * rates.cache_read_input_token_cost
float(cache_read_tokens - cached_audio_tokens - cached_image_tokens) * rates.cache_read_input_token_cost
+ float(cached_audio_tokens) * rates.cache_read_input_audio_token_cost
+ float(cached_image_tokens) * image_cache_read_rate
),
cache_creation_cost=cache_creation_cost,
rates=rates,

View file

@ -38,6 +38,7 @@ from litellm.types.utils import (
StreamingChoices,
TextChoices,
TextCompletionResponse,
TranscriptionDetectedLanguage,
TranscriptionResponse,
TranscriptionUsageDurationObject,
TranscriptionUsageTokensObject,
@ -772,9 +773,11 @@ def convert_to_model_response_object(
model_response_object.data = response_object["data"]
if "usage" in response_object and response_object["usage"] is not None:
model_response_object.usage.completion_tokens = response_object["usage"].get("completion_tokens", 0)
model_response_object.usage.prompt_tokens = response_object["usage"].get("prompt_tokens", 0)
model_response_object.usage.total_tokens = response_object["usage"].get("total_tokens", 0)
embedding_usage: Final = model_response_object.usage or Usage()
embedding_usage.completion_tokens = response_object["usage"].get("completion_tokens", 0)
embedding_usage.prompt_tokens = response_object["usage"].get("prompt_tokens", 0)
embedding_usage.total_tokens = response_object["usage"].get("total_tokens", 0)
model_response_object.usage = embedding_usage
if start_time is not None and end_time is not None:
model_response_object._response_ms = (
@ -817,6 +820,12 @@ def convert_to_model_response_object(
if key in response_object:
setattr(model_response_object, key, response_object[key])
if "languages" in response_object and response_object["languages"] is not None:
transcription_response: Final = model_response_object
transcription_response.languages = tuple(
TranscriptionDetectedLanguage.model_validate(language) for language in response_object["languages"]
)
if "usage" in response_object and response_object["usage"] is not None:
tr_usage_object: TranscriptionUsageDurationObject | TranscriptionUsageTokensObject | None = None

View file

@ -1,4 +1,5 @@
import asyncio
import base64
import json
import traceback
from collections.abc import Coroutine, Mapping, Sequence
@ -19,6 +20,8 @@ from litellm.types.llms.openai import (
OpenAIRealtimeResponseDelta,
OpenAIRealtimeStreamResponseBaseObject,
OpenAIRealtimeStreamSessionEvents,
OpenAIRealtimeTranslationClosedEvent,
OpenAIRealtimeTranslationDurationUsage,
)
from litellm.types.realtime import ALL_DELTA_TYPES
@ -137,6 +140,7 @@ class RealTimeStreaming:
force_transcription_model: str | None = None,
event_normalizer: RealtimeEventNormalizer | None = None,
logging_worker: _LoggingWorker = GLOBAL_LOGGING_WORKER,
translation_session: bool = False,
):
self.websocket: _ClientWebSocket = websocket
self.backend_ws = backend_ws
@ -148,6 +152,10 @@ class RealTimeStreaming:
self.input_messages: list[dict[str, str]] = []
self.session_tools: list[dict] = []
self.tool_calls: list[dict] = []
self._is_translation_session = translation_session
self._translation_output_audio_bytes = 0
self._translation_output_bytes_per_second = 48000.0
self._translation_usage_finalized = False
# Detect whether the client is explicitly opting into the beta protocol.
self._client_wants_beta = self._detect_beta_header(websocket)
@ -196,6 +204,7 @@ class RealTimeStreaming:
# their input_audio_transcription.completed usage drives duration-based cost.
self._force_transcription_model = force_transcription_model
self._is_transcription_session: bool = force_transcription_model is not None
self._bound_nested_transcription_model: str | None = None
# Optional per-provider GA event normalizer (e.g. XAIRealtimeNormalizer).
self._event_normalizer = event_normalizer
@ -410,6 +419,7 @@ class RealTimeStreaming:
async def log_messages(self):
"""Log messages in list"""
self._finalize_translation_usage()
if self.logging_obj:
if self.input_messages:
self.logging_obj.model_call_details["messages"] = self.input_messages
@ -424,6 +434,60 @@ class RealTimeStreaming:
)
self.logging_obj.model_call_details[REALTIME_SESSION_SUCCESS_LOGGED_KEY] = True
def _capture_translation_output_audio(self, event_obj: Mapping[str, object]) -> None:
if not self._is_translation_session:
return
self._capture_translation_output_format(event_obj)
if event_obj.get("type") != "session.output_audio.delta":
return
delta: Final = event_obj.get("delta")
if not isinstance(delta, str):
return
try:
decoded: Final = base64.b64decode(delta, validate=True)
except (ValueError, TypeError):
return
self._translation_output_audio_bytes += len(decoded)
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")
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)
def _finalize_translation_usage(self) -> None:
if self._translation_usage_finalized:
return
for event in self.messages:
if event.get("type") != "session.closed":
continue
event_usage = event.get("usage") # rebind-ok: each close event carries independent usage
if isinstance(event_usage, dict) and isinstance(event_usage.get("output_seconds"), (int, float)):
self._translation_usage_finalized = True
return
if self._translation_output_audio_bytes == 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)
self.messages.append(OpenAIRealtimeTranslationClosedEvent(type="session.closed", usage=synthetic_usage))
self._translation_usage_finalized = True
async def _send_to_backend(self, message: str) -> bool:
"""Send a message to the backend WebSocket.
@ -436,7 +500,7 @@ class RealTimeStreaming:
backend, False if the provider transformation produced no output and
the message was effectively dropped.
"""
message = self._enforce_transcription_session_model(message)
message = await self._apply_nested_transcription_model_policy(message)
if self.provider_config:
transformed: Final = self.provider_config.transform_realtime_request(
message, self.model, self.session_configuration_request
@ -478,6 +542,90 @@ class RealTimeStreaming:
await self.backend_ws.send(message)
return True
async def _apply_nested_transcription_model_policy(self, message: str) -> str:
if self._force_transcription_model is not None:
return self._enforce_transcription_session_model(message)
if self._is_translation_session:
return await self._enforce_translation_nested_transcription_model(message)
return message
def _session_update_message_obj(self, message: str) -> Mapping[str, object] | None:
try:
message_obj: Final = _decode_json_object(message)
except (json.JSONDecodeError, TypeError):
return None
if message_obj.get("type") not in (
"session.update",
"transcription_session.update",
):
return None
return message_obj
def _nested_transcription_models_from_session(
self,
session: Mapping[str, object],
) -> tuple[str, ...]:
audio: Final = session.get("audio")
audio_input: Final = audio.get("input") if isinstance(audio, dict) else None
nested_transcription: Final = audio_input.get("transcription") if isinstance(audio_input, dict) else None
nested_model: Final = self._transcription_model_value(nested_transcription)
flat_model: Final = self._transcription_model_value(session.get("input_audio_transcription"))
return tuple(dict.fromkeys(model for model in (nested_model, flat_model) if model is not None))
def _transcription_model_value(self, transcription_config: object) -> str | None:
if not isinstance(transcription_config, dict):
return None
model: Final = transcription_config.get("model")
if isinstance(model, str) and model:
return model
return None
def _rewrite_session_update_transcription_model(self, message: str, authorized_model: str) -> str:
message_obj: Final = self._session_update_message_obj(message)
if message_obj is None:
return message
session: Final = message_obj.get("session")
if not isinstance(session, dict):
return message
transcription: Final = session.get("input_audio_transcription")
rewrite_flat: Final = isinstance(transcription, dict) and transcription.get("model") != authorized_model
if isinstance(transcription, dict) and rewrite_flat:
session["input_audio_transcription"] = {
**transcription,
"model": authorized_model,
}
audio: Final = session.get("audio")
audio_input: Final = audio.get("input") if isinstance(audio, dict) else None
nested_transcription: Final = audio_input.get("transcription") if isinstance(audio_input, dict) else None
rewrite_nested: Final = (
isinstance(audio, dict)
and isinstance(audio_input, dict)
and isinstance(nested_transcription, dict)
and nested_transcription.get("model") != authorized_model
)
if (
isinstance(audio, dict)
and isinstance(audio_input, dict)
and isinstance(nested_transcription, dict)
and rewrite_nested
):
session["audio"] = {
**audio,
"input": {
**audio_input,
"transcription": {
**nested_transcription,
"model": authorized_model,
},
},
}
if not rewrite_flat and not rewrite_nested:
return message
return json.dumps(message_obj)
def _enforce_transcription_session_model(self, message: str) -> str:
"""Force client transcription session updates to the authorized model.
@ -495,56 +643,49 @@ class RealTimeStreaming:
if self._force_transcription_model is None:
return message
try:
message_obj: Final = _decode_json_object(message)
except (json.JSONDecodeError, TypeError):
message_obj: Final = self._session_update_message_obj(message)
if message_obj is None:
return message
session: Final = message_obj.get("session")
if isinstance(session, dict) and session.get("type") == "transcription":
self._is_transcription_session = True
return self._rewrite_session_update_transcription_model(message, self._force_transcription_model)
if message_obj.get("type") not in (
"session.update",
"transcription_session.update",
):
async def _enforce_translation_nested_transcription_model(self, message: str) -> str:
if self._bound_nested_transcription_model is not None:
return self._rewrite_session_update_transcription_model(message, self._bound_nested_transcription_model)
message_obj: Final = self._session_update_message_obj(message)
if message_obj is None:
return message
session: Final = message_obj.get("session")
if not isinstance(session, dict):
return message
if session.get("type") == "transcription":
self._is_transcription_session = True
authorized_model: Final = self._force_transcription_model
changed = False
transcription: Final = session.get("input_audio_transcription")
if isinstance(transcription, dict) and transcription.get("model") != authorized_model:
session["input_audio_transcription"] = {
**transcription,
"model": authorized_model,
}
changed = True
audio: Final = session.get("audio")
if isinstance(audio, dict):
audio_input: Final = audio.get("input")
if isinstance(audio_input, dict):
nested_transcription: Final = audio_input.get("transcription")
if isinstance(nested_transcription, dict) and nested_transcription.get("model") != authorized_model:
session["audio"] = {
**audio,
"input": {
**audio_input,
"transcription": {
**nested_transcription,
"model": authorized_model,
},
},
}
changed = True
if not changed:
nested_models: Final = self._nested_transcription_models_from_session(session)
if not nested_models:
return message
return json.dumps(message_obj)
valid_token: Final = self.user_api_key_dict
if valid_token is None:
return message
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
from litellm.proxy.proxy_server import llm_model_list, llm_router
if not isinstance(valid_token, UserAPIKeyAuth):
return message
for nested_model in nested_models:
await can_key_call_resolved_model(
model=nested_model,
valid_token=valid_token,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
bound_model: Final = nested_models[0]
self._bound_nested_transcription_model = bound_model
return self._rewrite_session_update_transcription_model(message, bound_model)
def _uses_deferred_backend_setup(self) -> bool:
"""True when setup is deferred until the client's first session.update."""
@ -942,7 +1083,10 @@ class RealTimeStreaming:
async def _handle_provider_config_message(self, raw_response: str) -> None:
"""Process a backend message when a provider_config is set (transformed path)."""
returned_object: Final = self.provider_config.transform_realtime_response(
provider_config: Final = self.provider_config
if provider_config is None:
raise RuntimeError("Provider response handling requires a provider configuration")
returned_object: Final = provider_config.transform_realtime_response(
raw_response,
self.model,
self.logging_obj,
@ -1103,6 +1247,7 @@ class RealTimeStreaming:
if self._should_drop_event_from_client(event):
continue
self._capture_translation_output_audio(event)
if await self._handle_raw_backend_message(event, raw_response):
continue
@ -1507,6 +1652,8 @@ class RealTimeStreaming:
session = client_event.get("session", {})
if isinstance(session, dict):
session = self._remap_beta_session_to_ga(session)
if self._is_translation_session:
session.pop("type", None)
msg_obj["session"] = session
message = json.dumps(msg_obj)

View file

@ -1,7 +1,7 @@
from collections.abc import Coroutine
from typing import TYPE_CHECKING, Any, Final
from openai import AsyncAzureOpenAI, AzureOpenAI
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from pydantic import BaseModel
from litellm._uuid import uuid
@ -49,6 +49,7 @@ class AzureAudioTranscription(AzureChatCompletion):
timeout=timeout,
api_key=api_key,
api_base=api_base,
api_version=api_version,
client=client,
max_retries=max_retries,
logging_obj=logging_obj,
@ -66,7 +67,7 @@ class AzureAudioTranscription(AzureChatCompletion):
client=client,
litellm_params=litellm_params,
)
if not isinstance(azure_client, AzureOpenAI):
if not isinstance(azure_client, (AzureOpenAI, OpenAI)):
raise AzureOpenAIError(
status_code=500,
message="azure_client is not an instance of AzureOpenAI",
@ -85,10 +86,13 @@ class AzureAudioTranscription(AzureChatCompletion):
)
response: Final = azure_client.audio.transcriptions.create(
**data,
**data, # pyright: ignore[reportArgumentType] # SDK TypedDict lags accepted transcription options
timeout=timeout,
)
if data.get("stream") is True:
return response
if isinstance(response, BaseModel):
stringified_response = response.model_dump()
else:
@ -137,7 +141,7 @@ class AzureAudioTranscription(AzureChatCompletion):
client=client,
litellm_params=litellm_params,
)
if not isinstance(async_azure_client, AsyncAzureOpenAI):
if not isinstance(async_azure_client, (AsyncAzureOpenAI, AsyncOpenAI)):
raise AzureOpenAIError(
status_code=500,
message="async_azure_client is not an instance of AsyncAzureOpenAI",
@ -155,8 +159,15 @@ class AzureAudioTranscription(AzureChatCompletion):
},
)
if data.get("stream") is True:
return await async_azure_client.audio.transcriptions.create(
**data, # pyright: ignore[reportArgumentType] # SDK TypedDict lags accepted transcription options
timeout=timeout,
)
raw_response: Final = await async_azure_client.audio.transcriptions.with_raw_response.create(
**data, timeout=timeout
**data, # pyright: ignore[reportArgumentType] # SDK TypedDict lags accepted transcription options
timeout=timeout,
)
headers: Final = dict(raw_response.headers)

View file

@ -83,6 +83,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
api_version: str | None,
realtime_protocol: str | None = None,
query_params: RealtimeQueryParams | None = None,
realtime_mode: str = "realtime",
) -> str:
"""
Construct Azure realtime WebSocket URL.
@ -114,18 +115,26 @@ class AzureOpenAIRealtime(AzureChatCompletion):
)
intent: Final = (query_params or {}).get("intent")
if _is_ga:
path = "/openai/v1/realtime"
query_parts = []
if intent != "transcription" and (query_params is None or "model" in query_params):
query_parts.append(urlencode({"model": model}))
else:
# Default to beta path for backwards compatibility
path = "/openai/realtime"
query_parts = [urlencode({"api-version": api_version, "deployment": model})]
path: Final = (
"/openai/v1/realtime/translations"
if realtime_mode == "translation"
else "/openai/v1/realtime"
if _is_ga
else "/openai/realtime"
)
base_query_parts: Final = (
(urlencode((("model", model),)),)
if realtime_mode == "translation"
else (
(urlencode((("model", model),)),)
if intent != "transcription" and (query_params is None or "model" in query_params)
else ()
)
if _is_ga
else (urlencode((("api-version", api_version), ("deployment", model))),)
)
if intent:
query_parts.append(urlencode({"intent": intent}))
query_parts: Final = (*base_query_parts, urlencode((("intent", intent),))) if intent else base_query_parts
qs: Final = "&".join(query_parts)
return f"{api_base}{path}?{qs}" if qs else f"{api_base}{path}"
@ -145,6 +154,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
query_params: RealtimeQueryParams | None = None,
user_api_key_dict: object | None = None,
litellm_metadata: dict | None = None,
realtime_mode: str = "realtime",
):
import websockets
from websockets.asyncio.client import ClientConnection
@ -161,6 +171,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
api_version,
realtime_protocol=realtime_protocol,
query_params=query_params,
realtime_mode=realtime_mode,
)
auth_headers: Final = self.get_auth_headers(api_key=api_key, azure_ad_token=azure_ad_token)
@ -184,6 +195,7 @@ class AzureOpenAIRealtime(AzureChatCompletion):
force_transcription_model=(
model if (query_params or {}).get("intent") == "transcription" else None
),
translation_session=realtime_mode == "translation",
)
await realtime_streaming.bidirectional_forward()

View file

@ -3,11 +3,16 @@
from typing import Final
import litellm
from litellm.constants import AZURE_GA_REALTIME_MODELS
from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig
from litellm.secret_managers.main import get_secret_str
class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig):
@staticmethod
def _uses_ga_api(model: str, api_version: str | None) -> bool:
return api_version in ("preview", "latest", "v1") or model in AZURE_GA_REALTIME_MODELS
def get_api_base(self, api_base: str | None, **kwargs) -> str:
return api_base or litellm.api_base or get_secret_str("AZURE_API_BASE") or ""
@ -16,6 +21,8 @@ class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig):
def get_complete_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
base: Final = self.get_api_base(api_base).rstrip("/")
if self._uses_ga_api(model, api_version):
return f"{base}/openai/v1/realtime/client_secrets"
version: Final = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17"
return f"{base}/openai/realtime/client_secrets?api-version={version}"
@ -25,22 +32,38 @@ class AzureRealtimeHTTPConfig(BaseRealtimeHTTPConfig):
model: str,
api_key: str | None = None,
) -> dict:
return {
validated_headers: Final = { # mutable-ok: provider authentication headers are extended before dispatch
**headers,
"api-key": api_key or "",
"Content-Type": "application/json",
}
if api_key:
validated_headers["api-key"] = api_key
return validated_headers
def get_realtime_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
base: Final = self.get_api_base(api_base).rstrip("/")
if self._uses_ga_api(model, api_version):
return f"{base}/openai/v1/realtime/calls"
version: Final = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17"
return f"{base}/openai/realtime/calls?api-version={version}"
def get_transcription_session_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
base: Final = self.get_api_base(api_base).rstrip("/")
if self._uses_ga_api(model, api_version):
return f"{base}/openai/v1/realtime/transcription_sessions"
version: Final = api_version or get_secret_str("AZURE_API_VERSION") or "2024-12-17"
return f"{base}/openai/realtime/transcription_sessions?api-version={version}"
def get_translation_client_secret_url(
self, api_base: str | None, model: str, api_version: str | None = None
) -> str:
base: Final = self.get_api_base(api_base).rstrip("/")
return f"{base}/openai/v1/realtime/translations/client_secrets"
def get_translation_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
base: Final = self.get_api_base(api_base).rstrip("/")
return f"{base}/openai/v1/realtime/translations/calls"
def get_realtime_calls_headers(self, ephemeral_key: str) -> dict:
return {
"api-key": ephemeral_key,

View file

@ -63,6 +63,12 @@ class BaseRealtimeHTTPConfig(ABC):
base = base.removesuffix("/v1")
return f"{base}/v1/realtime/transcription_sessions"
def get_translation_client_secret_url(
self, api_base: str | None, model: str, api_version: str | None = None
) -> str:
base: Final = (api_base or "").rstrip("/")
return f"{base}/v1/realtime/translations/client_secrets"
@abstractmethod
def validate_environment(
self,
@ -86,6 +92,10 @@ class BaseRealtimeHTTPConfig(ABC):
base: Final = (api_base or "").rstrip("/")
return f"{base}/v1/realtime/calls"
def get_translation_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
base: Final = (api_base or "").rstrip("/")
return f"{base}/v1/realtime/translations/calls"
def get_realtime_calls_headers(self, ephemeral_key: str) -> dict:
"""
Build headers for the realtime_calls POST.

View file

@ -23,7 +23,9 @@ from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
import httpx
from httpx import USE_CLIENT_DEFAULT
from httpx._types import FileContent
from openai import AsyncOpenAI
from openai.types.file_deleted import FileDeleted
from openai.types.realtime import RealtimeSessionCreateRequestParam
import litellm
import litellm.litellm_core_utils
@ -6186,6 +6188,7 @@ class BaseLLMHTTPHandler:
"BasePassthroughConfig",
"BaseContainerConfig",
BaseEvalsAPIConfig,
BaseRealtimeHTTPConfig,
],
):
received_status_code: Final = (
@ -6423,9 +6426,10 @@ class BaseLLMHTTPHandler:
timeout: float | httpx.Timeout,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
extra_headers: dict[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
extra_headers: Mapping[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | AsyncOpenAI | None = None,
api_version: str | None = None,
use_openai_sdk: bool = False,
) -> httpx.Response:
"""
Forward POST /v1/realtime/client_secrets to upstream provider.
@ -6433,6 +6437,52 @@ class BaseLLMHTTPHandler:
Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and
header auth when available; falls back to the legacy OpenAI-style defaults.
"""
if use_openai_sdk:
trimmed_api_base: Final = api_base.rstrip("/")
normalized_api_base: Final = (
trimmed_api_base if trimmed_api_base.endswith("/v1") else f"{trimmed_api_base}/v1"
)
owns_client: Final = not isinstance(client, AsyncOpenAI)
openai_client: Final = (
client
if isinstance(client, AsyncOpenAI)
else AsyncOpenAI(api_key=api_key, base_url=normalized_api_base, max_retries=0)
)
logging_obj.pre_call(
input=request_data,
api_key="",
additional_args={ # mutable-ok: logging owns a mutable request metadata payload
"complete_input_dict": request_data,
"api_base": normalized_api_base,
},
)
try:
configured_client: Final = openai_client.with_options(
timeout=timeout,
set_default_headers={ # mutable-ok: OpenAI SDK accepts a mutable custom-header mapping
key: str(value) # mutable-ok: SDK headers are materialized as a concrete string mapping
for key, value in (extra_headers or {}).items() # mutable-ok: SDK requires concrete headers
},
)
raw_response: Final = await configured_client.post(
"/realtime/client_secrets",
cast_to=httpx.Response,
body=request_data,
)
response_headers: Final = { # mutable-ok: httpx requires a concrete response-header mapping
key: value # mutable-ok: transport headers are materialized after filtering
for key, value in raw_response.headers.items() # mutable-ok: transport headers are materialized
if key.lower() not in ("content-encoding", "content-length", "transfer-encoding")
}
return httpx.Response(
status_code=raw_response.status_code,
headers=response_headers,
content=raw_response.content,
request=httpx.Request("POST", f"{normalized_api_base}/realtime/client_secrets"),
)
finally:
if owns_client:
await openai_client.close()
return await self._async_realtime_session_post(
endpoint="client_secrets",
api_base=api_base,
@ -6456,8 +6506,8 @@ class BaseLLMHTTPHandler:
timeout: float | httpx.Timeout,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
extra_headers: dict[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
extra_headers: Mapping[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | AsyncOpenAI | None = None,
api_version: str | None = None,
) -> httpx.Response:
"""Forward POST /v1/realtime/transcription_sessions to upstream provider."""
@ -6475,18 +6525,79 @@ class BaseLLMHTTPHandler:
api_version=api_version,
)
async def _async_realtime_session_post(
async def async_realtime_translation_client_secret_handler(
self,
endpoint: Literal["client_secrets", "transcription_sessions"],
api_base: str,
api_key: str,
request_data: dict[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
provider_config: Any | None = None,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
extra_headers: dict[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
extra_headers: Mapping[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | AsyncOpenAI | None = None,
api_version: str | None = None,
use_openai_sdk: bool = False,
) -> httpx.Response:
if use_openai_sdk:
normalized_api_base = api_base.rstrip("/")
if not normalized_api_base.endswith("/v1"):
normalized_api_base = f"{normalized_api_base}/v1"
owns_client: Final = not isinstance(client, AsyncOpenAI)
openai_client: Final = (
client
if isinstance(client, AsyncOpenAI)
else AsyncOpenAI(api_key=api_key, base_url=normalized_api_base, max_retries=0)
)
logging_obj.pre_call(
input=request_data,
api_key="",
additional_args={ # mutable-ok: logging owns a mutable request metadata payload
"complete_input_dict": request_data,
"api_base": normalized_api_base,
},
)
try:
configured_client: Final = openai_client.with_options(
timeout=timeout,
set_default_headers={ # mutable-ok: OpenAI SDK accepts a mutable custom-header mapping
key: str(value) for key, value in (extra_headers or {}).items()
},
)
return await configured_client.post(
"/realtime/translations/client_secrets",
cast_to=httpx.Response,
body=request_data,
)
finally:
if owns_client:
await openai_client.close()
return await self._async_realtime_session_post(
endpoint="translation_client_secrets",
api_base=api_base,
api_key=api_key,
request_data=request_data,
logging_obj=logging_obj,
timeout=timeout,
provider_config=provider_config,
model=model,
extra_headers=extra_headers,
client=client,
api_version=api_version,
)
async def _async_realtime_session_post(
self,
endpoint: Literal["client_secrets", "transcription_sessions", "translation_client_secrets"],
api_base: str,
api_key: str,
request_data: dict[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
extra_headers: Mapping[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | AsyncOpenAI | None = None,
api_version: str | None = None,
) -> httpx.Response:
"""
@ -6508,13 +6619,20 @@ class BaseLLMHTTPHandler:
url = provider_config.get_transcription_session_url(
api_base=api_base, model=model or "", api_version=api_version
)
elif endpoint == "translation_client_secrets":
url = provider_config.get_translation_client_secret_url(
api_base=api_base, model=model or "", api_version=api_version
)
else:
url = provider_config.get_complete_url(api_base=api_base, model=model or "", api_version=api_version)
headers: dict[str, object] = provider_config.validate_environment(
headers={}, model=model or "", api_key=api_key
)
else:
url = f"{api_base.rstrip('/')}/v1/realtime/{endpoint}"
endpoint_path: Final = (
"translations/client_secrets" if endpoint == "translation_client_secrets" else endpoint
)
url = f"{api_base.rstrip('/')}/v1/realtime/{endpoint_path}"
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
@ -6548,6 +6666,80 @@ class BaseLLMHTTPHandler:
)
raise
async def _async_realtime_calls_sdk(
self,
api_base: str,
openai_ephemeral_key: str,
sdp_text: str,
session_data: Mapping[str, object],
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
extra_headers: Mapping[str, object] | None,
client: object | None,
translation: bool,
) -> httpx.Response:
normalized_api_base = api_base.rstrip("/")
if not normalized_api_base.endswith("/v1"):
normalized_api_base = f"{normalized_api_base}/v1"
owns_client: Final = not isinstance(client, AsyncOpenAI)
openai_client: Final = (
client
if isinstance(client, AsyncOpenAI)
else AsyncOpenAI(api_key=openai_ephemeral_key, base_url=normalized_api_base, max_retries=0)
)
logging_obj.pre_call(
input="realtime_sdp_offer",
api_key="",
additional_args={ # mutable-ok: logging owns a mutable request metadata payload
"api_base": normalized_api_base,
"session": session_data,
},
)
try:
if translation:
configured_client: Final = openai_client.with_options(
timeout=timeout,
set_default_headers={ # mutable-ok: OpenAI SDK accepts a mutable custom-header mapping
"Content-Type": "application/sdp",
**{ # mutable-ok: caller headers are normalized into the SDK header mapping
key: str(value) for key, value in (extra_headers or {}).items()
},
},
)
return await configured_client.post(
"/realtime/translations/calls",
cast_to=httpx.Response,
content=sdp_text.encode("utf-8"),
)
realtime_session_data: Final = cast( # cast-ok: endpoint validation produced an OpenAI realtime session
RealtimeSessionCreateRequestParam,
session_data,
)
sdk_extra_headers: Final = { # mutable-ok: OpenAI SDK accepts a mutable custom-header mapping
key: str(value) for key, value in (extra_headers or {}).items()
}
raw_response: Final = await openai_client.realtime.calls.with_raw_response.create(
sdp=sdp_text,
session=realtime_session_data,
extra_headers=sdk_extra_headers,
timeout=timeout,
)
return httpx.Response(
status_code=raw_response.status_code,
headers=raw_response.headers,
content=raw_response.content,
request=httpx.Request("POST", f"{normalized_api_base}/realtime/calls"),
)
finally:
if owns_client:
await openai_client.close()
@staticmethod
def _get_realtime_async_http_client(client: object | None) -> AsyncHTTPHandler:
if isinstance(client, AsyncHTTPHandler):
return client
return get_async_httpx_client(llm_provider=litellm.LlmProviders.OPENAI)
async def async_realtime_calls_handler(
self,
api_base: str,
@ -6555,12 +6747,14 @@ class BaseLLMHTTPHandler:
sdp_body: bytes,
logging_obj: LiteLLMLoggingObj,
timeout: float | httpx.Timeout,
provider_config: Any | None = None,
provider_config: BaseRealtimeHTTPConfig | None = None,
model: str | None = None,
session_config: dict[str, object] | None = None,
extra_headers: dict[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
session_config: Mapping[str, object] | None = None,
extra_headers: Mapping[str, object] | None = None,
client: HTTPHandler | AsyncHTTPHandler | AsyncOpenAI | None = None,
api_version: str | None = None,
translation: bool = False,
use_openai_sdk: bool = False,
) -> httpx.Response:
"""
Forward POST /v1/realtime/calls (SDP exchange) to upstream provider.
@ -6572,18 +6766,45 @@ class BaseLLMHTTPHandler:
- sdp: the SDP offer (text)
- session: JSON string with {"type": "realtime", "model": "...", ...}
"""
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.OPENAI,
session_data: Final[dict[str, object]] = { # mutable-ok: model and session type are resolved locally
**(
session_config or {} # mutable-ok: absent session configuration starts from an empty provider payload
)
else:
async_httpx_client = client
}
if "type" not in session_data:
session_data["type"] = "translation" if translation else "realtime"
if "model" not in session_data and model:
session_data["model"] = model
sdp_text: Final = sdp_body.decode("utf-8") if isinstance(sdp_body, bytes) else sdp_body
if use_openai_sdk:
return await self._async_realtime_calls_sdk(
api_base=api_base,
openai_ephemeral_key=openai_ephemeral_key,
sdp_text=sdp_text,
session_data=session_data,
logging_obj=logging_obj,
timeout=timeout,
extra_headers=extra_headers,
client=client,
translation=translation,
)
async_httpx_client: Final = self._get_realtime_async_http_client(client)
if provider_config is not None:
url = provider_config.get_realtime_calls_url(api_base=api_base, model=model or "", api_version=api_version)
url = (
provider_config.get_translation_calls_url(api_base=api_base, model=model or "", api_version=api_version)
if translation
else provider_config.get_realtime_calls_url(
api_base=api_base, model=model or "", api_version=api_version
)
)
headers: dict[str, object] = provider_config.get_realtime_calls_headers(ephemeral_key=openai_ephemeral_key)
else:
url = f"{api_base.rstrip('/')}/v1/realtime/calls"
path: Final = "translations/calls" if translation else "calls"
url = f"{api_base.rstrip('/')}/v1/realtime/{path}"
headers = {
"Authorization": f"Bearer {openai_ephemeral_key}",
}
@ -6591,14 +6812,8 @@ class BaseLLMHTTPHandler:
if extra_headers:
headers.update(extra_headers)
# Build multipart form data: sdp + session JSON
session_data: Final = session_config or {}
if "type" not in session_data:
session_data["type"] = "realtime"
if "model" not in session_data and model:
session_data["model"] = model
sdp_text: Final = sdp_body.decode("utf-8") if isinstance(sdp_body, bytes) else sdp_body
if translation:
headers["Content-Type"] = "application/sdp"
files: Final = {
"sdp": (None, sdp_text, "text/plain"),
@ -6616,12 +6831,14 @@ class BaseLLMHTTPHandler:
)
try:
return await async_httpx_client.post(
url=url,
headers=headers,
files=files,
timeout=timeout,
)
if translation:
return await async_httpx_client.post(
url=url,
headers=headers,
content=sdp_text,
timeout=timeout,
)
return await async_httpx_client.post(url=url, headers=headers, files=files, timeout=timeout)
except Exception as e:
if provider_config is not None:
raise self._handle_error(

View file

@ -5,8 +5,17 @@ This requires websockets, and is currently only supported on LiteLLM Proxy.
"""
import ssl
from collections.abc import Mapping
from contextlib import AbstractAsyncContextManager
from types import TracebackType
from typing import Any, Final, cast
from openai import AsyncOpenAI, omit
from openai.resources.realtime.realtime import (
AsyncRealtimeConnection,
AsyncRealtimeConnectionManager,
)
from litellm._logging import _redact_string, verbose_logger
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
from litellm.types.realtime import RealtimeQueryParams
@ -22,6 +31,49 @@ from ....llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
from ..openai import OpenAIChatCompletion
class OpenAIRealtimeConnectionAdapter:
def __init__(self, connection: AsyncRealtimeConnection) -> None:
self._connection = connection
async def send(self, message: str) -> None:
await self._connection.send_raw(message)
async def recv(self, decode: bool = True) -> str | bytes:
message: Final = await self._connection.recv_bytes()
if decode:
return message.decode("utf-8")
return message
async def close(self) -> None:
await self._connection.close()
class OpenAIRealtimeSDKConnectionManager:
def __init__(
self,
manager: AsyncRealtimeConnectionManager,
owned_client: AsyncOpenAI | None = None,
) -> None:
self._manager = manager
self._owned_client = owned_client
async def __aenter__(self) -> OpenAIRealtimeConnectionAdapter:
connection: Final = await self._manager.__aenter__()
return OpenAIRealtimeConnectionAdapter(connection)
async def __aexit__(
self,
exc_type: type[BaseException] | None,
exc: BaseException | None,
traceback: TracebackType | None,
) -> None:
try:
await self._manager.__aexit__(exc_type, exc, traceback)
finally:
if self._owned_client is not None:
await self._owned_client.close()
class OpenAIRealtime(OpenAIChatCompletion):
"""
Base handler for OpenAI-compatible realtime WebSocket connections.
@ -82,7 +134,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
return ssl_config
def _construct_url(self, api_base: str, query_params: RealtimeQueryParams) -> str:
def _construct_url(self, api_base: str, query_params: RealtimeQueryParams, realtime_mode: str = "realtime") -> str:
"""
Construct the backend websocket URL with all query parameters (including 'model').
"""
@ -92,7 +144,8 @@ class OpenAIRealtime(OpenAIChatCompletion):
api_base = api_base.replace("http://", "ws://")
url = URL(api_base)
# Set the correct path
url = url.copy_with(path="/v1/realtime")
path: Final = "/v1/realtime/translations" if realtime_mode == "translation" else "/v1/realtime"
url = url.copy_with(path=path)
# Include all query parameters including 'model'
if query_params:
url = url.copy_with(params=query_params)
@ -106,6 +159,47 @@ class OpenAIRealtime(OpenAIChatCompletion):
"""
return None
def _create_connection_manager(
self,
api_base: str,
api_key: str,
model: str,
query_params: RealtimeQueryParams,
headers: Mapping[str, str],
timeout: float | None,
realtime_mode: str,
ssl_config: object,
client: object | None,
url: str,
) -> AbstractAsyncContextManager[object]:
import websockets
if realtime_mode == "translation" or client is None:
return websockets.connect(
url,
additional_headers=headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_config,
)
if not isinstance(client, AsyncOpenAI):
raise TypeError("client must be an AsyncOpenAI instance")
openai_client: Final = client
model_query: Final = query_params.get("model")
extra_query: Final = { # mutable-ok: OpenAI SDK accepts a mutable query-parameter mapping
key: value for key, value in query_params.items() if key != "model"
}
sdk_model: Final = omit if query_params.get("intent") == "transcription" else model_query or model
sdk_connection_manager: Final = openai_client.realtime.connect(
model=sdk_model,
extra_query=extra_query,
extra_headers=headers,
websocket_connection_options={ # mutable-ok: OpenAI SDK forwards a mutable options mapping
"max_size": REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
},
max_retries=0,
)
return OpenAIRealtimeSDKConnectionManager(sdk_connection_manager)
async def async_realtime(
self,
model: str,
@ -118,6 +212,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
query_params: RealtimeQueryParams | None = None,
user_api_key_dict: object | None = None,
litellm_metadata: dict | None = None,
realtime_mode: str = "realtime",
**kwargs: object,
):
import websockets
@ -131,7 +226,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
# Use all query params if provided, else fallback to just model
if query_params is None:
query_params = {"model": model}
url: Final = self._construct_url(api_base, query_params)
url: Final = self._construct_url(api_base, query_params, realtime_mode=realtime_mode)
try:
# Get provider-specific SSL configuration
@ -156,15 +251,25 @@ class OpenAIRealtime(OpenAIChatCompletion):
"complete_input_dict": {"query_params": query_params},
},
)
async with websockets.connect(
url,
additional_headers=headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_config,
) as backend_ws:
connection_manager: Final = self._create_connection_manager(
api_base=api_base,
api_key=api_key,
model=model,
query_params=query_params,
headers=headers,
timeout=timeout,
realtime_mode=realtime_mode,
ssl_config=ssl_config,
client=client,
url=url,
)
async with connection_manager as backend_ws:
realtime_streaming: Final = RealTimeStreaming(
websocket,
cast(ClientConnection, backend_ws),
cast( # cast-ok: both SDK and websockets adapters implement the streaming connection interface
ClientConnection, backend_ws
),
logging_obj,
model=model,
user_api_key_dict=user_api_key_dict,
@ -173,6 +278,7 @@ class OpenAIRealtime(OpenAIChatCompletion):
model if (query_params or {}).get("intent") == "transcription" else None
),
event_normalizer=self._make_event_normalizer(),
translation_session=realtime_mode == "translation",
)
await realtime_streaming.bidirectional_forward()

View file

@ -27,6 +27,18 @@ class OpenAIRealtimeHTTPConfig(BaseRealtimeHTTPConfig):
base = base.removesuffix("/v1")
return f"{base}/v1/realtime/transcription_sessions"
def get_translation_client_secret_url(
self, api_base: str | None, model: str, api_version: str | None = None
) -> str:
base = self.get_api_base(api_base).rstrip("/")
base = base.removesuffix("/v1")
return f"{base}/v1/realtime/translations/client_secrets"
def get_translation_calls_url(self, api_base: str | None, model: str, api_version: str | None = None) -> str:
base = self.get_api_base(api_base).rstrip("/")
base = base.removesuffix("/v1")
return f"{base}/v1/realtime/translations/calls"
def validate_environment(
self,
headers: dict,

View file

@ -14,7 +14,7 @@ class OpenAIGPTAudioTranscriptionConfig(OpenAIWhisperAudioTranscriptionConfig):
"""
Get the supported OpenAI params for the `gpt-4o-transcribe` models
"""
return [
return [ # mutable-ok: base transcription interface requires a mutable supported-parameter list
"language",
"prompt",
"response_format",
@ -37,3 +37,31 @@ class OpenAIGPTAudioTranscriptionConfig(OpenAIWhisperAudioTranscriptionConfig):
return AudioTranscriptionRequestData(
data=data,
)
class OpenAIGPTTranscribeAudioTranscriptionConfig(OpenAIGPTAudioTranscriptionConfig):
def get_supported_openai_params( # mutable-ok: base transcription interface returns a mutable parameter list
self, model: str
) -> list[OpenAIAudioTranscriptionOptionalParams]:
return [
"prompt",
"response_format",
"keywords",
"languages",
"stream",
]
def transform_audio_transcription_request(
self,
model: str,
audio_file: FileTypes,
optional_params: dict, # mutable-ok: base transformation interface supplies a mutable request payload
litellm_params: dict, # mutable-ok: base transformation interface supplies mutable provider parameters
) -> AudioTranscriptionRequestData:
data: Final = { # mutable-ok: OpenAI SDK consumes this multipart request mapping
"model": model,
"file": audio_file,
"response_format": "json",
**optional_params,
}
return AudioTranscriptionRequestData(data=data)

View file

@ -25,6 +25,23 @@ from ..openai import OpenAIChatCompletion
class OpenAIAudioTranscription(OpenAIChatCompletion):
# Audio Transcriptions
@staticmethod
def _sdk_compatible_request_data(data: dict) -> dict:
"""Route API fields that predate SDK support through ``extra_body``."""
extension_keys: Final = ("keywords", "languages")
extension_body: Final = {key: data[key] for key in extension_keys if key in data}
if not extension_body:
return data
existing_extra_body: Final = data.get("extra_body")
return { # mutable-ok: OpenAI SDK requires a mutable request mapping
**{key: value for key, value in data.items() if key not in extension_keys},
"extra_body": {
**(existing_extra_body if isinstance(existing_extra_body, dict) else {}),
**extension_body,
},
}
async def make_openai_audio_transcriptions_request(
self,
openai_aclient: AsyncOpenAI,
@ -37,11 +54,17 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
- call openai_aclient.audio.transcriptions.create by default
"""
try:
raw_response = await openai_aclient.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
sdk_data: Final = self._sdk_compatible_request_data(data)
if data.get("stream") is True:
stream_response: Final = await openai_aclient.audio.transcriptions.create(**sdk_data, timeout=timeout)
return {}, stream_response # mutable-ok: response headers use the existing mutable mapping contract
raw_response: Final = await openai_aclient.audio.transcriptions.with_raw_response.create(
**sdk_data, timeout=timeout
) # pyright: ignore[reportArgumentType] # SDK TypedDict lags accepted transcription options
headers: Final = dict(raw_response.headers)
response: Final = raw_response.parse()
parsed_response: Final = raw_response.parse()
return headers, response
return headers, parsed_response
except Exception as e:
raise e
@ -57,13 +80,17 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
- call openai_aclient.audio.transcriptions.create by default
"""
try:
sdk_data: Final = self._sdk_compatible_request_data(data)
if data.get("stream") is True:
response = openai_client.audio.transcriptions.create(**sdk_data, timeout=timeout)
return None, response
if litellm.return_response_headers is True:
raw_response = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
raw_response = openai_client.audio.transcriptions.with_raw_response.create(**sdk_data, timeout=timeout)
headers: Final = dict(raw_response.headers)
response = raw_response.parse()
return headers, response
else:
response = openai_client.audio.transcriptions.create(**data, timeout=timeout)
response = openai_client.audio.transcriptions.create(**sdk_data, timeout=timeout)
return None, response
except Exception as e:
raise e
@ -139,6 +166,9 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
timeout=timeout,
)
if data.get("stream") is True:
return response
if isinstance(response, BaseModel):
stringified_response = response.model_dump()
else:
@ -200,6 +230,8 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
timeout=timeout,
)
logging_obj.model_call_details["response_headers"] = headers
if data.get("stream") is True:
return response
if isinstance(response, BaseModel):
stringified_response = response.model_dump()
else:

View file

@ -37,6 +37,9 @@ if TYPE_CHECKING:
import dotenv
import httpx
import openai
import tiktoken
from openai import AsyncStream, Stream
from openai.types.audio import TranscriptionStreamEvent
from pydantic import BaseModel
from typing_extensions import overload
@ -7799,7 +7802,10 @@ async def amoderation(
@client
async def atranscription(*args, **kwargs) -> TranscriptionResponse:
async def atranscription(
*args, # noqa: ANN002 # public SDK wrapper preserves positional call compatibility
**kwargs, # noqa: ANN003 # kwargs-ok: public SDK wrapper preserves keyword call compatibility
) -> TranscriptionResponse | AsyncStream[TranscriptionStreamEvent]:
"""
Calls openai + azure whisper endpoints.
@ -7832,6 +7838,12 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse:
else:
# Call the synchronous function using run_in_executor
response = await loop.run_in_executor(None, func_with_context)
if kwargs.get("stream") is True and isinstance(response, AsyncStream):
if file is not None:
calculated_duration = calculate_request_duration(file)
if calculated_duration is not None:
setattr(response, "_litellm_audio_duration", calculated_duration)
return response
if not isinstance(response, TranscriptionResponse):
raise ValueError(
f"Invalid response from transcription provider, expected TranscriptionResponse, but got {type(response)}"
@ -7844,9 +7856,9 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse:
if response is not None and not isinstance(response, Coroutine) and file is not None:
existing_duration: Final = getattr(response, "duration", None)
if existing_duration is None:
calculated_duration: Final = calculate_request_duration(file)
if calculated_duration is not None:
response._hidden_params["audio_transcription_duration"] = calculated_duration
sync_calculated_duration: Final = calculate_request_duration(file)
if sync_calculated_duration is not None:
response.set_audio_transcription_duration(sync_calculated_duration)
return response
except Exception as e:
@ -7860,16 +7872,52 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse:
)
def _validate_gpt_transcription_request(
model: str,
custom_llm_provider: str,
language: str | None,
languages: Sequence[str] | None,
response_format: str | None,
api_version: str | None,
) -> str | None:
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":
raise litellm.UnsupportedParamsError(
message="gpt-live-transcribe 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"):
raise litellm.UnsupportedParamsError(
message="gpt-transcribe only supports response_format='json'",
model=model,
llm_provider=custom_llm_provider,
)
if custom_llm_provider == "azure" and model == "gpt-transcribe":
if api_version in (None, "v1", "latest", "preview"):
return litellm.AZURE_DEFAULT_API_VERSION
return api_version
return api_version
@client
def transcription(
model: str,
file: FileTypes,
## OPTIONAL OPENAI PARAMS ##
language: str | None = None,
languages: Sequence[str] | None = None,
keywords: Sequence[str] | None = None,
prompt: str | None = None,
response_format: Literal["json", "text", "srt", "verbose_json", "vtt"] | None = None,
timestamp_granularities: list[Literal["word", "segment"]] | None = None,
temperature: int | None = None, # openai defaults this to 0
stream: bool | None = None,
## LITELLM PARAMS ##
user: str | None = None,
timeout=600, # default to 10 minutes
@ -7879,7 +7927,11 @@ def transcription(
max_retries: int | None = None,
custom_llm_provider=None,
**kwargs,
) -> TranscriptionResponse | Coroutine[object, object, TranscriptionResponse]:
) -> (
TranscriptionResponse
| Stream[TranscriptionStreamEvent]
| Coroutine[Any, Any, TranscriptionResponse | AsyncStream[TranscriptionStreamEvent]]
):
"""
Calls openai + azure whisper endpoints.
@ -7917,13 +7969,25 @@ def transcription(
api_key = dynamic_api_key if dynamic_api_key is not None else api_key
validated_api_version: Final = _validate_gpt_transcription_request(
model=model,
custom_llm_provider=custom_llm_provider,
language=language,
languages=languages,
response_format=response_format,
api_version=api_version,
)
optional_params: Final = get_optional_params_transcription(
model=model,
language=language,
languages=languages,
keywords=keywords,
prompt=prompt,
response_format=response_format,
timestamp_granularities=timestamp_granularities,
temperature=temperature,
stream=stream,
custom_llm_provider=custom_llm_provider,
**non_default_params,
)
@ -7946,7 +8010,13 @@ def transcription(
custom_llm_provider=custom_llm_provider,
)
response: TranscriptionResponse | Coroutine[object, object, TranscriptionResponse] | None = None
response: (
TranscriptionResponse
| Stream[TranscriptionStreamEvent]
| AsyncStream[TranscriptionStreamEvent]
| Coroutine[Any, Any, TranscriptionResponse | AsyncStream[TranscriptionStreamEvent]]
| None
) = None
provider_config: Final = ProviderConfigManager.get_provider_audio_transcription_config(
model=model,
@ -7961,7 +8031,7 @@ def transcription(
# azure configs
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
api_version = api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
azure_api_version: Final = validated_api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
azure_ad_token: Final = kwargs.pop("azure_ad_token", None) or get_secret_str("AZURE_AD_TOKEN")
@ -7980,7 +8050,7 @@ def transcription(
logging_obj=litellm_logging_obj,
api_base=api_base,
api_key=api_key,
api_version=api_version,
api_version=azure_api_version,
azure_ad_token=azure_ad_token,
max_retries=max_retries,
litellm_params=litellm_params_dict,
@ -8114,11 +8184,12 @@ def transcription(
# Store duration in _hidden_params for cost calculation without
# exposing it in the response body (see sync path comment above).
if response is not None and not isinstance(response, Coroutine):
existing_duration: Final = getattr(response, "duration", None)
if existing_duration is None:
calculated_duration: Final = calculate_request_duration(file)
calculated_duration: Final = calculate_request_duration(file)
if isinstance(response, (Stream, AsyncStream)):
if calculated_duration is not None:
response._hidden_params["audio_transcription_duration"] = calculated_duration
setattr(response, "_litellm_audio_duration", calculated_duration)
elif getattr(response, "duration", None) is None and calculated_duration is not None:
response.set_audio_transcription_duration(calculated_duration)
if response is None:
raise ValueError("Unmapped provider passed in. Unable to get the response.")

View file

@ -31,6 +31,20 @@ def _include_router(attr_name: str = "router") -> Callable[["FastAPI", object],
return _register
def _include_router_prepend(attr_name: str = "router") -> Callable[["FastAPI", object], None]:
def _register(app: "FastAPI", module: object) -> None:
router: Final = app.router
route_count: Final = len(router.routes)
app.include_router(getattr(module, attr_name))
new_routes: Final = router.routes[route_count:]
router.routes[:] = [ # mutable-ok: FastAPI route registration requires in-place list replacement
*new_routes,
*router.routes[:route_count],
]
return _register
def _mount_app(prefix: str, attr_name: str = "app") -> Callable[["FastAPI", object], None]:
def _register(app: "FastAPI", module: object) -> None:
app.mount(path=prefix, app=getattr(module, attr_name))
@ -225,6 +239,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
name="realtime",
module_path="litellm.proxy.realtime_endpoints.endpoints",
path_prefixes=("/openai/v1/realtime", "/v1/realtime", "/realtime"),
register_fn=_include_router_prepend(),
),
LazyFeature(
name="anthropic_passthrough",

View file

@ -20245,6 +20245,51 @@
]
}
},
"/openai/v1/realtime/translations/calls": {
"post": {
"operationId": "proxy_realtime_calls_openai_v1_realtime_translations_calls_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Realtime Calls",
"tags": [
"llm_passthrough"
]
}
},
"/openai/v1/realtime/translations/client_secrets": {
"post": {
"operationId": "create_realtime_client_secret_openai_v1_realtime_translations_client_secrets_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/RealtimeClientSecretResponse"
}
}
},
"description": "Successful Response"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Create Realtime Client Secret",
"tags": [
"llm_passthrough"
]
}
},
"/openai/v1/responses": {
"post": {
"description": "Follows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses\n\nSupports background mode with polling_via_cache for partial response retrieval.\nWhen background=true and polling_via_cache is enabled, returns a polling_id immediately\nand streams the response in the background, updating Redis cache.\n\n```bash\n# Normal request\ncurl -X POST http://localhost:4000/v1/responses -H \"Content-Type: application/json\" -H \"Authorization: Bearer sk-1234\" -d '{\n \"model\": \"gpt-4o\",\n \"input\": \"Tell me about AI\"\n}'\n\n# Background request with polling\ncurl -X POST http://localhost:4000/v1/responses -H \"Content-Type: application/json\" -H \"Authorization: Bearer sk-1234\" -d '{\n \"model\": \"gpt-4o\",\n \"input\": \"Tell me about AI\",\n \"background\": true\n}'\n```",
@ -39286,6 +39331,51 @@
]
}
},
"/openai/v1/realtime/translations/calls": {
"post": {
"operationId": "proxy_realtime_calls_openai_v1_realtime_translations_calls_post_2",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Realtime Calls",
"tags": [
"realtime"
]
}
},
"/openai/v1/realtime/translations/client_secrets": {
"post": {
"operationId": "create_realtime_client_secret_openai_v1_realtime_translations_client_secrets_post_2",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/RealtimeClientSecretResponse"
}
}
},
"description": "Successful Response"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Create Realtime Client Secret",
"tags": [
"realtime"
]
}
},
"/realtime/calls": {
"post": {
"operationId": "proxy_realtime_calls_realtime_calls_post",
@ -39358,6 +39448,51 @@
]
}
},
"/realtime/translations/calls": {
"post": {
"operationId": "proxy_realtime_calls_realtime_translations_calls_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Realtime Calls",
"tags": [
"realtime"
]
}
},
"/realtime/translations/client_secrets": {
"post": {
"operationId": "create_realtime_client_secret_realtime_translations_client_secrets_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/RealtimeClientSecretResponse"
}
}
},
"description": "Successful Response"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Create Realtime Client Secret",
"tags": [
"realtime"
]
}
},
"/v1/realtime/calls": {
"post": {
"operationId": "proxy_realtime_calls_v1_realtime_calls_post",
@ -39429,6 +39564,51 @@
"realtime"
]
}
},
"/v1/realtime/translations/calls": {
"post": {
"operationId": "proxy_realtime_calls_v1_realtime_translations_calls_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"summary": "Proxy Realtime Calls",
"tags": [
"realtime"
]
}
},
"/v1/realtime/translations/client_secrets": {
"post": {
"operationId": "create_realtime_client_secret_v1_realtime_translations_client_secrets_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {
"$ref": "#/components/schemas/RealtimeClientSecretResponse"
}
}
},
"description": "Successful Response"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Create Realtime Client Secret",
"tags": [
"realtime"
]
}
}
}
},

View file

@ -417,6 +417,9 @@ class LiteLLMRoutes(enum.Enum):
"/realtime?{model}",
"/v1/realtime?{model}",
"/openai/v1/realtime?{model}",
"/realtime/translations",
"/v1/realtime/translations",
"/openai/v1/realtime/translations",
# realtime (GA WebRTC HTTP routes)
"/realtime/client_secrets",
"/v1/realtime/client_secrets",
@ -427,6 +430,12 @@ class LiteLLMRoutes(enum.Enum):
"/realtime/transcription_sessions",
"/v1/realtime/transcription_sessions",
"/openai/v1/realtime/transcription_sessions",
"/realtime/translations/client_secrets",
"/v1/realtime/translations/client_secrets",
"/openai/v1/realtime/translations/client_secrets",
"/realtime/translations/calls",
"/v1/realtime/translations/calls",
"/openai/v1/realtime/translations/calls",
# responses API
"/responses",
"/v1/responses",
@ -2625,6 +2634,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
admission_queue_timeout_seconds: float = Field(
1.0, gt=0, description="maximum time a request waits for a worker slot"
)
allow_non_billable_realtime_protocols: bool = Field(
False,
description="Allow Realtime WebRTC setup endpoints whose inference usage bypasses LiteLLM spend tracking and budget enforcement",
)
plugins: list[PluginConfig] | None = Field(
None, description="external services registered as embeddable UI plugins"
)

View file

@ -132,7 +132,9 @@ ProxyRouteType: TypeAlias = Literal[
"_arealtime",
"_aresponses_websocket",
"acreate_realtime_client_secret",
"acreate_realtime_translation_client_secret",
"arealtime_calls",
"arealtime_translation_calls",
"aget_responses",
"adelete_responses",
"acancel_responses",

View file

@ -399,14 +399,18 @@ async def get_form_data(request: Request) -> dict[str, Any]:
Handles when OpenAI SDKs pass form keys as `timestamp_granularities[]="word"` instead of `timestamp_granularities=["word", "sentence"]`
"""
form: Final = await request.form()
parsed_form_data: Final[dict[str, Any]] = {}
for key, value in form.multi_items(): # not dict(form), which keeps only the last repeat
if key.endswith("[]"):
clean_key = key[:-2]
parsed_form_data.setdefault(clean_key, []).append(value)
else:
parsed_form_data[key] = value
return parsed_form_data
form_items: Final = tuple(form.multi_items() if hasattr(form, "multi_items") else form.items())
array_keys: Final = frozenset(key[:-2] for key, _ in form_items if key.endswith("[]"))
normalized_items: Final = tuple((key.removesuffix("[]"), value) for key, value in form_items)
normalized_keys: Final = frozenset(key for key, _ in normalized_items)
return { # mutable-ok: request parsers expose a mutable form-data mapping to endpoint handlers
key: [ # mutable-ok: repeated multipart values follow the established mutable list contract
value for item_key, value in normalized_items if item_key == key
]
if key in array_keys
else next(value for item_key, value in reversed(normalized_items) if item_key == key)
for key in normalized_keys
}
async def convert_upload_files_to_file_data(

View file

@ -49,6 +49,7 @@ from typing import (
import anyio
import websockets
import websockets.exceptions
from openai.types.audio import TranscriptionStreamEvent
from pydantic import BaseModel, Json, JsonValue, TypeAdapter, ValidationError
from pydantic.fields import FieldInfo, PydanticUndefined
from typing_extensions import NotRequired, ReadOnly, assert_never
@ -12199,7 +12200,11 @@ async def audio_transcriptions(
try:
# Use orjson to parse JSON data, orjson speeds up requests significantly
form_data: Final = await get_form_data(request)
data = {key: value for key, value in form_data.items() if key != "file"} | data
data = {
key: value is True or str(value).lower() in ("1", "true") if key == "stream" else value
for key, value in form_data.items()
if key != "file"
} | data
# Include original request and headers in the data
data = await add_litellm_data_to_request(
@ -12265,6 +12270,29 @@ async def audio_transcriptions(
finally:
file_object.close() # close the file read in by io library
if data.get("stream") is True:
if not hasattr(response, "__aiter__"):
raise TypeError(f"Streaming transcription returned {type(response).__name__}, expected an async stream")
stream_response: Final = cast(AsyncIterator[TranscriptionStreamEvent], response)
async def transcription_event_stream(
stream: AsyncIterator[TranscriptionStreamEvent],
) -> AsyncGenerator[str, None]:
try:
async for event in stream:
yield f"data: {event.model_dump_json()}\n\n"
finally:
close: Final = getattr(stream, "aclose", None) or getattr(stream, "close", None)
if callable(close):
close_result: Final = close()
if inspect.isawaitable(close_result):
await close_result
return StreamingResponse(
transcription_event_stream(stream_response),
media_type="text/event-stream",
)
### ALERTING ###
asyncio.create_task(proxy_logging_obj.update_request_status(litellm_call_id=litellm_call_id, status="success"))
@ -12431,9 +12459,39 @@ async def _reject_realtime_session(
await _release_realtime_max_parallel_slot(user_api_key_dict)
def _resolve_realtime_route_model(
model: str | None,
intent: str | None,
is_translation: bool,
) -> str | None:
if model is not None:
return model
if is_translation:
return "gpt-realtime-translate"
if intent == "transcription":
return "gpt-realtime-whisper"
return None
def _resolve_realtime_upstream_query_model(
model: str | None,
intent: str | None,
is_translation: bool,
route_model: str,
) -> str | None:
if intent == "transcription":
return None
if is_translation:
return route_model
return model
@app.websocket("/openai/v1/realtime")
@app.websocket("/v1/realtime")
@app.websocket("/realtime")
@app.websocket("/openai/v1/realtime/translations")
@app.websocket("/v1/realtime/translations")
@app.websocket("/realtime/translations")
async def realtime_websocket_endpoint(
websocket: WebSocket,
model: str | None = fastapi.Query(None, description="The model to use for the websocket connection."),
@ -12451,15 +12509,13 @@ async def realtime_websocket_endpoint(
if requested_protocols:
accept_kwargs["subprotocol"] = requested_protocols[0]
route_model = model
is_translation: Final = websocket.url.path.endswith("/realtime/translations")
route_model: Final = _resolve_realtime_route_model(model, intent, is_translation)
if route_model is None:
if intent == "transcription":
route_model = "gpt-realtime-whisper"
else:
await _reject_realtime_session(
websocket, user_api_key_dict, code=1008, reason="model query parameter is required"
)
return
await _reject_realtime_session(
websocket, user_api_key_dict, code=1008, reason="model query parameter is required"
)
return
assert route_model is not None
try:
await can_key_call_resolved_model(
@ -12475,12 +12531,24 @@ async def realtime_websocket_endpoint(
await websocket.accept(**accept_kwargs)
# Only use explicit parameters, not all query params
query_params: Final = cast(RealtimeQueryParams, dict(_realtime_query_params_template(model, intent)))
query_model: Final = _resolve_realtime_upstream_query_model(
model=model,
intent=intent,
is_translation=is_translation,
route_model=route_model,
)
query_params: Final = cast( # cast-ok: cached tuples contain only the declared realtime query keys
RealtimeQueryParams,
dict( # mutable-ok: downstream realtime routing normalizes this request-scoped query mapping
_realtime_query_params_template(query_model, intent)
),
)
data: dict[str, object] = {
"model": route_model,
"websocket": websocket,
"query_params": query_params, # Only explicit params
"realtime_mode": "translation" if is_translation else "realtime",
}
# Pass guardrails into data so pre-call guardrail processing picks them up

View file

@ -2,6 +2,7 @@
import json
import time
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -37,7 +38,21 @@ router: Final = APIRouter()
_REALTIME_TOKEN_VERSION: Final = "realtime_v1"
_DEFAULT_REALTIME_MODEL: Final = "gpt-4o-realtime-preview"
_DEFAULT_TRANSCRIPTION_MODEL: Final = "gpt-realtime-whisper"
_ALLOWED_SESSION_TYPES: Final = ("realtime", "transcription")
_ALLOWED_SESSION_TYPES: Final = ("realtime", "transcription", "translation")
_NON_BILLABLE_REALTIME_PROTOCOL_SETTING: Final = "allow_non_billable_realtime_protocols"
_NON_BILLABLE_REALTIME_PROTOCOL_MESSAGE: Final = (
"Realtime WebRTC endpoints are disabled because provider usage bypasses LiteLLM billing. "
"Set general_settings.allow_non_billable_realtime_protocols to true to opt in"
)
def _enforce_non_billable_realtime_protocol_gate(general_settings: Mapping[str, object]) -> None:
if general_settings.get(_NON_BILLABLE_REALTIME_PROTOCOL_SETTING) is True:
return
raise HTTPException(
status_code=http_status.HTTP_403_FORBIDDEN,
detail=_NON_BILLABLE_REALTIME_PROTOCOL_MESSAGE,
)
def _coerce_realtime_session_type(session_type: str | None) -> str:
@ -120,19 +135,47 @@ def _set_transcription_model_on_session(
}
async def _authorize_and_bind_nested_transcription_models(
session_data: dict, # mutable-ok: session payload is rewritten in place for provider serialization
user_api_key_dict: UserAPIKeyAuth,
llm_model_list: list | None, # mutable-ok: inherited auth helper accepts the proxy model list
llm_router: Any,
) -> None:
nested_models: Final = tuple(_transcription_model_candidates_from_session(session_data))
for nested_model in nested_models:
await can_key_call_resolved_model(
model=nested_model,
valid_token=user_api_key_dict,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
if nested_models:
_set_transcription_model_on_session(
session=session_data,
model=nested_models[0],
)
async def _prepare_client_secret_session(
req: RealtimeClientSecretRequest,
user_api_key_dict: UserAPIKeyAuth,
llm_model_list: list | None,
llm_router: "Router | None",
forced_session_type: str | None = None,
) -> tuple[str, dict | None, str]:
session_type: Final = _coerce_realtime_session_type(req.session.type if req.session else None)
session_data: Final[dict | None] = req.session.model_dump(exclude_none=True) if req.session else None
requested_session_type: Final = req.session.type if req.session else None
if forced_session_type is None and requested_session_type == "translation":
raise HTTPException(status_code=400, detail="Translation sessions require the translations endpoint")
session_type: Final = forced_session_type or _coerce_realtime_session_type(requested_session_type)
session_data: Final[dict | None] = (
req.session.model_dump(exclude_none=True) if req.session else ({} if session_type == "translation" else None)
)
if session_data is not None:
session_data["type"] = session_type
session_model: Final = req.session.model if req.session else None
model: str = session_model or req.model or _DEFAULT_REALTIME_MODEL
default_model: Final = "gpt-realtime-translate" if session_type == "translation" else _DEFAULT_REALTIME_MODEL
model: str = session_model or req.model or default_model
if session_type != "transcription":
await can_key_call_resolved_model(
model=model,
@ -140,6 +183,15 @@ async def _prepare_client_secret_session(
llm_model_list=llm_model_list,
llm_router=llm_router,
)
if session_data is not None:
session_data["model"] = model
if session_type == "translation":
await _authorize_and_bind_nested_transcription_models(
session_data=session_data,
user_api_key_dict=user_api_key_dict,
llm_model_list=llm_model_list,
llm_router=llm_router,
)
return model, session_data, session_type
transcription_model_candidates: Final = _transcription_model_candidates_from_session(session_data or {})
@ -228,6 +280,21 @@ def _decode_realtime_token_payload(
dependencies=[Depends(user_api_key_auth)],
tags=["realtime"],
)
@router.post(
"/v1/realtime/translations/client_secrets",
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI decorator requires dependency lists
tags=["realtime"], # mutable-ok: FastAPI decorator requires tag lists
)
@router.post(
"/realtime/translations/client_secrets",
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI decorator requires dependency lists
tags=["realtime"], # mutable-ok: FastAPI decorator requires tag lists
)
@router.post(
"/openai/v1/realtime/translations/client_secrets",
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI decorator requires dependency lists
tags=["realtime"], # mutable-ok: FastAPI decorator requires tag lists
)
async def create_realtime_client_secret(
request: Request,
fastapi_response: Response,
@ -245,16 +312,19 @@ async def create_realtime_client_secret(
version,
)
_enforce_non_billable_realtime_protocol_gate(general_settings)
data: dict = {}
try:
body: Final = await _read_request_body(request=request)
req: Final = RealtimeClientSecretRequest(**body)
is_translation_request: Final = "/realtime/translations/client_secrets" in request.url.path
model, session_data, session_type = await _prepare_client_secret_session(
req=req,
user_api_key_dict=user_api_key_dict,
llm_model_list=llm_model_list,
llm_router=llm_router,
forced_session_type="translation" if is_translation_request else None,
)
data = {"model": model}
@ -278,17 +348,20 @@ async def create_realtime_client_secret(
proxy_config=proxy_config,
)
call_type: Final = (
"acreate_realtime_translation_client_secret" if is_translation_request else "acreate_realtime_client_secret"
)
data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=user_api_key_dict,
data=data,
call_type="acreate_realtime_client_secret",
call_type=call_type,
)
verbose_proxy_logger.debug("WebRTC: /v1/realtime/client_secrets (model=%s)", model)
llm_call: Final = await route_request(
data=data,
route_type="acreate_realtime_client_secret",
route_type=call_type,
llm_router=llm_router,
user_model=user_model,
)
@ -371,6 +444,18 @@ async def create_realtime_client_secret(
"/openai/v1/realtime/calls",
tags=["realtime"],
)
@router.post(
"/v1/realtime/translations/calls",
tags=["realtime"], # mutable-ok: FastAPI decorator requires tag lists
)
@router.post(
"/realtime/translations/calls",
tags=["realtime"], # mutable-ok: FastAPI decorator requires tag lists
)
@router.post(
"/openai/v1/realtime/translations/calls",
tags=["realtime"], # mutable-ok: FastAPI decorator requires tag lists
)
async def proxy_realtime_calls(
request: Request,
fastapi_response: Response,
@ -396,6 +481,7 @@ async def proxy_realtime_calls(
media_type="application/json",
)
is_translation_request: Final = "/realtime/translations/calls" in request.url.path
encrypted_token: Final = auth_header.removeprefix("Bearer ").strip()
decrypted_token_value: Final = decrypt_value_helper(
value=encrypted_token,
@ -408,26 +494,42 @@ async def proxy_realtime_calls(
media_type="application/json",
)
_enforce_non_billable_realtime_protocol_gate(general_settings)
sdp_body: Final[bytes] = await request.body()
decoded_payload: Final = _decode_realtime_token_payload(decrypted_token_value)
if decoded_payload is not None:
# Check token expiry
expires_at: Final = decoded_payload.get("expires_at")
if expires_at is not None and isinstance(expires_at, int):
if time.time() > expires_at:
return Response(
content=json.dumps({"error": "Token has expired"}),
status_code=http_status.HTTP_401_UNAUTHORIZED,
media_type="application/json",
)
if isinstance(expires_at, int) and time.time() > expires_at:
return Response(
content=json.dumps({"error": "Token has expired"}),
status_code=http_status.HTTP_401_UNAUTHORIZED,
media_type="application/json",
)
openai_ephemeral_key = decoded_payload.get("ephemeral_key", "")
model = decoded_payload.get("model_id") or request.query_params.get("model") or _DEFAULT_REALTIME_MODEL
user_id = decoded_payload.get("user_id") or None
team_id = decoded_payload.get("team_id") or None
session_type = _coerce_realtime_session_type(decoded_payload.get("session_type"))
raw_session_type: Final = decoded_payload.get("session_type")
session_type = _coerce_realtime_session_type(raw_session_type)
if is_translation_request != (raw_session_type == "translation"):
return Response(
content=json.dumps( # mutable-ok: JSON encoder requires the endpoint error payload mapping
{"error": "Token is not valid for this Realtime endpoint"}
),
status_code=http_status.HTTP_401_UNAUTHORIZED,
media_type="application/json",
)
else:
# Backward compatibility: older tokens contained only encrypted upstream key.
if is_translation_request:
return Response(
content=json.dumps( # mutable-ok: JSON encoder requires the endpoint error payload mapping
{"error": "Token is not valid for this Realtime endpoint"}
),
status_code=http_status.HTTP_401_UNAUTHORIZED,
media_type="application/json",
)
openai_ephemeral_key = decrypted_token_value
model = request.query_params.get("model", _DEFAULT_REALTIME_MODEL)
user_id = None
@ -471,17 +573,18 @@ async def proxy_realtime_calls(
proxy_config=proxy_config,
)
call_type: Final = "arealtime_translation_calls" if is_translation_request else "arealtime_calls"
data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=minimal_auth,
data=data,
call_type="arealtime_calls",
call_type=call_type,
)
verbose_proxy_logger.debug("WebRTC: /v1/realtime/calls (model=%s)", model)
llm_call: Final = await route_request(
data=data,
route_type="arealtime_calls",
route_type=call_type,
llm_router=llm_router,
user_model=user_model,
)
@ -557,6 +660,7 @@ async def create_realtime_transcription_session(
version,
)
_enforce_non_billable_realtime_protocol_gate(general_settings)
data: dict = {}
try:
body: Final = await _read_request_body(request=request)

View file

@ -342,7 +342,9 @@ RouteType = Literal[
"alist_input_items",
"_arealtime", # private function for realtime API
"acreate_realtime_client_secret",
"acreate_realtime_translation_client_secret",
"arealtime_calls",
"arealtime_translation_calls",
"acreate_realtime_transcription_session",
"_aresponses_websocket", # private function for responses WebSocket mode
"aimage_edit",

View file

@ -6,4 +6,16 @@ Supported endpoints:
Supported providers: OpenAI, Azure OpenAI, Bedrock, Vertex AI, xAI.
For user-facing documentation and usage examples, see the litellm-docs repo.
Billing visibility:
- WebSocket sessions pass provider usage events through LiteLLM and support local spend tracking
- Client-secret and SDP call endpoints only proxy session setup; subsequent WebRTC media and usage events travel over the peer connection, so LiteLLM cannot record inference spend or enforce spend-based budgets for those sessions
- Use the proxied WebSocket transport when LiteLLM spend logs and budgets must include Realtime inference
Non-billable Realtime protocols are disabled by default. Operators who accept the billing and budget-enforcement limitation can opt in:
```yaml
general_settings:
allow_non_billable_realtime_protocols: true
```
For user-facing documentation and usage examples, see the litellm-docs repo.

View file

@ -6,14 +6,18 @@ from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, cast
import httpx
import litellm
from litellm.constants import (
AZURE_GA_REALTIME_MODELS,
AZURE_OPENAI_AUDIO_PROVIDERS,
REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
request_timeout,
)
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.llms.base_llm.realtime.http_transformation import BaseRealtimeHTTPConfig
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.xai.common_utils import XAIModelInfo
@ -70,6 +74,30 @@ def _model_params_with_stored_credentials(model_params: Mapping[str, object]) ->
def _with_resolved_session_model(session: dict[str, object], model_name: str) -> dict[str, object]:
if session.get("type") == "transcription":
audio = session.get("audio")
audio = audio if isinstance(audio, dict) else {} # mutable-ok: nested session model is rebuilt locally
audio_input = audio.get("input")
audio_input = ( # mutable-ok: nested session model is rebuilt locally
audio_input if isinstance(audio_input, dict) else {}
)
transcription = audio_input.get("transcription")
transcription = ( # mutable-ok: nested session model is rebuilt locally
transcription if isinstance(transcription, dict) else {}
)
return { # mutable-ok: provider routing requires an independently mutable session payload
**session,
"audio": { # mutable-ok: provider routing rebuilds nested audio configuration
**audio,
"input": { # mutable-ok: provider routing rebuilds nested input configuration
**audio_input,
"transcription": { # mutable-ok: resolved deployment replaces only the transcription model
**transcription,
"model": model_name,
},
},
},
}
if "model" not in session:
return session
return {**session, "model": model_name}
@ -84,6 +112,21 @@ def _build_litellm_metadata(kwargs: dict) -> dict:
return metadata
def _resolve_azure_realtime_protocol(
model: str,
realtime_protocol: str | None,
query_params: RealtimeQueryParams | None,
realtime_mode: str,
) -> str:
if model in AZURE_GA_REALTIME_MODELS:
if realtime_protocol is not None and realtime_protocol.upper() not in ("GA", "V1"):
raise ValueError(f"{model} requires the Azure OpenAI v1 Realtime API")
return "GA"
if realtime_mode == "translation" or (query_params or {}).get("intent") == "transcription":
return "GA"
return realtime_protocol or "beta"
def _get_realtime_http_provider_config(
custom_llm_provider: str,
dynamic_api_base: str | None,
@ -97,10 +140,6 @@ def _get_realtime_http_provider_config(
Uses ProviderConfigManager so each provider keeps its credential-resolution
and URL-construction logic in its own transformation class.
"""
from litellm.llms.base_llm.realtime.http_transformation import (
BaseRealtimeHTTPConfig,
)
provider_config: BaseRealtimeHTTPConfig | None = None
if custom_llm_provider in LlmProviders._member_map_.values():
provider_config = ProviderConfigManager.get_provider_realtime_http_config(
@ -124,11 +163,27 @@ def _get_realtime_http_provider_config(
return provider_config, resolved_api_base.rstrip("/"), resolved_api_key
def _get_realtime_http_extra_headers(
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
resolved_api_key: str,
extra_headers: Mapping[str, object] | None,
) -> Mapping[str, object] | None:
resolved_headers: Final = { # mutable-ok: Azure authentication may extend caller-supplied headers
**(extra_headers or {})
}
if custom_llm_provider == "azure" and not resolved_api_key:
azure_ad_token: Final = get_azure_ad_token(litellm_params)
if azure_ad_token:
resolved_headers["Authorization"] = f"Bearer {azure_ad_token}"
return resolved_headers or None
@wrapper_client
async def acreate_realtime_client_secret(
model: str | None = None,
session: dict[str, Any] | None = None,
expires_after: dict[str, Any] | None = None,
session: Mapping[str, Any] | None = None,
expires_after: Mapping[str, Any] | None = None,
timeout: float | None = None,
**kwargs,
):
@ -137,30 +192,48 @@ async def acreate_realtime_client_secret(
session=RealtimeSessionConfig(**session) if session else None,
expires_after=RealtimeExpiresAfter(**expires_after) if expires_after else None,
)
model_name = (req.session.model if req.session is not None else None) or req.model or "gpt-4o-realtime-preview"
transcription_model: Final = (
req.session.audio.input.transcription.model
if req.session is not None
and req.session.audio is not None
and req.session.audio.input is not None
and req.session.audio.input.transcription is not None
else None
)
provider_qualified_model: Final = (
req.model
if req.model is not None
and "/" in req.model
and req.model.split("/", 1)[0] in LlmProviders._member_map_.values()
else None
)
requested_model_name: Final = (
provider_qualified_model
or transcription_model
or (req.session.model if req.session is not None else None)
or req.model
or "gpt-4o-realtime-preview"
)
litellm_logging_obj: Final[LiteLLMLogging] = kwargs.get("litellm_logging_obj")
litellm_params: Final = GenericLiteLLMParams(**kwargs)
(
model_name,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = get_llm_provider(
model=model_name,
model_name, custom_llm_provider, dynamic_api_key, dynamic_api_base = get_llm_provider(
model=requested_model_name,
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
)
(
provider_config,
resolved_api_base,
resolved_api_key,
) = _get_realtime_http_provider_config(
provider_config, resolved_api_base, resolved_api_key = _get_realtime_http_provider_config(
custom_llm_provider=custom_llm_provider,
dynamic_api_base=dynamic_api_base,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
resolved_extra_headers: Final = _get_realtime_http_extra_headers(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
resolved_api_key=resolved_api_key,
extra_headers=kwargs.get("extra_headers"),
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
@ -171,6 +244,11 @@ async def acreate_realtime_client_secret(
request_data: Final = req.model_dump(exclude_none=True, exclude={"model"})
if isinstance(request_data.get("session"), dict):
request_data["session"] = _with_resolved_session_model(request_data["session"], model_name)
elif req.model is not None:
request_data["session"] = { # mutable-ok: OpenAI SDK consumes this request-scoped session payload
"type": "realtime",
"model": model_name,
}
return await base_llm_http_handler.async_realtime_client_secret_handler(
api_base=resolved_api_base,
api_key=resolved_api_key,
@ -179,9 +257,86 @@ async def acreate_realtime_client_secret(
timeout=timeout or request_timeout,
provider_config=provider_config,
model=model_name,
extra_headers=kwargs.get("extra_headers"),
extra_headers=resolved_extra_headers,
client=kwargs.get("client"),
api_version=litellm_params.api_version,
use_openai_sdk=custom_llm_provider == "openai",
)
@wrapper_client
async def acreate_realtime_translation_client_secret(
model: str | None = None,
session: Mapping[str, Any] | None = None,
expires_after: Mapping[str, Any] | None = None,
timeout: float | None = None,
**kwargs, # noqa: ANN003 # kwargs-ok: public client wrapper forwards provider-specific options
) -> httpx.Response:
requested_model_name: Final = model or (session or {}).get("model") or "gpt-realtime-translate"
session_config: Final = RealtimeSessionConfig.model_validate(
{ # mutable-ok: Pydantic validates this request-scoped translation session payload
**(session or {}),
"type": "translation",
"model": requested_model_name,
}
)
req: Final = RealtimeClientSecretRequest(
model=requested_model_name,
session=session_config,
expires_after=RealtimeExpiresAfter(**expires_after) if expires_after else None,
)
litellm_logging_obj: Final = kwargs.get("litellm_logging_obj")
if not isinstance(litellm_logging_obj, LiteLLMLogging):
raise TypeError("litellm_logging_obj must be a LiteLLM Logging instance")
litellm_params: Final = GenericLiteLLMParams(**kwargs)
model_name, custom_llm_provider, dynamic_api_key, dynamic_api_base = get_llm_provider(
model=requested_model_name,
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
)
provider_config, resolved_api_base, resolved_api_key = _get_realtime_http_provider_config(
custom_llm_provider=custom_llm_provider,
dynamic_api_base=dynamic_api_base,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
resolved_extra_headers: Final = _get_realtime_http_extra_headers(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
resolved_api_key=resolved_api_key,
extra_headers=kwargs.get("extra_headers"),
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
optional_params={ # mutable-ok: logging owns a mutable request metadata payload
"expires_after": expires_after,
"session": session,
},
litellm_params={ # mutable-ok: logging owns a mutable provider metadata payload
"api_base": resolved_api_base
},
custom_llm_provider=custom_llm_provider,
)
request_data: Final = req.model_dump(
exclude_none=True,
exclude={"model"}, # mutable-ok: Pydantic requires a mutable field-exclusion set
)
request_data["session"] = _with_resolved_session_model(request_data["session"], model_name)
request_data["session"].pop("type", None)
return await base_llm_http_handler.async_realtime_translation_client_secret_handler(
api_base=resolved_api_base,
api_key=resolved_api_key,
request_data=request_data,
logging_obj=litellm_logging_obj,
timeout=timeout or request_timeout,
provider_config=provider_config,
model=model_name,
extra_headers=resolved_extra_headers,
client=kwargs.get("client"),
api_version=litellm_params.api_version,
use_openai_sdk=custom_llm_provider == "openai",
)
@ -229,6 +384,12 @@ async def acreate_realtime_transcription_session(
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
resolved_extra_headers: Final = _get_realtime_http_extra_headers(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
resolved_api_key=resolved_api_key,
extra_headers=kwargs.get("extra_headers"),
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
@ -251,7 +412,7 @@ async def acreate_realtime_transcription_session(
timeout=timeout or request_timeout,
provider_config=provider_config,
model=model_name,
extra_headers=kwargs.get("extra_headers"),
extra_headers=resolved_extra_headers,
client=kwargs.get("client"),
api_version=litellm_params.api_version,
)
@ -307,6 +468,70 @@ async def arealtime_calls(
extra_headers=kwargs.get("extra_headers"),
client=kwargs.get("client"),
api_version=litellm_params.api_version,
use_openai_sdk=custom_llm_provider == "openai",
)
@wrapper_client
async def arealtime_translation_calls(
openai_ephemeral_key: str,
sdp_body: bytes,
model: str | None = None,
session: Mapping[str, Any] | None = None,
timeout: float | None = None,
**kwargs, # noqa: ANN003 # kwargs-ok: public client wrapper forwards provider-specific options
) -> httpx.Response:
requested_model_name: Final = model or "gpt-realtime-translate"
litellm_logging_obj: Final = kwargs.get("litellm_logging_obj")
if not isinstance(litellm_logging_obj, LiteLLMLogging):
raise TypeError("litellm_logging_obj must be a LiteLLM Logging instance")
litellm_params: Final = GenericLiteLLMParams(**kwargs)
model_name, custom_llm_provider, dynamic_api_key, dynamic_api_base = get_llm_provider(
model=requested_model_name,
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
)
provider_config, resolved_api_base, _ = _get_realtime_http_provider_config(
custom_llm_provider=custom_llm_provider,
dynamic_api_base=dynamic_api_base,
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
session_config: Final = _with_resolved_session_model(
{ # mutable-ok: provider routing requires an independently mutable session payload
**(session or {}),
"type": "translation",
"model": model_name,
},
model_name,
)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
optional_params={ # mutable-ok: logging owns a mutable request metadata payload
"realtime_translation_calls": True,
"session": session_config,
},
litellm_params={ # mutable-ok: logging owns a mutable provider metadata payload
"api_base": resolved_api_base
},
custom_llm_provider=custom_llm_provider,
)
return await base_llm_http_handler.async_realtime_calls_handler(
api_base=resolved_api_base,
openai_ephemeral_key=openai_ephemeral_key,
sdp_body=sdp_body,
logging_obj=litellm_logging_obj,
timeout=timeout or request_timeout,
provider_config=provider_config,
model=model_name,
session_config=session_config,
extra_headers=kwargs.get("extra_headers"),
client=kwargs.get("client"),
api_version=litellm_params.api_version,
translation=True,
use_openai_sdk=custom_llm_provider == "openai",
)
@ -356,6 +581,7 @@ async def _arealtime(
client: object | None = None,
timeout: float | None = None,
query_params: RealtimeQueryParams | None = None,
realtime_mode: str = "realtime",
**kwargs,
):
"""
@ -423,6 +649,9 @@ async def _arealtime(
api_base = dynamic_api_base or litellm_params.api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
# set API KEY
api_key = dynamic_api_key or litellm.api_key or litellm.openai_key or get_secret_str("AZURE_API_KEY")
resolved_azure_ad_token = azure_ad_token or litellm_params.azure_ad_token
if not api_key and not resolved_azure_ad_token:
resolved_azure_ad_token = get_azure_ad_token(litellm_params)
api_version = api_version or litellm_params.api_version or "2024-10-01-preview"
@ -431,11 +660,17 @@ async def _arealtime(
or litellm_params.get("realtime_protocol")
or os.environ.get("LITELLM_AZURE_REALTIME_PROTOCOL")
)
realtime_protocol: Final = azure_realtime_protocol_for_client(
configured_realtime_protocol, query_params=query_params, websocket=websocket
)
resolved_azure_ad_token: Final = (
None if api_key else get_azure_ad_token(GenericLiteLLMParams(**kwargs, azure_ad_token=azure_ad_token))
realtime_protocol: Final = _resolve_azure_realtime_protocol(
model=model,
realtime_protocol=(
configured_realtime_protocol
if model in AZURE_GA_REALTIME_MODELS or realtime_mode == "translation"
else azure_realtime_protocol_for_client(
configured_realtime_protocol, query_params=query_params, websocket=websocket
)
),
query_params=query_params,
realtime_mode=realtime_mode,
)
await azure_realtime.async_realtime(
model=model,
@ -449,6 +684,7 @@ async def _arealtime(
logging_obj=litellm_logging_obj,
realtime_protocol=realtime_protocol,
query_params=query_params,
realtime_mode=realtime_mode,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
)
@ -463,9 +699,10 @@ async def _arealtime(
logging_obj=litellm_logging_obj,
api_base=api_base,
api_key=api_key,
client=None,
client=client,
timeout=timeout,
query_params=query_params,
realtime_mode=realtime_mode,
user_api_key_dict=kwargs.get("user_api_key_dict"),
litellm_metadata=_build_litellm_metadata(kwargs),
)

View file

@ -176,6 +176,8 @@ _ERROR_CODE_HTTP_STATUS: Final[Mapping[str, int]] = MappingProxyType(
"insufficient_quota": 429,
"vector_store_timeout": 504,
"invalid_prompt": 400,
"data_residency_mismatch": 400,
"bio_policy": 400,
"invalid_image": 400,
"invalid_image_format": 400,
"invalid_base64_image": 400,

View file

@ -1188,12 +1188,14 @@ class ResponseAPILoggingUtils:
else:
prompt_tokens_details = PromptTokensDetailsWrapper(
cached_tokens=getattr(response_api_usage.input_tokens_details, "cached_tokens", None),
cached_tokens_details=getattr(
response_api_usage.input_tokens_details,
"cached_tokens_details",
None,
),
audio_tokens=getattr(response_api_usage.input_tokens_details, "audio_tokens", None),
text_tokens=getattr(response_api_usage.input_tokens_details, "text_tokens", None),
image_tokens=getattr(response_api_usage.input_tokens_details, "image_tokens", None),
cached_tokens_details=getattr(
response_api_usage.input_tokens_details, "cached_tokens_details", None
),
video_tokens=getattr(response_api_usage.input_tokens_details, "video_tokens", None),
cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None),
web_search_requests=getattr(response_api_usage.input_tokens_details, "web_search_requests", None),

View file

@ -1893,6 +1893,14 @@ class Router:
self.acreate_realtime_transcription_session = self.factory_function(
litellm.acreate_realtime_transcription_session, call_type="acreate_realtime_transcription_session"
)
self.acreate_realtime_translation_client_secret = self.factory_function(
litellm.acreate_realtime_translation_client_secret,
call_type="acreate_realtime_translation_client_secret",
)
self.arealtime_translation_calls = self.factory_function(
litellm.arealtime_translation_calls,
call_type="arealtime_translation_calls",
)
self._aresponses_websocket = self.factory_function(
litellm._aresponses_websocket, call_type="_aresponses_websocket"
)
@ -6516,6 +6524,8 @@ class Router:
"acreate_realtime_client_secret",
"arealtime_calls",
"acreate_realtime_transcription_session",
"acreate_realtime_translation_client_secret",
"arealtime_translation_calls",
"_aresponses_websocket",
"acreate_fine_tuning_job",
"acancel_fine_tuning_job",
@ -6777,6 +6787,8 @@ class Router:
"acreate_realtime_client_secret",
"arealtime_calls",
"acreate_realtime_transcription_session",
"acreate_realtime_translation_client_secret",
"arealtime_translation_calls",
):
return await self._ageneric_api_call_with_fallbacks(
original_function=original_function,

View file

@ -1154,9 +1154,12 @@ AllEmbeddingInputValues = str | list[str] | list[int] | list[list[int]]
OpenAIAudioTranscriptionOptionalParams = Literal[
"language",
"languages",
"keywords",
"prompt",
"temperature",
"response_format",
"stream",
"timestamp_granularities",
"include",
]
@ -2308,6 +2311,16 @@ class OpenAIRealtimeResponseUsage(TypedDict):
output_token_details: NotRequired[ReadOnly[OpenAIRealtimeUsageTokenDetails]]
class OpenAIRealtimeTranslationDurationUsage(TypedDict):
type: ReadOnly[Literal["duration"]]
output_seconds: ReadOnly[float]
class OpenAIRealtimeTranslationClosedEvent(TypedDict):
type: ReadOnly[Literal["session.closed"]]
usage: ReadOnly[OpenAIRealtimeTranslationDurationUsage]
class OpenAIRealtimeEventTypes(Enum):
SESSION_CREATED = "session.created"
# Beta delta event names
@ -2350,6 +2363,7 @@ OpenAIRealtimeEvents = (
| OpenAIRealtimeInputAudioTranscriptionCompleted
| OpenAIRealtimeTranscriptionSessionCreated
| OpenAIRealtimeErrorEvent
| OpenAIRealtimeTranslationClosedEvent
)
OpenAIRealtimeStreamList = list[OpenAIRealtimeEvents]

View file

@ -1,6 +1,7 @@
from typing import Any, Literal
from collections.abc import Mapping, Sequence
from typing import Any, Final, Literal
from pydantic import BaseModel
from pydantic import BaseModel, ConfigDict
from typing_extensions import ReadOnly, TypedDict
from .llms.openai import (
@ -10,6 +11,7 @@ from .llms.openai import (
)
ALL_DELTA_TYPES = Literal["text", "audio"]
RealtimeSessionType: Final = Literal["realtime", "transcription", "translation"]
class RealtimeResponseTransformInput(TypedDict):
@ -64,6 +66,41 @@ class RealtimeExpiresAfter(BaseModel):
seconds: int | None = None
class RealtimeAudioTranscriptionConfig(BaseModel):
model_config = ConfigDict(extra="allow")
model: str | None = None
delay: Literal["minimal", "low", "medium", "high", "xhigh"] | None = None
keywords: Sequence[str] | None = None
language: str | None = None
languages: Sequence[str] | None = None
prompt: str | None = None
class RealtimeAudioInputConfig(BaseModel):
model_config = ConfigDict(extra="allow")
format: str | Mapping[str, object] | None = None
noise_reduction: Mapping[str, object] | None = None
transcription: RealtimeAudioTranscriptionConfig | None = None
turn_detection: Mapping[str, object] | None = None
class RealtimeAudioOutputConfig(BaseModel):
model_config = ConfigDict(extra="allow")
format: str | Mapping[str, object] | None = None
language: str | None = None
voice: str | Mapping[str, object] | None = None
class RealtimeSessionAudioConfig(BaseModel):
model_config = ConfigDict(extra="allow")
input: RealtimeAudioInputConfig | None = None
output: RealtimeAudioOutputConfig | None = None
class RealtimeSessionConfig(BaseModel):
"""
Session configuration nested inside the client_secrets request body.
@ -75,10 +112,10 @@ class RealtimeSessionConfig(BaseModel):
model_config = {"extra": "allow"}
type: str | None = None
type: RealtimeSessionType | None = None
model: str | None = None
instructions: str | None = None
audio: dict[str, object] | None = None
audio: RealtimeSessionAudioConfig | None = None
include: list[str] | None = None
max_output_tokens: int | str | None = None
output_modalities: list[str] | None = None
@ -132,12 +169,15 @@ class RealtimeTranscriptionSessionRequest(BaseModel):
# LiteLLM-only routing hint — stripped before forwarding upstream.
model: str | None = None
input_audio_transcription: dict[str, Any] | None = None
audio: RealtimeSessionAudioConfig | None = None
def resolved_model(self) -> str | None:
if self.model:
return self.model
if self.input_audio_transcription:
return self.input_audio_transcription.get("model")
if self.audio and self.audio.input and self.audio.input.transcription:
return self.audio.input.transcription.model
return None

View file

@ -449,6 +449,11 @@ class CallTypes(str, Enum):
asearch = "asearch"
arealtime = "_arealtime"
aresponses_websocket = "_aresponses_websocket"
acreate_realtime_client_secret = "acreate_realtime_client_secret"
arealtime_calls = "arealtime_calls"
acreate_realtime_transcription_session = "acreate_realtime_transcription_session"
acreate_realtime_translation_client_secret = "acreate_realtime_translation_client_secret"
arealtime_translation_calls = "arealtime_translation_calls"
create_batch = "create_batch"
acreate_batch = "acreate_batch"
aretrieve_batch = "aretrieve_batch"
@ -677,10 +682,18 @@ CallTypesLiteral = Literal[
"acreate_realtime_client_secret",
"arealtime_calls",
"acreate_realtime_transcription_session",
"acreate_realtime_translation_client_secret",
"arealtime_translation_calls",
]
# Mapping of API routes to their corresponding call types
API_ROUTE_TO_CALL_TYPES: Final[Mapping[str, Sequence[CallTypes]]] = {
"/v1/realtime/translations/client_secrets": (CallTypes.acreate_realtime_translation_client_secret,),
"/realtime/translations/client_secrets": (CallTypes.acreate_realtime_translation_client_secret,),
"/openai/v1/realtime/translations/client_secrets": (CallTypes.acreate_realtime_translation_client_secret,),
"/v1/realtime/translations/calls": (CallTypes.arealtime_translation_calls,),
"/realtime/translations/calls": (CallTypes.arealtime_translation_calls,),
"/openai/v1/realtime/translations/calls": (CallTypes.arealtime_translation_calls,),
# Chat Completions
"/chat/completions": [CallTypes.acompletion, CallTypes.completion],
"/v1/chat/completions": [CallTypes.acompletion, CallTypes.completion],
@ -993,9 +1006,12 @@ API_ROUTE_TO_CALL_TYPES: Final[Mapping[str, Sequence[CallTypes]]] = {
"/v1/responses/{response_id}/input_items": (CallTypes.alist_input_items,),
"/openai/v1/responses/{response_id}/input_items": (CallTypes.alist_input_items,),
# Realtime API
"/realtime": [CallTypes.arealtime],
"/v1/realtime": [CallTypes.arealtime],
"/openai/v1/realtime": [CallTypes.arealtime],
"/realtime": (CallTypes.arealtime,),
"/v1/realtime": (CallTypes.arealtime,),
"/openai/v1/realtime": (CallTypes.arealtime,),
"/realtime/translations": (CallTypes.arealtime,),
"/v1/realtime/translations": (CallTypes.arealtime,),
"/openai/v1/realtime/translations": (CallTypes.arealtime,),
# Provider-specific routes
"/anthropic/v1/messages": [CallTypes.anthropic_messages],
# Google GenAI routes
@ -1724,6 +1740,8 @@ class PromptTokensDetailsWrapper(
image_tokens: int | None = None
"""Image tokens sent to the model."""
cached_tokens_details: CachedTokensDetails | None = None
video_tokens: int | None = None
"""Video tokens sent to the model."""
@ -2708,18 +2726,26 @@ class TranscriptionUsageTokensObject(BaseModel):
input_tokens: int
output_tokens: int
total_tokens: int
input_token_details: TranscriptionUsageInputTokenDetailsObject
input_token_details: TranscriptionUsageInputTokenDetailsObject | None = None
class TranscriptionDetectedLanguage(BaseModel):
code: str
class TranscriptionResponse(OpenAIObject):
text: str | None = None
usage: TranscriptionUsageDurationObject | TranscriptionUsageTokensObject | None = None
languages: Sequence[TranscriptionDetectedLanguage] | None = None
_hidden_params: dict = {}
_response_headers: dict | None = None
def __init__(self, text=None) -> None:
super().__init__(text=text)
def __init__(self, text=None, usage=None, languages=None, **kwargs) -> None: # noqa: ANN003 # OpenAI-compatible response accepts provider extension fields
super().__init__(text=text, usage=usage, languages=languages, **kwargs)
def set_audio_transcription_duration(self, duration: float) -> None:
self._hidden_params["audio_transcription_duration"] = duration
def __contains__(self, key) -> bool:
# Define custom behavior for the 'in' operator

View file

@ -1242,6 +1242,8 @@ def function_setup(
applied_guardrails=applied_guardrails,
supports_correlation_logging=is_async_call,
)
if logging_obj is None:
raise RuntimeError("LiteLLM logging initialization returned no logger")
## check if metadata is passed in
litellm_params: Final[dict[str, object]] = {"api_base": ""}
@ -1760,6 +1762,12 @@ def client(original_function):
chunks.append(chunk)
return litellm.stream_chunk_builder(chunks, messages=kwargs.get("messages", None))
else:
if call_type == CallTypes.transcription.value and isinstance(result, openai.Stream):
from litellm.litellm_core_utils.audio_utils.transcription_streaming import (
wrap_transcription_stream,
)
result = wrap_transcription_stream(result, logging_obj, start_time)
# RETURN RESULT
update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata")
update_response_metadata(
@ -2062,6 +2070,12 @@ def client(original_function):
chunks.append(chunk)
return litellm.stream_chunk_builder(chunks, messages=kwargs.get("messages", None))
else:
if call_type == CallTypes.atranscription.value and isinstance(result, openai.AsyncStream):
from litellm.litellm_core_utils.audio_utils.transcription_streaming import (
wrap_transcription_stream,
)
result = wrap_transcription_stream(result, logging_obj, start_time)
_update_response_metadata(
result=result,
logging_obj=logging_obj,
@ -3441,10 +3455,13 @@ def get_optional_params_transcription(
model: str,
custom_llm_provider: str,
language: str | None = None,
languages: Sequence[str] | None = None,
keywords: Sequence[str] | None = None,
prompt: str | None = None,
response_format: str | None = None,
temperature: int | None = None,
timestamp_granularities: list[Literal["word", "segment"]] | None = None,
stream: bool | None = None,
drop_params: bool | None = None,
**kwargs,
):
@ -3454,6 +3471,7 @@ def get_optional_params_transcription(
passed_params: Final = locals()
passed_params.pop("OPENAI_TRANSCRIPTION_PARAMS")
passed_params.pop("model")
custom_llm_provider = passed_params.pop("custom_llm_provider")
drop_params = normalize_drop_params(passed_params.pop("drop_params"))
special_params: Final[Mapping[str, object]] = passed_params.pop("kwargs")
@ -3462,10 +3480,13 @@ def get_optional_params_transcription(
default_params: Final = {
"language": None,
"languages": None,
"keywords": None,
"prompt": None,
"response_format": None,
"temperature": None, # openai defaults this to 0
"timestamp_granularities": None,
"stream": None,
}
non_default_params = {k: v for k, v in passed_params.items() if (k in default_params and v != default_params[k])}
@ -3522,6 +3543,9 @@ def get_optional_params_transcription(
openai_params=OPENAI_TRANSCRIPTION_PARAMS,
additional_drop_params=kwargs.get("additional_drop_params", None),
)
extra_body: Final = optional_params.get("extra_body")
if isinstance(extra_body, dict) and not extra_body:
optional_params.pop("extra_body")
return optional_params
@ -6083,6 +6107,7 @@ def _get_model_info_helper(
),
cache_read_input_token_cost=_model_info.get("cache_read_input_token_cost", None),
cache_read_input_audio_token_cost=_model_info.get("cache_read_input_audio_token_cost", None),
cache_read_input_image_token_cost=_model_info.get("cache_read_input_image_token_cost", None),
prompt_cache_min_tokens=_model_info.get("prompt_cache_min_tokens", None),
cache_read_input_token_cost_above_200k_tokens=_model_info.get(
"cache_read_input_token_cost_above_200k_tokens", None
@ -8903,7 +8928,13 @@ class ProviderConfigManager:
return XAIAudioTranscriptionConfig()
elif litellm.LlmProviders.OPENAI == provider:
if "gpt-4o" in model:
if model == "gpt-transcribe":
from litellm.llms.openai.transcriptions.gpt_transformation import (
OpenAIGPTTranscribeAudioTranscriptionConfig,
)
return OpenAIGPTTranscribeAudioTranscriptionConfig()
elif "gpt-4o" in model:
return litellm.OpenAIGPTAudioTranscriptionConfig()
else:
return litellm.OpenAIWhisperAudioTranscriptionConfig()

View file

@ -72,6 +72,10 @@
- {id: llm.bedrock_native.bedrock_invoke.basic.stream.works, module: llm, tier: P1, subject_endpoint: bedrock_native, route: bedrock_invoke, capability: basic, streaming: stream, assertions: [works], source: "vendor strategy §9.12 / LIT-4778", rationale: "Bedrock native invoke stream"}
- {id: llm.bedrock_native.bedrock_invoke.input_validation.nonstream.works, module: llm, tier: P1, subject_endpoint: bedrock_native, route: bedrock_invoke, capability: input_validation, streaming: nonstream, assertions: [works], source: "vendor strategy §9.12 / LIT-4778", rationale: "Bedrock invoke missing fields and invalid temperature"}
- {id: llm.ocr.openai.input_validation.nonstream.works, module: llm, tier: P1, subject_endpoint: ocr, route: openai, capability: input_validation, streaming: nonstream, assertions: [works], source: "vendor strategy §9.13 / LIT-4778", rationale: "OCR missing document rejected"}
- {id: llm.realtime.openai.translation.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: openai, capability: translation, streaming: stream, assertions: [works], source: "realtime_endpoints/endpoints.py", rationale: "Dedicated translation client-secret, raw SDP, and WebSocket paths emit translated audio and transcript deltas"}
- {id: llm.realtime.azure_openai.translation.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: azure_openai, capability: translation, streaming: stream, assertions: [works], source: "azure/realtime/handler.py", rationale: "Azure GA translation session emits translated audio and transcript deltas"}
- {id: llm.realtime.openai.transcription.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: openai, capability: transcription, streaming: stream, assertions: [works], source: "openai/realtime/handler.py", rationale: "gpt-live-transcribe and gpt-realtime-whisper emit live transcript deltas"}
- {id: llm.realtime.azure_openai.transcription.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: azure_openai, capability: transcription, streaming: stream, assertions: [works], source: "azure/realtime/handler.py", rationale: "Azure GA live transcription emits transcript deltas"}
- {id: llm.rerank.bedrock.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "llms/bedrock/rerank/handler.py", rationale: "Bedrock rerank"}
- {id: llm.rerank.together_ai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: together_ai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/together_ai/rerank/handler.py", rationale: "Together rerank"}
- {id: llm.images_generations.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_image_generation_e2e.py:22", rationale: "OpenAI image gen, b64/url"}
@ -89,6 +93,8 @@
- {id: llm.audio_speech.vertex.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_speech, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "vertex_ai/text_to_speech/text_to_speech_handler.py", rationale: "Vertex TTS"}
- {id: llm.audio_transcriptions.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_transcriptions, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "openai/transcriptions/handler.py", rationale: "OpenAI Whisper"}
- {id: llm.audio_transcriptions.openai.input_validation.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_transcriptions, route: openai, capability: input_validation, streaming: nonstream, assertions: [works], source: "vendor strategy §9.7 / LIT-4778", rationale: "Transcription empty file and missing model are rejected"}
- {id: llm.audio_transcriptions.openai.transcription.stream.works, module: llm, tier: P0, subject_endpoint: audio_transcriptions, route: openai, capability: transcription, streaming: stream, assertions: [works], source: "openai/transcriptions/handler.py", rationale: "gpt-transcribe streams typed transcript delta and done events"}
- {id: llm.audio_transcriptions.azure_openai.transcription.stream.works, module: llm, tier: P0, subject_endpoint: audio_transcriptions, route: azure_openai, capability: transcription, streaming: stream, assertions: [works], source: "azure/audio_transcriptions.py", rationale: "Azure gpt-transcribe streams typed transcript delta and done events over the v1 API"}
- {id: llm.audio_transcriptions.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: audio_transcriptions, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "azure/audio_transcriptions.py", rationale: "Azure STT"}
- {id: llm.audio_transcriptions.soniox.basic.nonstream.works, module: llm, tier: P2, subject_endpoint: audio_transcriptions, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "soniox/audio_transcription/handler.py", rationale: "Soniox via OpenAI-compat (smoke)"}
- {id: llm.audio_transcriptions.nvidia_riva.basic.nonstream.works, module: llm, tier: P2, subject_endpoint: audio_transcriptions, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "nvidia_riva/audio_transcription/handler.py", rationale: "NVIDIA Riva (smoke)"}

View file

@ -84,6 +84,8 @@ LlmCapability = Literal[
"tool_search",
"tool_search_history",
"tool_use",
"transcription",
"translation",
"vision",
"web_search",
"web_search_server_tool",
@ -153,14 +155,7 @@ class OtherCell(_Base):
Cell = Annotated[
LlmCell
| MgmtCell
| McpCell
| ReliabilityCell
| QuotaCell
| LoggingCell
| GuardrailCell
| OtherCell,
LlmCell | MgmtCell | McpCell | ReliabilityCell | QuotaCell | LoggingCell | GuardrailCell | OtherCell,
Field(discriminator="module"),
]

View file

@ -17,6 +17,7 @@ from litellm.litellm_core_utils.realtime_streaming import (
client_sent_openai_beta_realtime_header,
)
from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
@ -602,6 +603,42 @@ async def test_client_ack_messages_keeps_beta_session_shape_for_beta_backend():
assert "audio" not in session
@pytest.mark.asyncio
async def test_translation_session_update_omits_session_type():
client_ws = MagicMock()
client_ws.scope = {"headers": []}
client_ws.receive_text = AsyncMock(
side_effect=[
json.dumps(
{
"type": "session.update",
"session": {
"type": "translation",
"audio": {"output": {"language": "fr"}},
},
}
),
Exception("connection closed"),
]
)
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.pre_call = MagicMock()
streaming = RealTimeStreaming(
client_ws,
backend_ws,
logging_obj,
translation_session=True,
)
await streaming.client_ack_messages()
sent_to_backend = json.loads(backend_ws.send.call_args_list[0].args[0])
assert "type" not in sent_to_backend["session"]
assert sent_to_backend["session"]["audio"]["output"]["language"] == "fr"
def test_translate_event_to_beta_renames_delta_types():
ev = RealTimeStreaming._translate_event_to_beta(
{"type": "response.output_audio.delta", "delta": "abc", "event_id": "e1"}
@ -1023,6 +1060,93 @@ async def test_transcription_session_update_enforces_authorized_nested_model():
assert streaming._is_transcription_session is True
@pytest.mark.asyncio
async def test_translation_session_update_rejects_disallowed_nested_transcription_model() -> None:
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
streaming = RealTimeStreaming(
MagicMock(),
backend_ws,
MagicMock(),
model="gpt-realtime-translate",
user_api_key_dict=UserAPIKeyAuth(models=["gpt-realtime-translate"]),
translation_session=True,
)
with pytest.raises(Exception, match="Tried to access gpt-live-transcribe"):
await streaming._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {
"type": "translation",
"audio": {
"input": {
"transcription": {"model": "gpt-live-transcribe"},
}
},
},
}
)
)
backend_ws.send.assert_not_awaited()
assert streaming._is_transcription_session is False
@pytest.mark.asyncio
async def test_translation_session_update_binds_nested_transcription_model() -> None:
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
streaming = RealTimeStreaming(
MagicMock(),
backend_ws,
MagicMock(),
model="gpt-realtime-translate",
user_api_key_dict=UserAPIKeyAuth(models=["gpt-realtime-translate", "gpt-realtime-whisper"]),
translation_session=True,
)
await streaming._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {
"type": "translation",
"audio": {
"input": {
"transcription": {"model": "gpt-realtime-whisper", "language": "en"},
}
},
},
}
)
)
await streaming._send_to_backend(
json.dumps(
{
"type": "session.update",
"session": {
"type": "translation",
"audio": {
"input": {
"transcription": {"model": "gpt-live-transcribe", "language": "fr"},
}
},
},
}
)
)
first_sent = json.loads(backend_ws.send.await_args_list[0].args[0])
second_sent = json.loads(backend_ws.send.await_args_list[1].args[0])
assert first_sent["session"]["audio"]["input"]["transcription"]["model"] == "gpt-realtime-whisper"
assert second_sent["session"]["audio"]["input"]["transcription"] == {
"model": "gpt-realtime-whisper",
"language": "fr",
}
assert streaming._is_transcription_session is False
@pytest.mark.asyncio
async def test_normal_realtime_session_keeps_nested_transcription_model():
backend_ws = MagicMock()
@ -2786,6 +2910,69 @@ def test_store_message_skips_pydantic_for_unlogged_audio_delta():
assert streaming.messages == []
def test_translation_audio_duration_is_finalized_once():
import base64
streaming = RealTimeStreaming(
websocket=MagicMock(),
backend_ws=MagicMock(),
logging_obj=MagicMock(),
model="gpt-realtime-translate",
translation_session=True,
)
payload = base64.b64encode(bytes(48000)).decode()
streaming._capture_translation_output_audio({"type": "session.output_audio.delta", "delta": payload})
streaming._finalize_translation_usage()
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": 1.0}
def test_translation_audio_duration_uses_session_output_format():
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.created",
"session": {"audio": {"output": {"format": {"type": "audio/pcmu", "rate": 8000}}}},
}
)
streaming._capture_translation_output_audio(
{
"type": "session.output_audio.delta",
"delta": base64.b64encode(bytes(8000)).decode(),
}
)
streaming._finalize_translation_usage()
assert streaming.messages[-1]["usage"] == {"type": "duration", "output_seconds": 1.0}
def test_translation_does_not_duplicate_provider_duration_usage():
streaming = RealTimeStreaming(
websocket=MagicMock(),
backend_ws=MagicMock(),
logging_obj=MagicMock(),
model="gpt-realtime-translate",
translation_session=True,
)
streaming._translation_output_audio_bytes = 48000
streaming.messages.append({"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
@pytest.mark.asyncio
async def test_audio_delta_frame_parsed_at_most_once():
client_ws = _beta_client_ws()

View file

@ -474,6 +474,11 @@ def test_default_max_retries_env_var_reaches_azure_sdk_client():
"avector_store_search",
"acreate_skill",
"acreate_interaction",
"acreate_realtime_client_secret",
"acreate_realtime_transcription_session",
"acreate_realtime_translation_client_secret",
"arealtime_calls",
"arealtime_translation_calls",
]
],
)

View file

@ -1,16 +1,31 @@
import json
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from litellm.llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
from openai import AsyncOpenAI, omit
class DummySDKConnectionManager:
def __init__(self, connection):
self.connection = connection
@pytest.mark.parametrize(
"api_base", ["https://api.openai.com/v1", "https://api.openai.com"]
)
async def __aenter__(self):
return self.connection
async def __aexit__(self, exc_type, exc, tb):
return None
def make_realtime_sdk_client():
connection = MagicMock()
connection.send_raw = AsyncMock()
connection.recv_bytes = AsyncMock()
connection.close = AsyncMock()
client = MagicMock(spec=AsyncOpenAI)
client.realtime.connect = MagicMock(return_value=DummySDKConnectionManager(connection))
return client
@pytest.mark.parametrize("api_base", ["https://api.openai.com/v1", "https://api.openai.com"])
def test_openai_realtime_handler_url_construction(api_base):
from litellm.llms.openai.realtime.handler import OpenAIRealtime
@ -59,12 +74,8 @@ def test_openai_realtime_handler_model_parameter_inclusion():
api_base = "https://api.openai.com/"
# Test with just model parameter
query_params_model_only: RealtimeQueryParams = {
"model": "gpt-4o-mini-realtime-preview"
}
url = handler._construct_url(
api_base=api_base, query_params=query_params_model_only
)
query_params_model_only: RealtimeQueryParams = {"model": "gpt-4o-mini-realtime-preview"}
url = handler._construct_url(api_base=api_base, query_params=query_params_model_only)
# Verify the URL structure
assert url.startswith("wss://api.openai.com/v1/realtime?")
@ -75,9 +86,7 @@ def test_openai_realtime_handler_model_parameter_inclusion():
"model": "gpt-4o-mini-realtime-preview",
"intent": "chat",
}
url_with_extras = handler._construct_url(
api_base=api_base, query_params=query_params_with_extras
)
url_with_extras = handler._construct_url(api_base=api_base, query_params=query_params_with_extras)
# Verify both parameters are included
assert url_with_extras.startswith("wss://api.openai.com/v1/realtime?")
@ -91,11 +100,6 @@ def test_openai_realtime_handler_model_parameter_inclusion():
assert expected_pattern in url_with_extras
import asyncio
import pytest
@pytest.mark.asyncio
async def test_async_realtime_success():
from litellm.llms.openai.realtime.handler import OpenAIRealtime
@ -109,27 +113,10 @@ async def test_async_realtime_success():
dummy_websocket = AsyncMock()
dummy_logging_obj = MagicMock()
mock_backend_ws = AsyncMock()
class DummyAsyncContextManager:
def __init__(self, value):
self.value = value
async def __aenter__(self):
return self.value
async def __aexit__(self, exc_type, exc, tb):
return None
shared_context = get_shared_realtime_ssl_context()
with (
patch(
"websockets.connect", return_value=DummyAsyncContextManager(mock_backend_ws)
) as mock_ws_connect,
patch(
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming,
):
sdk_client = make_realtime_sdk_client()
with patch( # test-quality-ok: orchestration test replaces the unbounded streaming loop
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming:
mock_streaming_instance = MagicMock()
mock_realtime_streaming.return_value = mock_streaming_instance
mock_streaming_instance.bidirectional_forward = AsyncMock()
@ -141,6 +128,7 @@ async def test_async_realtime_success():
api_base=api_base,
api_key=api_key,
query_params=query_params,
client=sdk_client,
)
mock_realtime_streaming.assert_called_once()
@ -164,28 +152,10 @@ async def test_async_realtime_url_contains_model():
dummy_websocket = AsyncMock()
dummy_logging_obj = MagicMock()
mock_backend_ws = AsyncMock()
class DummyAsyncContextManager:
def __init__(self, value):
self.value = value
async def __aenter__(self):
return self.value
async def __aexit__(self, exc_type, exc, tb):
return None
shared_context = get_shared_realtime_ssl_context()
with (
patch(
"websockets.connect", return_value=DummyAsyncContextManager(mock_backend_ws)
) as mock_ws_connect,
patch(
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming,
):
sdk_client = make_realtime_sdk_client()
with patch( # test-quality-ok: orchestration test replaces the unbounded streaming loop
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming:
mock_streaming_instance = MagicMock()
mock_realtime_streaming.return_value = mock_streaming_instance
mock_streaming_instance.bidirectional_forward = AsyncMock()
@ -197,30 +167,48 @@ async def test_async_realtime_url_contains_model():
api_base=api_base,
api_key=api_key,
query_params=query_params,
client=sdk_client,
)
# Verify websockets.connect was called with the correct URL
mock_ws_connect.assert_called_once()
called_url = mock_ws_connect.call_args[0][0]
# Verify the URL contains the model parameter
assert called_url.startswith("wss://api.openai.com/v1/realtime?")
assert f"model={model}" in called_url
# Verify proper headers were set (GA default: no OpenAI-Beta unless client sent it)
called_kwargs = mock_ws_connect.call_args[1]
assert "additional_headers" in called_kwargs
additional_headers = called_kwargs["additional_headers"]
sdk_client.realtime.connect.assert_called_once()
called_kwargs = sdk_client.realtime.connect.call_args.kwargs
assert called_kwargs["model"] == model
additional_headers = called_kwargs["extra_headers"]
assert additional_headers["Authorization"] == f"Bearer {api_key}"
assert "OpenAI-Beta" not in additional_headers
# Verify SSL is configured (should be an SSLContext or True, not None or False)
assert called_kwargs["ssl"] is not None
assert called_kwargs["ssl"] is not False
assert called_kwargs["max_retries"] == 0
mock_realtime_streaming.assert_called_once()
mock_streaming_instance.bidirectional_forward.assert_awaited_once()
@pytest.mark.asyncio
async def test_async_realtime_transcription_omits_sdk_model_query():
from litellm.llms.openai.realtime.handler import OpenAIRealtime
handler = OpenAIRealtime()
websocket = AsyncMock()
logging_obj = MagicMock()
sdk_client = make_realtime_sdk_client()
with patch( # test-quality-ok: orchestration test replaces the unbounded streaming loop
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming:
mock_realtime_streaming.return_value.bidirectional_forward = AsyncMock()
await handler.async_realtime(
model="gpt-live-transcribe",
websocket=websocket,
logging_obj=logging_obj,
api_key="test-key",
query_params={"intent": "transcription"},
client=sdk_client,
)
called_kwargs = sdk_client.realtime.connect.call_args.kwargs
assert called_kwargs["model"] is omit
assert called_kwargs["extra_query"] == {"intent": "transcription"}
@pytest.mark.asyncio
async def test_async_realtime_forwards_openai_beta_header_when_client_sends_it():
"""Upstream WS gets OpenAI-Beta: realtime=v1 only when the client WebSocket included it."""
@ -240,26 +228,10 @@ async def test_async_realtime_forwards_openai_beta_header_when_client_sends_it()
]
}
dummy_logging_obj = MagicMock()
mock_backend_ws = AsyncMock()
class DummyAsyncContextManager:
def __init__(self, value):
self.value = value
async def __aenter__(self):
return self.value
async def __aexit__(self, exc_type, exc, tb):
return None
with (
patch(
"websockets.connect", return_value=DummyAsyncContextManager(mock_backend_ws)
) as mock_ws_connect,
patch(
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming,
):
sdk_client = make_realtime_sdk_client()
with patch( # test-quality-ok: orchestration test replaces the unbounded streaming loop
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming:
mock_streaming_instance = MagicMock()
mock_realtime_streaming.return_value = mock_streaming_instance
mock_streaming_instance.bidirectional_forward = AsyncMock()
@ -271,11 +243,12 @@ async def test_async_realtime_forwards_openai_beta_header_when_client_sends_it()
api_base=api_base,
api_key=api_key,
query_params=query_params,
client=sdk_client,
)
mock_ws_connect.assert_called_once()
called_kwargs = mock_ws_connect.call_args[1]
additional_headers = called_kwargs["additional_headers"]
sdk_client.realtime.connect.assert_called_once()
called_kwargs = sdk_client.realtime.connect.call_args.kwargs
additional_headers = called_kwargs["extra_headers"]
assert additional_headers["Authorization"] == f"Bearer {api_key}"
assert additional_headers["OpenAI-Beta"] == "realtime=v1"
@ -300,28 +273,10 @@ async def test_async_realtime_uses_max_size_parameter():
dummy_websocket = AsyncMock()
dummy_logging_obj = MagicMock()
mock_backend_ws = AsyncMock()
class DummyAsyncContextManager:
def __init__(self, value):
self.value = value
async def __aenter__(self):
return self.value
async def __aexit__(self, exc_type, exc, tb):
return None
shared_context = get_shared_realtime_ssl_context()
with (
patch(
"websockets.connect", return_value=DummyAsyncContextManager(mock_backend_ws)
) as mock_ws_connect,
patch(
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming,
):
sdk_client = make_realtime_sdk_client()
with patch( # test-quality-ok: orchestration test replaces the unbounded streaming loop
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming:
mock_streaming_instance = MagicMock()
mock_realtime_streaming.return_value = mock_streaming_instance
mock_streaming_instance.bidirectional_forward = AsyncMock()
@ -333,20 +288,14 @@ async def test_async_realtime_uses_max_size_parameter():
api_base=api_base,
api_key=api_key,
query_params=query_params,
client=sdk_client,
)
# Verify websockets.connect was called with the max_size parameter
mock_ws_connect.assert_called_once()
called_kwargs = mock_ws_connect.call_args[1]
# Verify max_size is set (default None for unlimited, matching OpenAI's SDK)
assert "max_size" in called_kwargs
assert called_kwargs["max_size"] is None
# Verify SSL is configured (should be an SSLContext or True, not None or False)
assert called_kwargs["ssl"] is not None
assert called_kwargs["ssl"] is not False
# Default should be None (unlimited) to match OpenAI's official agents SDK
# https://github.com/openai/openai-agents-python/blob/cf1b933660e44fd37b4350c41febab8221801409/src/agents/realtime/openai_realtime.py#L235
sdk_client.realtime.connect.assert_called_once()
called_kwargs = sdk_client.realtime.connect.call_args.kwargs
connection_options = called_kwargs["websocket_connection_options"]
assert connection_options["max_size"] is REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
assert "ssl" not in connection_options
mock_realtime_streaming.assert_called_once()
mock_streaming_instance.bidirectional_forward.assert_awaited_once()
@ -371,27 +320,10 @@ async def test_async_realtime_ws_url_has_no_ssl():
dummy_websocket = AsyncMock()
dummy_logging_obj = MagicMock()
mock_backend_ws = AsyncMock()
class DummyAsyncContextManager:
def __init__(self, value):
self.value = value
async def __aenter__(self):
return self.value
async def __aexit__(self, exc_type, exc, tb):
return None
with (
patch(
"websockets.connect", return_value=DummyAsyncContextManager(mock_backend_ws)
) as mock_ws_connect,
patch(
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming,
):
sdk_client = make_realtime_sdk_client()
with patch( # test-quality-ok: orchestration test replaces the unbounded streaming loop
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as mock_realtime_streaming:
mock_streaming_instance = MagicMock()
mock_realtime_streaming.return_value = mock_streaming_instance
mock_streaming_instance.bidirectional_forward = AsyncMock()
@ -403,19 +335,13 @@ async def test_async_realtime_ws_url_has_no_ssl():
api_base=api_base,
api_key=api_key,
query_params=query_params,
client=sdk_client,
)
# Verify websockets.connect was called
mock_ws_connect.assert_called_once()
called_url = mock_ws_connect.call_args[0][0]
called_kwargs = mock_ws_connect.call_args[1]
# Verify URL was converted from http:// to ws://
assert called_url.startswith("ws://localhost:8113/v1/realtime?")
assert f"model={model}" in called_url
# Verify ssl is None for ws:// URLs (the fix for issue #19222)
assert called_kwargs["ssl"] is None
sdk_client.realtime.connect.assert_called_once()
called_kwargs = sdk_client.realtime.connect.call_args.kwargs
assert called_kwargs["model"] == model
assert "ssl" not in called_kwargs["websocket_connection_options"]
@pytest.mark.asyncio
@ -465,3 +391,65 @@ async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_
assert event["error"]["type"] == "server_error"
assert "401" in event["error"]["message"]
assert closed and closed[0][0] == 1008
def test_translation_url_uses_dedicated_path():
from litellm.llms.openai.realtime.handler import OpenAIRealtime
handler = OpenAIRealtime()
url = handler._construct_url(
api_base="https://api.openai.com/v1",
query_params={"model": "gpt-realtime-translate"},
realtime_mode="translation",
)
assert url == "wss://api.openai.com/v1/realtime/translations?model=gpt-realtime-translate"
@pytest.mark.asyncio
async def test_translation_websocket_uses_direct_transport():
from litellm.llms.openai.realtime.handler import OpenAIRealtime
backend = AsyncMock()
class TranslationConnectionManager:
async def __aenter__(self):
return backend
async def __aexit__(self, exc_type, exc, tb):
return None
websocket = MagicMock()
websocket.scope = {"headers": []}
websocket.close = AsyncMock()
logging_obj = MagicMock()
handler = OpenAIRealtime()
expected_url = "wss://api.openai.com/v1/realtime/translations?model=gpt-realtime-translate"
assert (
handler._construct_url(
api_base="https://api.openai.com/v1",
query_params={"model": "gpt-realtime-translate"},
realtime_mode="translation",
)
== expected_url
)
with (
patch("websockets.connect", return_value=TranslationConnectionManager()) as connect,
patch( # test-quality-ok: transport test replaces the unbounded streaming loop
"litellm.llms.openai.realtime.handler.RealTimeStreaming"
) as streaming,
):
streaming.return_value.bidirectional_forward = AsyncMock()
await handler.async_realtime(
model="gpt-realtime-translate",
websocket=websocket,
logging_obj=logging_obj,
api_base="https://api.openai.com/v1",
api_key="sk-test",
query_params={"model": "gpt-realtime-translate"},
realtime_mode="translation",
)
connect.assert_called_once()
assert connect.call_args.args[0] == expected_url
assert streaming.call_args.kwargs["translation_session"] is True

View file

@ -21,9 +21,7 @@ from litellm.types.realtime import RealtimeTranscriptionSessionRequest
def test_openai_transcription_session_url():
cfg = OpenAIRealtimeHTTPConfig()
assert (
cfg.get_transcription_session_url(
api_base="https://api.openai.com", model="gpt-realtime-whisper"
)
cfg.get_transcription_session_url(api_base="https://api.openai.com", model="gpt-realtime-whisper")
== "https://api.openai.com/v1/realtime/transcription_sessions"
)
@ -32,9 +30,7 @@ def test_openai_transcription_session_url_strips_trailing_v1():
"""A /v1 suffix must not be duplicated in the path."""
cfg = OpenAIRealtimeHTTPConfig()
assert (
cfg.get_transcription_session_url(
api_base="https://api.openai.com/v1", model="gpt-realtime-whisper"
)
cfg.get_transcription_session_url(api_base="https://api.openai.com/v1", model="gpt-realtime-whisper")
== "https://api.openai.com/v1/realtime/transcription_sessions"
)
@ -46,9 +42,18 @@ def test_azure_transcription_session_url_uses_deployment_and_api_version():
model="whisper-deploy",
api_version="2025-04-01-preview",
)
assert (
url
== "https://my.openai.azure.com/openai/realtime/transcription_sessions?api-version=2025-04-01-preview"
assert url == "https://my.openai.azure.com/openai/realtime/transcription_sessions?api-version=2025-04-01-preview"
@pytest.mark.parametrize("api_version", [None, "v1", "latest", "preview"])
def test_azure_ga_realtime_http_urls(api_version):
cfg = AzureRealtimeHTTPConfig()
base = "https://my.openai.azure.com"
assert cfg.get_complete_url(base, "gpt-realtime-2", api_version) == (f"{base}/openai/v1/realtime/client_secrets")
assert cfg.get_realtime_calls_url(base, "gpt-realtime-2", api_version) == (f"{base}/openai/v1/realtime/calls")
assert cfg.get_transcription_session_url(base, "gpt-live-transcribe", api_version) == (
f"{base}/openai/v1/realtime/transcription_sessions"
)
@ -141,6 +146,30 @@ async def test_client_secret_handler_still_targets_client_secrets_url():
assert kwargs["url"] == "https://api.openai.com/v1/realtime/client_secrets"
@pytest.mark.asyncio
async def test_azure_client_secret_prefers_provider_qualified_routing_model():
import litellm
mock_response = MagicMock(spec=httpx.Response)
mock_client = MagicMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=mock_response)
result = await litellm.acreate_realtime_client_secret(
model="azure/gpt-realtime-2.1",
session={"type": "realtime", "model": "gpt-realtime-2.1"},
api_base="https://my.openai.azure.com",
api_key="azure-test-key",
api_version="v1",
client=mock_client,
)
assert result is mock_response
request = mock_client.post.call_args.kwargs
assert request["url"] == "https://my.openai.azure.com/openai/v1/realtime/client_secrets"
assert request["headers"]["api-key"] == "azure-test-key"
assert request["json"]["session"]["model"] == "gpt-realtime-2.1"
@pytest.mark.asyncio
async def test_sdk_fn_routes_openai_transcription_session(monkeypatch):
"""
@ -169,18 +198,14 @@ async def test_sdk_fn_routes_openai_transcription_session(monkeypatch):
assert kwargs["url"].endswith("/v1/realtime/transcription_sessions")
# The litellm-only routing hint must not be forwarded upstream.
assert "model" not in kwargs["json"]
assert kwargs["json"]["input_audio_transcription"] == {
"model": "gpt-realtime-whisper"
}
assert kwargs["json"]["input_audio_transcription"] == {"model": "gpt-realtime-whisper"}
def test_append_query_params_skips_existing_keys():
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
url = "wss://example.com/v1/realtime?model=gpt-4o"
result = BaseLLMHTTPHandler._append_query_params(
url, {"model": "ignored", "intent": "transcription"}
)
result = BaseLLMHTTPHandler._append_query_params(url, {"model": "ignored", "intent": "transcription"})
assert "model=ignored" not in result
assert "intent=transcription" in result

View file

@ -0,0 +1,265 @@
import gzip
import json
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from openai import AsyncOpenAI
from litellm.llms.azure.realtime.http_transformation import AzureRealtimeHTTPConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig
from litellm.types.realtime import RealtimeSessionConfig
def test_realtime_session_config_supports_translation_and_live_transcription_fields():
session = RealtimeSessionConfig(
type="translation",
model="gpt-realtime-translate",
audio={
"input": {
"transcription": {
"model": "gpt-live-transcribe",
"delay": "minimal",
"languages": ["en", "fr"],
"keywords": ["LiteLLM"],
}
},
"output": {"language": "es"},
},
)
assert session.audio is not None
assert session.audio.input is not None
assert session.audio.input.transcription is not None
assert session.audio.input.transcription.delay == "minimal"
assert session.audio.input.transcription.languages == ["en", "fr"]
assert session.audio.input.transcription.keywords == ["LiteLLM"]
assert session.audio.output is not None
assert session.audio.output.language == "es"
@pytest.mark.parametrize(
"api_base,expected",
[
(
"https://api.openai.com",
"https://api.openai.com/v1/realtime/translations/client_secrets",
),
(
"https://api.openai.com/v1",
"https://api.openai.com/v1/realtime/translations/client_secrets",
),
],
)
def test_openai_translation_client_secret_url(api_base: str, expected: str):
config = OpenAIRealtimeHTTPConfig()
assert config.get_translation_client_secret_url(api_base, "gpt-realtime-translate") == expected
def test_openai_translation_calls_url():
config = OpenAIRealtimeHTTPConfig()
assert (
config.get_translation_calls_url("https://api.openai.com/v1", "gpt-realtime-translate")
== "https://api.openai.com/v1/realtime/translations/calls"
)
def test_azure_translation_urls_use_ga_paths():
config = AzureRealtimeHTTPConfig()
assert (
config.get_translation_client_secret_url("https://example.openai.azure.com", "translate-deployment")
== "https://example.openai.azure.com/openai/v1/realtime/translations/client_secrets"
)
assert (
config.get_translation_calls_url("https://example.openai.azure.com", "translate-deployment")
== "https://example.openai.azure.com/openai/v1/realtime/translations/calls"
)
@pytest.mark.asyncio
async def test_translation_client_secret_uses_custom_translation_path():
client = MagicMock(spec=AsyncHTTPHandler)
client.post = AsyncMock(
return_value=httpx.Response(
200,
json={"value": "ek_test"},
request=httpx.Request("POST", "https://api.openai.com/v1/realtime/translations/client_secrets"),
)
)
logging_obj = MagicMock()
handler = BaseLLMHTTPHandler()
request_data = {"session": {"type": "translation", "model": "gpt-realtime-translate"}}
response = await handler.async_realtime_translation_client_secret_handler(
api_base="https://api.openai.com",
api_key="sk-test",
request_data=request_data,
logging_obj=logging_obj,
timeout=10,
provider_config=OpenAIRealtimeHTTPConfig(),
model="gpt-realtime-translate",
client=client,
)
assert response.status_code == 200
call = client.post.call_args.kwargs
assert call["url"] == "https://api.openai.com/v1/realtime/translations/client_secrets"
assert call["json"] == request_data
@pytest.mark.asyncio
async def test_azure_translation_client_secret_supports_entra_bearer_auth():
client = MagicMock(spec=AsyncHTTPHandler)
client.post = AsyncMock(
return_value=httpx.Response(
200,
json={"value": "ek_test"},
request=httpx.Request(
"POST",
"https://example.openai.azure.com/openai/v1/realtime/translations/client_secrets",
),
)
)
handler = BaseLLMHTTPHandler()
await handler.async_realtime_translation_client_secret_handler(
api_base="https://example.openai.azure.com",
api_key="",
request_data={"session": {"type": "translation", "model": "translate-deployment"}},
logging_obj=MagicMock(),
timeout=10,
provider_config=AzureRealtimeHTTPConfig(),
model="translate-deployment",
extra_headers={"Authorization": "Bearer entra-token"},
client=client,
)
headers = client.post.call_args.kwargs["headers"]
assert headers["Authorization"] == "Bearer entra-token"
assert "api-key" not in headers
@pytest.mark.asyncio
async def test_translation_calls_use_translation_session_and_path():
client = MagicMock(spec=AsyncHTTPHandler)
client.post = AsyncMock(
return_value=httpx.Response(
201,
content=b"v=0\r\n",
request=httpx.Request("POST", "https://api.openai.com/v1/realtime/translations/calls"),
)
)
logging_obj = MagicMock()
handler = BaseLLMHTTPHandler()
response = await handler.async_realtime_calls_handler(
api_base="https://api.openai.com",
openai_ephemeral_key="ek_test",
sdp_body=b"v=0\r\n",
logging_obj=logging_obj,
timeout=10,
provider_config=OpenAIRealtimeHTTPConfig(),
model="gpt-realtime-translate",
client=client,
translation=True,
)
assert response.status_code == 201
call = client.post.call_args.kwargs
assert call["url"] == "https://api.openai.com/v1/realtime/translations/calls"
assert call["headers"]["Content-Type"] == "application/sdp"
assert call["content"] == "v=0\r\n"
@pytest.mark.asyncio
async def test_standard_client_secret_uses_openai_sdk_resource():
async def send_response(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/v1/realtime/client_secrets"
return httpx.Response(
200,
content=gzip.compress(b'{"value":"ek_test"}'),
headers={"content-encoding": "gzip"},
)
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client)
handler = BaseLLMHTTPHandler()
logging_obj = MagicMock()
response = await handler.async_realtime_client_secret_handler(
api_base="https://example.com",
api_key="sk-test",
request_data={"session": {"type": "realtime", "model": "gpt-realtime-2.1"}},
logging_obj=logging_obj,
timeout=10,
client=openai_client,
use_openai_sdk=True,
)
await openai_client.close()
assert response.status_code == 200
assert response.json() == {"value": "ek_test"}
assert "content-encoding" not in response.headers
@pytest.mark.asyncio
async def test_translation_client_secret_uses_openai_sdk_custom_post():
async def send_response(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/v1/realtime/translations/client_secrets"
assert json.loads(request.content)["session"]["audio"]["output"]["language"] == "es"
return httpx.Response(200, json={"value": "ek_translation"})
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client)
handler = BaseLLMHTTPHandler()
response = await handler.async_realtime_translation_client_secret_handler(
api_base="https://example.com",
api_key="sk-test",
request_data={
"session": {
"model": "gpt-realtime-translate",
"audio": {"output": {"language": "es"}},
}
},
logging_obj=MagicMock(),
timeout=10,
client=openai_client,
use_openai_sdk=True,
)
await openai_client.close()
assert response.status_code == 200
assert response.json() == {"value": "ek_translation"}
@pytest.mark.asyncio
async def test_translation_calls_use_openai_sdk_custom_post():
async def send_response(request: httpx.Request) -> httpx.Response:
assert request.url.path == "/v1/realtime/translations/calls"
body = await request.aread()
assert request.headers["content-type"] == "application/sdp"
assert body == b"v=0\r\n"
return httpx.Response(201, content=b"v=0\r\n")
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
openai_client = AsyncOpenAI(api_key="ek_test", base_url="https://example.com/v1", http_client=http_client)
handler = BaseLLMHTTPHandler()
response = await handler.async_realtime_calls_handler(
api_base="https://example.com",
openai_ephemeral_key="ek_test",
sdp_body=b"v=0\r\n",
logging_obj=MagicMock(),
timeout=10,
model="gpt-realtime-translate",
client=openai_client,
translation=True,
use_openai_sdk=True,
)
await openai_client.close()
assert response.status_code == 201
assert response.text == "v=0\r\n"

View file

@ -693,6 +693,37 @@ def test_realtime_webrtc_http_routes_classified_as_llm_api(route):
assert RouteChecks.is_management_route(route=route) is False
@pytest.mark.parametrize(
"route",
[
"/realtime/translations",
"/v1/realtime/translations",
"/openai/v1/realtime/translations",
"/realtime/translations/client_secrets",
"/v1/realtime/translations/client_secrets",
"/openai/v1/realtime/translations/client_secrets",
"/realtime/translations/calls",
"/v1/realtime/translations/calls",
"/openai/v1/realtime/translations/calls",
],
)
def test_realtime_translation_routes_allowed_for_openai_virtual_keys(route):
valid_token = UserAPIKeyAuth(
user_id="test_user",
allowed_routes=["openai_routes"],
)
assert RouteChecks.is_llm_api_route(route=route) is True
assert RouteChecks.is_management_route(route=route) is False
assert (
RouteChecks.is_virtual_key_allowed_to_call_route(
route=route,
valid_token=valid_token,
)
is True
)
def test_virtual_key_allowed_routes_with_litellm_routes_member_name_denied():
"""Test that virtual key is denied when route is not in the allowed LiteLLMRoutes group"""

View file

@ -119,9 +119,7 @@ def patched_transcription(monkeypatch):
return data
monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", _add_data)
monkeypatch.setattr(
proxy_server, "check_file_size_under_limit", lambda **kwargs: True
)
monkeypatch.setattr(proxy_server, "check_file_size_under_limit", lambda **kwargs: True)
async def _form_data(request):
from starlette.datastructures import FormData, UploadFile
@ -153,6 +151,47 @@ def patched_transcription_error(monkeypatch, patched_transcription):
yield
@pytest.fixture
def patched_transcription_stream(monkeypatch, patched_transcription):
class _FakeEvent:
def model_dump_json(self):
return '{"type":"transcript.text.done","text":"hello world"}'
class _FakeAsyncStream:
def __init__(self):
self.closed = False
def __aiter__(self):
async def _events():
yield _FakeEvent()
return _events()
async def aclose(self):
self.closed = True
async def _form_data(request):
from starlette.datastructures import FormData, UploadFile
upload = UploadFile(
filename="audio.mp3",
file=io.BytesIO(b"\x00\x01\x02"),
)
return FormData([("file", upload), ("model", "gpt-transcribe"), ("stream", "true")])
stream = _FakeAsyncStream()
async def _llm_call():
return stream
async def _fake_route_request(*args, **kwargs):
return _llm_call()
monkeypatch.setattr(proxy_server, "get_form_data", _form_data)
monkeypatch.setattr(proxy_server, "route_request", _fake_route_request)
yield stream
@pytest.mark.parametrize("path", ["/v1/audio/speech", "/audio/speech"])
def test_audio_speech_happy_path(client, auth_as, patched_speech, path):
"""Pins ``POST /v1/audio/speech`` and ``POST /audio/speech`` (happy)."""
@ -254,3 +293,15 @@ def test_audio_transcription_error(client, auth_as, patched_transcription_error,
response = client.post(path, files=files, data=data)
assert response.status_code == 500
assert len(response.content) > 0
@pytest.mark.parametrize("path", ["/v1/audio/transcriptions", "/audio/transcriptions"])
def test_audio_transcription_stream_returns_sse(client, auth_as, patched_transcription_stream, path):
files = {"file": ("audio.mp3", b"\x00\x01\x02", "audio/mpeg")}
data = {"model": "gpt-transcribe", "stream": "true"}
with auth_as():
response = client.post(path, files=files, data=data)
assert response.status_code == 200
assert response.headers["content-type"].startswith("text/event-stream")
assert response.text == 'data: {"type":"transcript.text.done","text":"hello world"}\n\n'
assert patched_transcription_stream.closed is True

View file

@ -15,16 +15,25 @@ import pytest
from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import ConfigGeneralSettings, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.realtime_endpoints.endpoints import (
_ALLOWED_SESSION_TYPES,
_coerce_realtime_session_type,
_decode_realtime_token_payload,
_encode_realtime_token_payload,
_prepare_client_secret_session,
)
from litellm.types.realtime import (
RealtimeAudioInputConfig,
RealtimeAudioTranscriptionConfig,
RealtimeClientSecretRequest,
RealtimeSessionAudioConfig,
RealtimeSessionConfig,
)
# --- Unit tests: token encode/decode helpers ---
@ -117,18 +126,107 @@ def proxy_app(monkeypatch):
from litellm.proxy import proxy_server
monkeypatch.setattr(proxy_server, "master_key", "sk-test-master-key")
monkeypatch.setitem(proxy_server.general_settings, "allow_non_billable_realtime_protocols", True)
return proxy_server.app
def test_non_billable_realtime_protocols_default_to_disabled():
assert ConfigGeneralSettings().allow_non_billable_realtime_protocols is False
@pytest.mark.parametrize(
("path", "body"),
(
(
"/v1/realtime/client_secrets",
{"model": "gpt-realtime-2"},
),
(
"/v1/realtime/translations/client_secrets",
{
"model": "gpt-realtime-translate",
"session": {"type": "translation", "model": "gpt-realtime-translate"},
},
),
(
"/v1/realtime/transcription_sessions",
{"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
),
),
)
def test_non_billable_realtime_credential_endpoints_require_opt_in(
proxy_app,
monkeypatch,
path,
body,
):
from litellm.proxy import proxy_server
monkeypatch.setitem(proxy_server.general_settings, "allow_non_billable_realtime_protocols", False)
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="test-user")
try:
client = TestClient(proxy_app, raise_server_exceptions=False)
with patch( # test-quality-ok: endpoint gate must prove routing is never reached
"litellm.proxy.proxy_server.route_request"
) as mock_route_request:
response = client.post(
path,
headers={"Authorization": "Bearer sk-test-master-key"},
json=body,
)
assert response.status_code == 403
assert "bypasses LiteLLM billing" in response.json()["detail"]
mock_route_request.assert_not_called()
finally:
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
@pytest.mark.parametrize(
"path",
(
"/v1/realtime/calls",
"/v1/realtime/translations/calls",
),
)
def test_non_billable_realtime_sdp_endpoints_require_opt_in(
proxy_app,
monkeypatch,
path,
):
from litellm.proxy import proxy_server
monkeypatch.setitem(proxy_server.general_settings, "allow_non_billable_realtime_protocols", False)
token_payload = _encode_realtime_token_payload(
ephemeral_key="epk",
model_id="gpt-realtime-2",
user_id=None,
team_id=None,
expires_at=int(time.time()) + 3600,
)
encrypted_token = encrypt_value_helper(token_payload)
client = TestClient(proxy_app, raise_server_exceptions=False)
with patch( # test-quality-ok: endpoint gate must prove routing is never reached
"litellm.proxy.proxy_server.route_request"
) as mock_route_request:
response = client.post(
path,
headers={"Authorization": f"Bearer {encrypted_token}"},
content=b"v=0\r\n",
)
assert response.status_code == 403
assert "bypasses LiteLLM billing" in response.json()["detail"]
mock_route_request.assert_not_called()
@pytest.fixture
def mock_route_request_client_secrets():
"""Mock route_request to return a fake upstream client_secrets response."""
future_expires_at = int(time.time()) + 3600
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.text = (
f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}'
)
mock_resp.text = f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}'
mock_resp.content = f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}'.encode()
mock_resp.headers = {}
mock_resp.json.return_value = {
@ -215,9 +313,7 @@ async def test_client_secrets_success_with_mock(
mock_pre_call_hook,
):
"""POST /v1/realtime/client_secrets returns 200 with valid auth and mocked upstream."""
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user", team_id="test-team"
)
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="test-user", team_id="test-team")
try:
client = TestClient(proxy_app)
with (
@ -275,13 +371,7 @@ async def test_client_secrets_transcription_rejects_disallowed_nested_model(
"session": {
"type": "transcription",
"model": "gpt-4o-realtime-preview",
"audio": {
"input": {
"transcription": {
"model": "gpt-realtime-whisper"
}
}
},
"audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}},
},
},
)
@ -312,12 +402,8 @@ async def test_client_secrets_transcription_routes_on_nested_model(
async def _inner():
resp = MagicMock(spec=httpx.Response)
resp.status_code = 200
resp.text = (
f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}'
)
resp.content = (
f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}'
).encode()
resp.text = f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}'
resp.content = (f'{{"value":"upstream_ephemeral_key","expires_at":{future_expires_at}}}').encode()
resp.headers = {}
resp.json.return_value = {
"value": "upstream_ephemeral_key",
@ -351,13 +437,7 @@ async def test_client_secrets_transcription_routes_on_nested_model(
"session": {
"type": "transcription",
"model": "gpt-4o-realtime-preview",
"audio": {
"input": {
"transcription": {
"model": "gpt-realtime-whisper"
}
}
},
"audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}},
},
},
)
@ -367,10 +447,7 @@ async def test_client_secrets_transcription_routes_on_nested_model(
session = captured["data"]["session"]
assert session["type"] == "transcription"
assert "model" not in session
assert (
session["audio"]["input"]["transcription"]["model"]
== "gpt-realtime-whisper"
)
assert session["audio"]["input"]["transcription"]["model"] == "gpt-realtime-whisper"
encrypted_value = response.json()["value"]
decoded = _decode_realtime_token_payload(
decrypt_value_helper(
@ -530,10 +607,7 @@ async def test_realtime_calls_replays_transcription_session_type(
)
assert captured["session"]["type"] == "transcription"
assert (
captured["session"]["audio"]["input"]["transcription"]["model"]
== "gpt-realtime-whisper"
)
assert captured["session"]["audio"]["input"]["transcription"]["model"] == "gpt-realtime-whisper"
# --- transcription_sessions endpoint ---
@ -605,9 +679,7 @@ async def test_transcription_sessions_rejects_disallowed_resolved_model(
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
json={
"input_audio_transcription": {"model": "gpt-realtime-whisper"}
},
json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
)
assert response.status_code == 403
@ -651,9 +723,7 @@ async def test_transcription_sessions_rejects_disallowed_team_model_scope(
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
json={
"input_audio_transcription": {"model": "gpt-realtime-whisper"}
},
json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
)
assert response.status_code == 403
@ -696,9 +766,7 @@ async def test_transcription_sessions_rejects_disallowed_project_model_scope(
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
json={
"input_audio_transcription": {"model": "gpt-realtime-whisper"}
},
json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
)
assert response.status_code == 403
@ -751,9 +819,7 @@ async def test_transcription_sessions_rejects_disallowed_team_member_model_scope
response = client.post(
"/v1/realtime/transcription_sessions",
headers={"Authorization": "Bearer sk-test-master-key"},
json={
"input_audio_transcription": {"model": "gpt-realtime-whisper"}
},
json={"input_audio_transcription": {"model": "gpt-realtime-whisper"}},
)
assert response.status_code == 403
@ -786,6 +852,20 @@ async def test_realtime_transcription_websocket_default_model_checks_key_scope()
assert "is not available for this API key" in close_kwargs["reason"]
def test_realtime_transcription_upstream_query_omits_model():
from litellm.proxy import proxy_server
assert (
proxy_server._resolve_realtime_upstream_query_model(
model="gpt-live-transcribe",
intent="transcription",
is_translation=False,
route_model="gpt-live-transcribe",
)
is None
)
@pytest.mark.asyncio
async def test_realtime_transcription_websocket_default_model_checks_team_scope():
from litellm.proxy import proxy_server
@ -947,9 +1027,7 @@ async def test_transcription_sessions_encrypts_client_secret(
POST /v1/realtime/transcription_sessions returns 200 and the ephemeral key
under client_secret.value must be encrypted (never the raw upstream key).
"""
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user", team_id="test-team"
)
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="test-user", team_id="test-team")
captured_route_type = {}
async def _capturing_route(*args, **kwargs):
@ -993,16 +1071,12 @@ async def test_transcription_sessions_encrypts_client_secret(
assert decrypted is not None
assert "upstream_ephemeral_key" in decrypted
# Routed through the dedicated transcription_sessions route type.
assert (
captured_route_type["route_type"]
== "acreate_realtime_transcription_session"
)
assert captured_route_type["route_type"] == "acreate_realtime_transcription_session"
finally:
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
def test_session_type_coerced_for_unknown_value():
"""An unrecognized session_type in the token falls back to 'realtime'."""
payload = _encode_realtime_token_payload(
ephemeral_key="epk",
model_id="gpt-4o",
@ -1011,12 +1085,222 @@ def test_session_type_coerced_for_unknown_value():
expires_at=None,
session_type="INJECTED_TYPE",
)
# Force-deserialize and check the coercion that happens in proxy_realtime_calls.
decoded = json.loads(payload)
session_type = decoded.get("session_type") or "realtime"
if session_type not in ("realtime", "transcription"):
session_type = "realtime"
assert session_type == "realtime"
assert decoded["session_type"] == "INJECTED_TYPE"
assert _coerce_realtime_session_type("INJECTED_TYPE") == "realtime"
assert _coerce_realtime_session_type(None) == "realtime"
for allowed_session_type in _ALLOWED_SESSION_TYPES:
assert _coerce_realtime_session_type(allowed_session_type) == allowed_session_type
@pytest.mark.asyncio
async def test_translation_client_secret_rejects_disallowed_nested_transcription_model() -> None:
req = RealtimeClientSecretRequest(
model="gpt-realtime-translate",
session=RealtimeSessionConfig(
type="translation",
model="gpt-realtime-translate",
audio=RealtimeSessionAudioConfig(
input=RealtimeAudioInputConfig(
transcription=RealtimeAudioTranscriptionConfig(model="gpt-live-transcribe"),
)
),
),
)
with pytest.raises(Exception, match="Tried to access gpt-live-transcribe"):
await _prepare_client_secret_session(
req=req,
user_api_key_dict=UserAPIKeyAuth(models=["gpt-realtime-translate"]),
llm_model_list=None,
llm_router=None,
forced_session_type="translation",
)
@pytest.mark.asyncio
async def test_translation_client_secret_binds_authorized_nested_transcription_model() -> None:
req = RealtimeClientSecretRequest(
model="gpt-realtime-translate",
session=RealtimeSessionConfig(
type="translation",
model="gpt-realtime-translate",
audio=RealtimeSessionAudioConfig(
input=RealtimeAudioInputConfig(
transcription=RealtimeAudioTranscriptionConfig(model="gpt-realtime-whisper"),
)
),
),
)
model, session_data, session_type = await _prepare_client_secret_session(
req=req,
user_api_key_dict=UserAPIKeyAuth(models=["gpt-realtime-translate", "gpt-realtime-whisper"]),
llm_model_list=None,
llm_router=None,
forced_session_type="translation",
)
assert model == "gpt-realtime-translate"
assert session_type == "translation"
assert session_data is not None
assert session_data["model"] == "gpt-realtime-translate"
assert session_data["audio"]["input"]["transcription"]["model"] == "gpt-realtime-whisper"
@pytest.mark.parametrize(
"path",
[
"/v1/realtime/translations/client_secrets",
"/realtime/translations/client_secrets",
"/openai/v1/realtime/translations/client_secrets",
],
)
def test_translation_client_secret_aliases_bind_token_family(
proxy_app,
mock_route_request_client_secrets,
mock_add_litellm_data,
mock_pre_call_hook,
path,
):
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user",
models=["gpt-realtime-translate"],
)
captured = {}
async def capture_route(*args, **kwargs):
captured["route_type"] = kwargs["route_type"]
captured["data"] = kwargs["data"]
return await mock_route_request_client_secrets(*args, **kwargs)
try:
client = TestClient(proxy_app)
with (
patch( # test-quality-ok: endpoint test captures the proxy routing boundary
"litellm.proxy.proxy_server.route_request", side_effect=capture_route
),
patch( # test-quality-ok: endpoint test isolates request metadata enrichment
"litellm.proxy.proxy_server.add_litellm_data_to_request",
side_effect=mock_add_litellm_data,
),
patch( # test-quality-ok: endpoint test isolates the process-wide proxy logger
"litellm.proxy.proxy_server.proxy_logging_obj"
) as logging,
):
logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook)
logging.post_call_failure_hook = AsyncMock()
response = client.post(
path,
headers={"Authorization": "Bearer sk-test-master-key"},
json={"model": "gpt-realtime-translate"},
)
assert response.status_code == 200
assert captured["route_type"] == "acreate_realtime_translation_client_secret"
assert captured["data"]["session"] == {
"type": "translation",
"model": "gpt-realtime-translate",
}
decrypted = decrypt_value_helper(
response.json()["value"],
key="client_secret.value",
exception_type="debug",
)
decoded = _decode_realtime_token_payload(decrypted or "")
assert decoded is not None
assert decoded["session_type"] == "translation"
finally:
proxy_app.dependency_overrides.pop(user_api_key_auth, None)
@pytest.mark.parametrize(
"path",
[
"/v1/realtime/translations/calls",
"/realtime/translations/calls",
"/openai/v1/realtime/translations/calls",
],
)
def test_translation_calls_aliases_route_translation_session(
proxy_app,
mock_route_request_realtime_calls,
mock_add_litellm_data,
mock_pre_call_hook,
path,
):
token = encrypt_value_helper(
_encode_realtime_token_payload(
ephemeral_key="ek_test",
model_id="gpt-realtime-translate",
user_id="test-user",
team_id=None,
expires_at=int(time.time()) + 3600,
session_type="translation",
)
)
captured = {}
async def capture_route(*args, **kwargs):
captured["route_type"] = kwargs["route_type"]
captured["data"] = kwargs["data"]
return await mock_route_request_realtime_calls(*args, **kwargs)
client = TestClient(proxy_app)
with (
patch( # test-quality-ok: endpoint test captures the proxy routing boundary
"litellm.proxy.proxy_server.route_request", side_effect=capture_route
),
patch( # test-quality-ok: endpoint test isolates request metadata enrichment
"litellm.proxy.proxy_server.add_litellm_data_to_request",
side_effect=mock_add_litellm_data,
),
patch( # test-quality-ok: endpoint test isolates the process-wide proxy logger
"litellm.proxy.proxy_server.proxy_logging_obj"
) as logging,
):
logging.pre_call_hook = AsyncMock(side_effect=mock_pre_call_hook)
logging.post_call_failure_hook = AsyncMock()
response = client.post(
path,
headers={"Authorization": f"Bearer {token}"},
content=b"v=0\r\n",
)
assert response.status_code == 201
assert captured["route_type"] == "arealtime_translation_calls"
assert captured["data"]["session"] == {
"type": "translation",
"model": "gpt-realtime-translate",
}
@pytest.mark.parametrize(
"session_type,path",
[
("realtime", "/v1/realtime/translations/calls"),
("translation", "/v1/realtime/calls"),
],
)
def test_realtime_calls_reject_cross_family_token(proxy_app, session_type, path):
model = "gpt-realtime-translate" if session_type == "translation" else "gpt-realtime-2"
token = encrypt_value_helper(
_encode_realtime_token_payload(
ephemeral_key="ek_test",
model_id=model,
user_id=None,
team_id=None,
expires_at=int(time.time()) + 3600,
session_type=session_type,
)
)
response = TestClient(proxy_app).post(
path,
headers={"Authorization": f"Bearer {token}"},
content=b"v=0\r\n",
)
assert response.status_code == 401
assert response.json()["error"] == "Token is not valid for this Realtime endpoint"
@pytest.mark.asyncio
@ -1142,9 +1426,7 @@ async def test_transcription_sessions_returns_upstream_error_verbatim(
return _inner()
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user", team_id="test-team"
)
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="test-user", team_id="test-team")
try:
client = TestClient(proxy_app)
with (
@ -1184,9 +1466,7 @@ async def test_transcription_sessions_wraps_route_exception(
async def _raise_http(*args, **kwargs):
raise HTTPException(status_code=403, detail="Model not allowed")
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_id="test-user"
)
proxy_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="test-user")
try:
client = TestClient(proxy_app, raise_server_exceptions=False)
with (

View file

@ -25,6 +25,7 @@ from fastapi import FastAPI, HTTPException, Request
from fastapi.encoders import jsonable_encoder
from fastapi.staticfiles import StaticFiles
from fastapi.testclient import TestClient
from starlette.datastructures import URL
import litellm
import litellm.proxy.proxy_server as proxy_server_module
@ -10757,7 +10758,7 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory():
def test_realtime_websocket_route_aliases_registered():
"""Realtime sessions reach the proxy via three path aliases stacked on
"""Realtime sessions reach the proxy via six path aliases stacked on
`realtime_websocket_endpoint`. Dropping any of them silently 405s
WebSocket upgrades because the catch-all `/openai/{endpoint:path}`
HTTP passthrough only declares HTTP methods. The aliases must also be
@ -10773,7 +10774,14 @@ def test_realtime_websocket_route_aliases_registered():
websocket_paths = {route.path for route in app.routes if isinstance(route, WebSocketRoute)}
openai_routes = LiteLLMRoutes.openai_routes.value
for expected in ("/openai/v1/realtime", "/v1/realtime", "/realtime"):
for expected in (
"/openai/v1/realtime",
"/v1/realtime",
"/realtime",
"/openai/v1/realtime/translations",
"/v1/realtime/translations",
"/realtime/translations",
):
assert expected in websocket_paths, (
f"{expected!r} missing from registered WebSocket routes; the "
f"realtime endpoint will 405 for clients hitting this path."
@ -10792,7 +10800,7 @@ def _lit6973_fake_realtime_ws() -> MagicMock:
ws = MagicMock()
ws.headers = {}
ws.scope = {"headers": [], "type": "websocket"}
ws.url = "ws://testserver/v1/realtime"
ws.url = URL("ws://testserver/v1/realtime")
ws.accept = AsyncMock()
ws.send_text = AsyncMock()
ws.close = AsyncMock()

View file

@ -538,6 +538,8 @@ def test_every_openai_sdk_response_error_code_has_explicit_status_mapping():
("insufficient_quota", 429),
("vector_store_timeout", 504),
("invalid_prompt", 400),
("data_residency_mismatch", 400),
("bio_policy", 400),
("invalid_image", 400),
("invalid_image_format", 400),
("invalid_base64_image", 400),

View file

@ -979,6 +979,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"/v1/images/generations",
"/v1/realtime",
"/v1/realtime/transcription_sessions",
"/v1/realtime/translations",
"/v1/realtime/translations/client_secrets",
"/v1/realtime/translations/calls",
"/v1/images/variations",
"/v1/images/edits",
"/v1/batch",

View file

@ -0,0 +1,297 @@
import io
import json
import wave
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from openai import AsyncOpenAI, AsyncStream, AzureOpenAI
import litellm
from litellm.litellm_core_utils.audio_utils.transcription_streaming import wrap_transcription_stream
from litellm.llms.azure.audio_transcriptions import AzureAudioTranscription
from litellm.llms.openai.transcriptions.gpt_transformation import (
OpenAIGPTTranscribeAudioTranscriptionConfig,
)
from litellm.llms.openai.transcriptions.handler import OpenAIAudioTranscription
from litellm.main import _validate_gpt_transcription_request
from litellm.types.utils import TranscriptionResponse
from litellm.utils import get_optional_params_transcription
def test_gpt_transcribe_config_uses_native_parameters_and_json():
config = OpenAIGPTTranscribeAudioTranscriptionConfig()
supported = config.get_supported_openai_params("gpt-transcribe")
assert supported == ["prompt", "response_format", "keywords", "languages", "stream"]
audio_file = io.BytesIO(b"audio")
request = config.transform_audio_transcription_request(
model="gpt-transcribe",
audio_file=audio_file,
optional_params={"keywords": ["LiteLLM"], "languages": ["en", "fr"], "stream": True},
litellm_params={},
)
assert request.data["response_format"] == "json"
assert request.data["keywords"] == ["LiteLLM"]
assert request.data["languages"] == ["en", "fr"]
assert request.data["stream"] is True
def test_gpt_transcribe_optional_params_are_preserved():
params = get_optional_params_transcription(
model="gpt-transcribe",
custom_llm_provider="openai",
keywords=["LiteLLM", "Realtime API"],
languages=["en", "fr"],
stream=True,
)
assert params == {
"keywords": ["LiteLLM", "Realtime API"],
"languages": ["en", "fr"],
"stream": True,
}
def test_transcription_response_preserves_empty_languages():
response = TranscriptionResponse(text="hello", languages=[])
assert response.model_dump()["languages"] == []
@pytest.mark.asyncio
async def test_openai_handler_returns_native_typed_stream():
async def send_response(request: httpx.Request) -> httpx.Response:
body = await request.aread()
assert b'name="keywords[]"' in body
assert b'name="languages[]"' in body
assert b'name="stream"' in body
events = (
{"type": "transcript.text.delta", "delta": "hello "},
{
"type": "transcript.text.done",
"text": "hello world",
"languages": [],
"usage": {
"type": "tokens",
"input_tokens": 10,
"output_tokens": 2,
"total_tokens": 12,
},
},
)
content = "".join(f"data: {json.dumps(event)}\n\n" for event in events).encode()
return httpx.Response(200, content=content, headers={"content-type": "text/event-stream"})
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client)
logging_obj = MagicMock()
logging_obj.model_call_details = {}
audio_file = io.BytesIO(b"audio")
audio_file.name = "sample.wav"
handler = OpenAIAudioTranscription()
result = handler.audio_transcriptions(
model="gpt-transcribe",
audio_file=audio_file,
optional_params={"keywords": ["LiteLLM"], "languages": ["en"], "stream": True},
litellm_params={},
model_response=TranscriptionResponse(),
timeout=10,
max_retries=0,
logging_obj=logging_obj,
api_key="sk-test",
api_base="https://example.com/v1",
client=openai_client,
atranscription=True,
provider_config=OpenAIGPTTranscribeAudioTranscriptionConfig(),
)
stream = await result
assert isinstance(stream, AsyncStream)
logging_obj.async_success_handler = AsyncMock()
logging_obj.async_failure_handler = AsyncMock()
wrapped_stream = wrap_transcription_stream(stream, logging_obj, datetime.now())
received = [event async for event in wrapped_stream]
await wrapped_stream.close()
await openai_client.close()
assert [event.type for event in received] == ["transcript.text.delta", "transcript.text.done"]
assert received[-1].languages == []
logging_obj.async_success_handler.assert_awaited_once()
logged_response = logging_obj.async_success_handler.await_args.kwargs["result"]
assert logged_response.text == "hello world"
assert logged_response.languages == []
@pytest.mark.asyncio
async def test_atranscription_stream_preserves_duration_for_callback_cost():
async def send_response(request: httpx.Request) -> httpx.Response:
events = (
{"type": "transcript.text.delta", "delta": "hello "},
{
"type": "transcript.text.done",
"text": "hello world",
"usage": {
"type": "tokens",
"input_tokens": 10,
"output_tokens": 2,
"total_tokens": 12,
},
},
)
content = "".join(f"data: {json.dumps(event)}\n\n" for event in events).encode()
return httpx.Response(200, content=content, headers={"content-type": "text/event-stream"})
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client)
logging_obj = MagicMock()
logging_obj.model_call_details = {}
logging_obj.async_success_handler = AsyncMock()
logging_obj.async_failure_handler = AsyncMock()
audio_file = io.BytesIO()
with wave.open(audio_file, "wb") as wav_file:
wav_file.setnchannels(1)
wav_file.setsampwidth(2)
wav_file.setframerate(16000)
wav_file.writeframes(b"\x00\x00" * 16000)
audio_file.name = "sample.wav"
stream = await litellm.atranscription(
model="openai/gpt-transcribe",
file=audio_file,
stream=True,
client=openai_client,
litellm_logging_obj=logging_obj,
)
received = [event async for event in stream]
await stream.close()
await openai_client.close()
assert [event.type for event in received] == ["transcript.text.delta", "transcript.text.done"]
logging_obj.async_success_handler.assert_awaited_once()
logged_response = logging_obj.async_success_handler.await_args.kwargs["result"]
assert logged_response._hidden_params["audio_transcription_duration"] == pytest.approx(1.0)
def test_gpt_transcribe_rejects_conflicting_language_inputs():
audio_file = io.BytesIO(b"audio")
audio_file.name = "sample.wav"
with pytest.raises(litellm.UnsupportedParamsError, match="cannot be used together"):
litellm.transcription(
model="gpt-transcribe",
file=audio_file,
language="en",
languages=["fr"],
api_key="sk-test",
)
def test_gpt_transcribe_rejects_whisper_response_formats():
audio_file = io.BytesIO(b"audio")
audio_file.name = "sample.wav"
with pytest.raises(litellm.UnsupportedParamsError, match="only supports response_format='json'"):
litellm.transcription(
model="gpt-transcribe",
file=audio_file,
response_format="verbose_json",
api_key="sk-test",
)
def test_gpt_live_transcribe_rejects_file_transcription():
audio_file = io.BytesIO(b"audio")
audio_file.name = "sample.wav"
with pytest.raises(litellm.UnsupportedParamsError, match="Realtime API"):
litellm.transcription(
model="gpt-live-transcribe",
file=audio_file,
api_key="sk-test",
)
def test_azure_async_gpt_transcribe_forwards_v1_api_version():
handler = AzureAudioTranscription()
handler.async_audio_transcriptions = MagicMock(return_value=MagicMock())
handler.audio_transcriptions(
model="gpt-transcribe",
audio_file=io.BytesIO(b"audio"),
optional_params={"stream": True},
logging_obj=MagicMock(),
model_response=TranscriptionResponse(),
timeout=10,
max_retries=0,
api_key="sk-test",
api_base="https://example.openai.azure.com",
api_version="v1",
atranscription=True,
)
assert handler.async_audio_transcriptions.call_args.kwargs["api_version"] == "v1"
@pytest.mark.parametrize("api_version", [None, "v1", "latest", "preview"])
def test_azure_gpt_transcribe_uses_deployment_scoped_api_version(api_version: str | None):
resolved_api_version = _validate_gpt_transcription_request(
model="gpt-transcribe",
custom_llm_provider="azure",
language=None,
languages=None,
response_format="json",
api_version=api_version,
)
assert resolved_api_version == litellm.AZURE_DEFAULT_API_VERSION
def test_azure_gpt_transcribe_uses_deployment_scoped_route():
def send_response(request: httpx.Request) -> httpx.Response:
assert str(request.url) == (
"https://example.openai.azure.com/openai/deployments/gpt-transcribe/audio/transcriptions"
f"?api-version={litellm.AZURE_DEFAULT_API_VERSION}"
)
return httpx.Response(
200,
json={"text": "hello", "languages": [{"code": "en"}], "usage": {"type": "duration", "seconds": 1}},
)
http_client = httpx.Client(transport=httpx.MockTransport(send_response))
client = AzureOpenAI(
api_key="azure-test-key",
azure_endpoint="https://example.openai.azure.com",
api_version=litellm.AZURE_DEFAULT_API_VERSION,
http_client=http_client,
)
audio_file = io.BytesIO(b"audio")
audio_file.name = "sample.wav"
response = AzureAudioTranscription().audio_transcriptions(
model="gpt-transcribe",
audio_file=audio_file,
optional_params={"response_format": "json"},
logging_obj=MagicMock(),
model_response=TranscriptionResponse(),
timeout=10,
max_retries=0,
api_key="azure-test-key",
api_base="https://example.openai.azure.com",
api_version=litellm.AZURE_DEFAULT_API_VERSION,
client=client,
)
assert response.text == "hello"
assert response.languages is not None
assert [language.code for language in response.languages] == ["en"]
client.close()
def test_azure_gpt_transcribe_preserves_dated_api_version():
resolved_api_version = _validate_gpt_transcription_request(
model="gpt-transcribe",
custom_llm_provider="azure",
language=None,
languages=None,
response_format="json",
api_version="2025-04-01-preview",
)
assert resolved_api_version == "2025-04-01-preview"

View file

@ -9,6 +9,9 @@ TranscriptionVerbose/Diarized type.
from unittest.mock import patch
import pytest
import litellm
from litellm.cost_calculator import completion_cost
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
convert_to_model_response_object,
@ -55,10 +58,7 @@ class TestDiarizedJsonUsageParsing:
assert result.usage.seconds == 295.8
def test_usage_duration_object_accepts_float_seconds(self):
assert (
TranscriptionUsageDurationObject(type="duration", seconds=295.8).seconds
== 295.8
)
assert TranscriptionUsageDurationObject(type="duration", seconds=295.8).seconds == 295.8
class TestTranscriptionDurationNotInResponseBody:
@ -126,6 +126,60 @@ class TestTranscriptionDurationNotInResponseBody:
class TestCostCalculatorReadsDurationFromHiddenParams:
"""The cost calculator should read duration from _hidden_params via completion_cost()."""
@patch( # test-quality-ok: isolates duration selection from the model pricing table
"litellm.cost_calculator.openai_cost_per_second"
)
def test_completion_cost_prefers_provider_usage_duration(self, mock_cost_fn):
mock_cost_fn.return_value = (0.001, 0.0)
response = TranscriptionResponse(
text="test",
usage={"type": "duration", "seconds": 60.0},
)
response._hidden_params = {
"audio_transcription_duration": 17.5,
"custom_llm_provider": "openai",
}
completion_cost(
completion_response=response,
model="gpt-transcribe",
call_type="atranscription",
)
mock_cost_fn.assert_called_once()
_, kwargs = mock_cost_fn.call_args
assert kwargs["duration"] == 60.0
@pytest.mark.parametrize(
("model", "provider"),
[
("gpt-transcribe", "openai"),
("azure/gpt-transcribe", "azure"),
],
)
def test_gpt_transcribe_provider_usage_duration_costs_without_local_decode(
self,
monkeypatch,
model: str,
provider: str,
):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
response = TranscriptionResponse(
text="test",
usage={"type": "duration", "seconds": 60.0},
)
response._hidden_params = {"custom_llm_provider": provider}
actual = completion_cost(
completion_response=response,
model=model,
call_type="atranscription",
custom_llm_provider=provider,
)
assert actual == pytest.approx(0.0045)
@patch("litellm.cost_calculator.openai_cost_per_second")
def test_completion_cost_uses_hidden_params_duration(self, mock_cost_fn):
"""