From bd43c233ef9a4b402793c912e85b543119f0dcb7 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Tue, 15 Sep 2026 12:12:59 -0500 Subject: [PATCH] feat(realtime): support latest OpenAI audio models --- cookbook/gpt_realtime_translate.py | 242 ++++++++++ litellm/__init__.py | 18 +- litellm/constants.py | 34 ++ litellm/cost_calculator.py | 76 +++- .../audio_utils/transcription_streaming.py | 204 +++++++++ .../litellm_core_utils/llm_cost_calc/utils.py | 79 +++- .../convert_dict_to_response.py | 15 +- .../litellm_core_utils/realtime_streaming.py | 237 ++++++++-- litellm/llms/azure/audio_transcriptions.py | 21 +- litellm/llms/azure/realtime/handler.py | 34 +- .../azure/realtime/http_transformation.py | 27 +- .../base_llm/realtime/http_transformation.py | 10 + litellm/llms/custom_httpx/llm_http_handler.py | 287 ++++++++++-- litellm/llms/openai/realtime/handler.py | 126 +++++- .../openai/realtime/http_transformation.py | 12 + .../transcriptions/gpt_transformation.py | 30 +- litellm/llms/openai/transcriptions/handler.py | 42 +- litellm/main.py | 95 +++- litellm/proxy/_lazy_features.py | 15 + litellm/proxy/_lazy_openapi_snapshot.json | 180 ++++++++ litellm/proxy/_types.py | 13 + litellm/proxy/common_request_processing.py | 2 + .../proxy/common_utils/http_parsing_utils.py | 20 +- litellm/proxy/proxy_server.py | 88 +++- litellm/proxy/realtime_endpoints/endpoints.py | 138 +++++- litellm/proxy/route_llm_request.py | 2 + litellm/realtime_api/README.md | 14 +- litellm/realtime_api/main.py | 291 ++++++++++-- litellm/responses/streaming_iterator.py | 2 + litellm/responses/utils.py | 8 +- litellm/router.py | 12 + litellm/types/llms/openai.py | 14 + litellm/types/realtime.py | 48 +- litellm/types/utils.py | 38 +- litellm/utils.py | 33 +- .../llm_nonconversational.yaml | 6 + tests/e2e/coverage_registry/schema.py | 11 +- .../test_realtime_streaming.py | 187 ++++++++ .../llms/azure/test_azure_common_utils.py | 5 + .../realtime/test_openai_realtime_handler.py | 324 +++++++------- .../realtime/test_transcription_sessions.py | 55 ++- .../llms/openai/realtime/test_translation.py | 265 +++++++++++ .../proxy/auth/test_route_checks.py | 31 ++ .../proxy/proxy_server/test_routes_audio.py | 57 ++- .../test_realtime_webrtc_endpoints.py | 414 +++++++++++++++--- tests/test_litellm/proxy/test_proxy_server.py | 14 +- .../test_streaming_iterator_error_events.py | 2 + tests/test_litellm/test_utils.py | 3 + .../transcriptions/test_gpt_transcribe.py | 297 +++++++++++++ .../test_transcription_duration_hidden.py | 62 ++- 50 files changed, 3744 insertions(+), 496 deletions(-) create mode 100644 cookbook/gpt_realtime_translate.py create mode 100644 litellm/litellm_core_utils/audio_utils/transcription_streaming.py create mode 100644 tests/test_litellm/llms/openai/realtime/test_translation.py create mode 100644 tests/unit/llms/openai/transcriptions/test_gpt_transcribe.py diff --git a/cookbook/gpt_realtime_translate.py b/cookbook/gpt_realtime_translate.py new file mode 100644 index 00000000000..73c91eaa1e9 --- /dev/null +++ b/cookbook/gpt_realtime_translate.py @@ -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()) diff --git a/litellm/__init__.py b/litellm/__init__.py index c8df4394a06..4ec7aeea8fb 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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, diff --git a/litellm/constants.py b/litellm/constants.py index e5b662bd515..855f51811a4 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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 = { diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 7990832dc48..2e4448475a7 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -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, diff --git a/litellm/litellm_core_utils/audio_utils/transcription_streaming.py b/litellm/litellm_core_utils/audio_utils/transcription_streaming.py new file mode 100644 index 00000000000..5ae246f3468 --- /dev/null +++ b/litellm/litellm_core_utils/audio_utils/transcription_streaming.py @@ -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) diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 46bf2ec2960..d5a27988762 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -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, diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 87524d86c61..c73b74a3c9a 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -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 diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index d2fbb26bb02..5e9e04ae54f 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -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) diff --git a/litellm/llms/azure/audio_transcriptions.py b/litellm/llms/azure/audio_transcriptions.py index 564ec94ba6b..fb575a3713c 100644 --- a/litellm/llms/azure/audio_transcriptions.py +++ b/litellm/llms/azure/audio_transcriptions.py @@ -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) diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index df74975ad0f..393499c0fcd 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -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() diff --git a/litellm/llms/azure/realtime/http_transformation.py b/litellm/llms/azure/realtime/http_transformation.py index 633c2b3b6d0..aa2b6d9cda0 100644 --- a/litellm/llms/azure/realtime/http_transformation.py +++ b/litellm/llms/azure/realtime/http_transformation.py @@ -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, diff --git a/litellm/llms/base_llm/realtime/http_transformation.py b/litellm/llms/base_llm/realtime/http_transformation.py index 43a80edb493..0d64ed3931d 100644 --- a/litellm/llms/base_llm/realtime/http_transformation.py +++ b/litellm/llms/base_llm/realtime/http_transformation.py @@ -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. diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 8cbc28362a8..052b9ee9e87 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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( diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index bdc3a6c7908..df2d9ca4c75 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -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() diff --git a/litellm/llms/openai/realtime/http_transformation.py b/litellm/llms/openai/realtime/http_transformation.py index 61dbf20397f..726a9545e11 100644 --- a/litellm/llms/openai/realtime/http_transformation.py +++ b/litellm/llms/openai/realtime/http_transformation.py @@ -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, diff --git a/litellm/llms/openai/transcriptions/gpt_transformation.py b/litellm/llms/openai/transcriptions/gpt_transformation.py index e7f8dd04516..81f3abd217a 100644 --- a/litellm/llms/openai/transcriptions/gpt_transformation.py +++ b/litellm/llms/openai/transcriptions/gpt_transformation.py @@ -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) diff --git a/litellm/llms/openai/transcriptions/handler.py b/litellm/llms/openai/transcriptions/handler.py index 701b3d30362..2c703a16d4e 100644 --- a/litellm/llms/openai/transcriptions/handler.py +++ b/litellm/llms/openai/transcriptions/handler.py @@ -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: diff --git a/litellm/main.py b/litellm/main.py index 7f4b34d28a0..376645f55f4 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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.") diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 5be87a8bf4d..ed09149904e 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -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", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 71dbf0da239..9a64be44330 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -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" + ] + } } } }, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 54574ed64e3..de0eb77d762 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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" ) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c40090233be..991da2b200f 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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", diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 1c2bd7ea217..37805ff0459 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -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( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 92a75bf953a..6b95865c8ff 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index be2ac2ff33e..bee2a55f5b5 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -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) diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 536c58df65a..399ecc19d27 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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", diff --git a/litellm/realtime_api/README.md b/litellm/realtime_api/README.md index d810de2f24f..d9699f57da7 100644 --- a/litellm/realtime_api/README.md +++ b/litellm/realtime_api/README.md @@ -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. \ No newline at end of file +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. diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index acc42c44c04..02055988565 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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), ) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 59655800af6..044e15339a5 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -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, diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index a2642795cea..d94b9940f24 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -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), diff --git a/litellm/router.py b/litellm/router.py index 7267f6eb3ba..ba35ebebe88 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 599b1a76249..083b7beda41 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -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] diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index 855cccc8ddd..f6ab5fc31db 100644 --- a/litellm/types/realtime.py +++ b/litellm/types/realtime.py @@ -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 diff --git a/litellm/types/utils.py b/litellm/types/utils.py index dfc98a9d89d..6667e2ee400 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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 diff --git a/litellm/utils.py b/litellm/utils.py index dd35c17809f..d39d045468e 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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() diff --git a/tests/e2e/coverage_registry/llm_nonconversational.yaml b/tests/e2e/coverage_registry/llm_nonconversational.yaml index 3e389acc2a9..df12b6feea0 100644 --- a/tests/e2e/coverage_registry/llm_nonconversational.yaml +++ b/tests/e2e/coverage_registry/llm_nonconversational.yaml @@ -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)"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index e009b02b69c..43226b637eb 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -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"), ] diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 7e6d4d24905..38b27895a35 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -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() diff --git a/tests/test_litellm/llms/azure/test_azure_common_utils.py b/tests/test_litellm/llms/azure/test_azure_common_utils.py index caf941ebd19..26f66d4055c 100644 --- a/tests/test_litellm/llms/azure/test_azure_common_utils.py +++ b/tests/test_litellm/llms/azure/test_azure_common_utils.py @@ -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", ] ], ) diff --git a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py index f7a88b5ba63..ebf93df7191 100644 --- a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py +++ b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py @@ -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 diff --git a/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py b/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py index 54f206d098d..c663efa341b 100644 --- a/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py +++ b/tests/test_litellm/llms/openai/realtime/test_transcription_sessions.py @@ -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 diff --git a/tests/test_litellm/llms/openai/realtime/test_translation.py b/tests/test_litellm/llms/openai/realtime/test_translation.py new file mode 100644 index 00000000000..6eb2a618b93 --- /dev/null +++ b/tests/test_litellm/llms/openai/realtime/test_translation.py @@ -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" diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index e5179387f82..1780bb5a101 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -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""" diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_audio.py b/tests/test_litellm/proxy/proxy_server/test_routes_audio.py index de76c7257cf..a756454405e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_audio.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_audio.py @@ -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 diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py index f5c97142dde..10c2543021e 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_realtime_webrtc_endpoints.py @@ -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 ( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 89fd9c5c9d4..6f1e5dc9ef8 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -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() diff --git a/tests/test_litellm/responses/test_streaming_iterator_error_events.py b/tests/test_litellm/responses/test_streaming_iterator_error_events.py index e7cf09909fe..1d884e39682 100644 --- a/tests/test_litellm/responses/test_streaming_iterator_error_events.py +++ b/tests/test_litellm/responses/test_streaming_iterator_error_events.py @@ -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), diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index ce280cc3513..6428b0c96ee 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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", diff --git a/tests/unit/llms/openai/transcriptions/test_gpt_transcribe.py b/tests/unit/llms/openai/transcriptions/test_gpt_transcribe.py new file mode 100644 index 00000000000..6ccbeaf690a --- /dev/null +++ b/tests/unit/llms/openai/transcriptions/test_gpt_transcribe.py @@ -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" diff --git a/tests/unit/llms/openai/transcriptions/test_transcription_duration_hidden.py b/tests/unit/llms/openai/transcriptions/test_transcription_duration_hidden.py index 703fa13cbc9..a15ba6d7543 100644 --- a/tests/unit/llms/openai/transcriptions/test_transcription_duration_hidden.py +++ b/tests/unit/llms/openai/transcriptions/test_transcription_duration_hidden.py @@ -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): """