mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
feat(realtime): support latest OpenAI audio models
This commit is contained in:
parent
bc3b5b1d5b
commit
bd43c233ef
50 changed files with 3744 additions and 496 deletions
242
cookbook/gpt_realtime_translate.py
Normal file
242
cookbook/gpt_realtime_translate.py
Normal file
|
|
@ -0,0 +1,242 @@
|
|||
#!/usr/bin/env python3
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import wave
|
||||
from collections.abc import Iterator, Sequence
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlencode, urlsplit, urlunsplit
|
||||
|
||||
import websockets
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
SAMPLE_RATE = 24_000
|
||||
CHANNELS = 1
|
||||
SAMPLE_WIDTH = 2
|
||||
CHUNK_DURATION_SECONDS = 0.1
|
||||
CHUNK_BYTES = int(SAMPLE_RATE * CHANNELS * SAMPLE_WIDTH * CHUNK_DURATION_SECONDS)
|
||||
OUTPUT_IDLE_TIMEOUT_SECONDS = 3.0
|
||||
INITIAL_RESPONSE_TIMEOUT_SECONDS = 30.0
|
||||
AUDIO_EVENT_TYPES = frozenset(
|
||||
{
|
||||
"session.output_audio.delta",
|
||||
"response.audio.delta",
|
||||
"response.output_audio.delta",
|
||||
}
|
||||
)
|
||||
TRANSCRIPT_EVENT_TYPES = frozenset(
|
||||
{
|
||||
"session.output_transcript.delta",
|
||||
"response.text.delta",
|
||||
"response.output_audio_transcript.delta",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Settings:
|
||||
input_wav: Path
|
||||
output_wav: Path
|
||||
base_url: str
|
||||
model: str
|
||||
target_language: str
|
||||
trailing_silence_seconds: float
|
||||
api_key: str
|
||||
|
||||
|
||||
def write_stdout(message: str = "", *, end: str = "\n", flush: bool = False) -> None:
|
||||
sys.stdout.write(f"{message}{end}")
|
||||
if flush:
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
def write_stderr(message: str) -> None:
|
||||
sys.stderr.write(f"{message}\n")
|
||||
|
||||
|
||||
def parse_args(argv: Sequence[str] | None = None) -> Settings | str:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Stream a 24 kHz PCM16 WAV through gpt-realtime-translate and save the translated audio",
|
||||
)
|
||||
parser.add_argument("input_wav", type=Path)
|
||||
parser.add_argument("--output", type=Path, default=Path("translated.wav"))
|
||||
parser.add_argument("--base-url", default=os.getenv("LITELLM_BASE_URL", "http://localhost:4000"))
|
||||
parser.add_argument("--model", default=os.getenv("REALTIME_TRANSLATE_MODEL", "gpt-realtime-translate"))
|
||||
parser.add_argument("--target-language", default="fr")
|
||||
parser.add_argument("--trailing-silence", type=float, default=1.5)
|
||||
parsed = parser.parse_args(argv)
|
||||
api_key = os.getenv("LITELLM_API_KEY") or os.getenv("OPENAI_API_KEY")
|
||||
if not api_key:
|
||||
return "Set LITELLM_API_KEY or OPENAI_API_KEY before running the script"
|
||||
if parsed.trailing_silence < 0:
|
||||
return "--trailing-silence must be zero or greater"
|
||||
return Settings(
|
||||
input_wav=parsed.input_wav,
|
||||
output_wav=parsed.output,
|
||||
base_url=parsed.base_url,
|
||||
model=parsed.model,
|
||||
target_language=parsed.target_language,
|
||||
trailing_silence_seconds=parsed.trailing_silence,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
|
||||
def translation_url(base_url: str, model: str) -> str | None:
|
||||
parsed = urlsplit(base_url.rstrip("/"))
|
||||
scheme = {"http": "ws", "https": "wss", "ws": "ws", "wss": "wss"}.get(parsed.scheme)
|
||||
if not scheme or not parsed.netloc:
|
||||
return None
|
||||
base_path = parsed.path.rstrip("/")
|
||||
realtime_path = (
|
||||
f"{base_path}/realtime/translations" if base_path.endswith("/v1") else f"{base_path}/v1/realtime/translations"
|
||||
)
|
||||
return urlunsplit((scheme, parsed.netloc, realtime_path, urlencode({"model": model}), ""))
|
||||
|
||||
|
||||
def read_pcm16_wav(path: Path) -> bytes | str:
|
||||
try:
|
||||
with wave.open(str(path), "rb") as source:
|
||||
actual_format = (
|
||||
source.getnchannels(),
|
||||
source.getsampwidth(),
|
||||
source.getframerate(),
|
||||
source.getcomptype(),
|
||||
)
|
||||
expected_format = (CHANNELS, SAMPLE_WIDTH, SAMPLE_RATE, "NONE")
|
||||
if actual_format != expected_format:
|
||||
return (
|
||||
f"{path} must be mono, 16-bit PCM, 24 kHz WAV; received "
|
||||
f"channels={actual_format[0]}, sample_width={actual_format[1]}, "
|
||||
f"sample_rate={actual_format[2]}, compression={actual_format[3]}"
|
||||
)
|
||||
return source.readframes(source.getnframes())
|
||||
except (OSError, EOFError, wave.Error) as exc:
|
||||
return f"Unable to read {path}: {exc}"
|
||||
|
||||
|
||||
def audio_chunks(audio: bytes) -> Iterator[bytes]:
|
||||
return (audio[offset : offset + CHUNK_BYTES] for offset in range(0, len(audio), CHUNK_BYTES))
|
||||
|
||||
|
||||
def audio_message(audio: bytes) -> str:
|
||||
return json.dumps(
|
||||
{
|
||||
"type": "session.input_audio_buffer.append",
|
||||
"audio": base64.b64encode(audio).decode("ascii"),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def configure_session(connection: ClientConnection, target_language: str) -> str | None:
|
||||
await connection.send(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "session.update",
|
||||
"session": {"audio": {"output": {"language": target_language}}},
|
||||
}
|
||||
)
|
||||
)
|
||||
while True:
|
||||
raw_event = await asyncio.wait_for(connection.recv(), timeout=20)
|
||||
event = json.loads(raw_event)
|
||||
event_type = event.get("type")
|
||||
if event_type == "session.created":
|
||||
write_stdout(f"Session: {event.get('session', {}).get('id', 'created')}")
|
||||
if event_type == "session.updated":
|
||||
return None
|
||||
if event_type == "error":
|
||||
return f"Realtime API error: {json.dumps(event.get('error', event), ensure_ascii=False)}"
|
||||
|
||||
|
||||
async def send_audio(
|
||||
connection: ClientConnection, pcm: bytes, trailing_silence_seconds: float, finished: asyncio.Event
|
||||
) -> None:
|
||||
silence = bytes(round(trailing_silence_seconds * SAMPLE_RATE * CHANNELS * SAMPLE_WIDTH))
|
||||
try:
|
||||
for chunk in audio_chunks(pcm + silence):
|
||||
await connection.send(audio_message(chunk))
|
||||
await asyncio.sleep(len(chunk) / (SAMPLE_RATE * CHANNELS * SAMPLE_WIDTH))
|
||||
finally:
|
||||
finished.set()
|
||||
|
||||
|
||||
async def receive_translation(
|
||||
connection: ClientConnection, output_path: Path, sender_finished: asyncio.Event
|
||||
) -> str | None:
|
||||
audio_received = asyncio.Event()
|
||||
try:
|
||||
with wave.open(str(output_path), "wb") as output:
|
||||
output.setnchannels(CHANNELS)
|
||||
output.setsampwidth(SAMPLE_WIDTH)
|
||||
output.setframerate(SAMPLE_RATE)
|
||||
write_stdout("Translation: ", end="", flush=True)
|
||||
while True:
|
||||
timeout = OUTPUT_IDLE_TIMEOUT_SECONDS if sender_finished.is_set() else INITIAL_RESPONSE_TIMEOUT_SECONDS
|
||||
try:
|
||||
raw_event = await asyncio.wait_for(connection.recv(), timeout=timeout)
|
||||
except TimeoutError:
|
||||
if sender_finished.is_set() and audio_received.is_set():
|
||||
write_stdout()
|
||||
return None
|
||||
return "The translation stream ended without translated audio"
|
||||
event = json.loads(raw_event)
|
||||
event_type = event.get("type")
|
||||
if event_type in AUDIO_EVENT_TYPES:
|
||||
output.writeframes(base64.b64decode(event.get("delta", ""), validate=True))
|
||||
audio_received.set()
|
||||
elif event_type in TRANSCRIPT_EVENT_TYPES:
|
||||
write_stdout(event.get("delta", event.get("text", "")), end="", flush=True)
|
||||
elif event_type == "error":
|
||||
return f"Realtime API error: {json.dumps(event.get('error', event), ensure_ascii=False)}"
|
||||
except (OSError, wave.Error) as exc:
|
||||
return f"Unable to write {output_path}: {exc}"
|
||||
|
||||
|
||||
async def translate(settings: Settings, pcm: bytes) -> str | None:
|
||||
url = translation_url(settings.base_url, settings.model)
|
||||
if not url:
|
||||
return f"Invalid --base-url: {settings.base_url}"
|
||||
sender_finished = asyncio.Event()
|
||||
try:
|
||||
async with websockets.connect(
|
||||
url,
|
||||
additional_headers={"Authorization": f"Bearer {settings.api_key}"},
|
||||
proxy=None,
|
||||
open_timeout=20,
|
||||
close_timeout=5,
|
||||
) as connection:
|
||||
configuration_error = await configure_session(connection, settings.target_language)
|
||||
if configuration_error:
|
||||
return configuration_error
|
||||
async with asyncio.TaskGroup() as tasks:
|
||||
receiver = tasks.create_task(receive_translation(connection, settings.output_wav, sender_finished))
|
||||
tasks.create_task(send_audio(connection, pcm, settings.trailing_silence_seconds, sender_finished))
|
||||
return receiver.result()
|
||||
except Exception as exc:
|
||||
return f"Translation failed: {type(exc).__name__}: {exc}"
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
settings = parse_args(argv)
|
||||
if isinstance(settings, str):
|
||||
write_stderr(settings)
|
||||
return 2
|
||||
pcm = read_pcm16_wav(settings.input_wav)
|
||||
if isinstance(pcm, str):
|
||||
write_stderr(pcm)
|
||||
return 2
|
||||
error = asyncio.run(translate(settings, pcm))
|
||||
if error:
|
||||
write_stderr(error)
|
||||
return 1
|
||||
write_stdout(f"Translated audio: {settings.output_wav}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -6,4 +6,16 @@ Supported endpoints:
|
|||
|
||||
Supported providers: OpenAI, Azure OpenAI, Bedrock, Vertex AI, xAI.
|
||||
|
||||
For user-facing documentation and usage examples, see the litellm-docs repo.
|
||||
Billing visibility:
|
||||
- WebSocket sessions pass provider usage events through LiteLLM and support local spend tracking
|
||||
- Client-secret and SDP call endpoints only proxy session setup; subsequent WebRTC media and usage events travel over the peer connection, so LiteLLM cannot record inference spend or enforce spend-based budgets for those sessions
|
||||
- Use the proxied WebSocket transport when LiteLLM spend logs and budgets must include Realtime inference
|
||||
|
||||
Non-billable Realtime protocols are disabled by default. Operators who accept the billing and budget-enforcement limitation can opt in:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
allow_non_billable_realtime_protocols: true
|
||||
```
|
||||
|
||||
For user-facing documentation and usage examples, see the litellm-docs repo.
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)"}
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
265
tests/test_litellm/llms/openai/realtime/test_translation.py
Normal file
265
tests/test_litellm/llms/openai/realtime/test_translation.py
Normal file
|
|
@ -0,0 +1,265 @@
|
|||
import gzip
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from litellm.llms.azure.realtime.http_transformation import AzureRealtimeHTTPConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig
|
||||
from litellm.types.realtime import RealtimeSessionConfig
|
||||
|
||||
|
||||
def test_realtime_session_config_supports_translation_and_live_transcription_fields():
|
||||
session = RealtimeSessionConfig(
|
||||
type="translation",
|
||||
model="gpt-realtime-translate",
|
||||
audio={
|
||||
"input": {
|
||||
"transcription": {
|
||||
"model": "gpt-live-transcribe",
|
||||
"delay": "minimal",
|
||||
"languages": ["en", "fr"],
|
||||
"keywords": ["LiteLLM"],
|
||||
}
|
||||
},
|
||||
"output": {"language": "es"},
|
||||
},
|
||||
)
|
||||
|
||||
assert session.audio is not None
|
||||
assert session.audio.input is not None
|
||||
assert session.audio.input.transcription is not None
|
||||
assert session.audio.input.transcription.delay == "minimal"
|
||||
assert session.audio.input.transcription.languages == ["en", "fr"]
|
||||
assert session.audio.input.transcription.keywords == ["LiteLLM"]
|
||||
assert session.audio.output is not None
|
||||
assert session.audio.output.language == "es"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base,expected",
|
||||
[
|
||||
(
|
||||
"https://api.openai.com",
|
||||
"https://api.openai.com/v1/realtime/translations/client_secrets",
|
||||
),
|
||||
(
|
||||
"https://api.openai.com/v1",
|
||||
"https://api.openai.com/v1/realtime/translations/client_secrets",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_openai_translation_client_secret_url(api_base: str, expected: str):
|
||||
config = OpenAIRealtimeHTTPConfig()
|
||||
assert config.get_translation_client_secret_url(api_base, "gpt-realtime-translate") == expected
|
||||
|
||||
|
||||
def test_openai_translation_calls_url():
|
||||
config = OpenAIRealtimeHTTPConfig()
|
||||
assert (
|
||||
config.get_translation_calls_url("https://api.openai.com/v1", "gpt-realtime-translate")
|
||||
== "https://api.openai.com/v1/realtime/translations/calls"
|
||||
)
|
||||
|
||||
|
||||
def test_azure_translation_urls_use_ga_paths():
|
||||
config = AzureRealtimeHTTPConfig()
|
||||
assert (
|
||||
config.get_translation_client_secret_url("https://example.openai.azure.com", "translate-deployment")
|
||||
== "https://example.openai.azure.com/openai/v1/realtime/translations/client_secrets"
|
||||
)
|
||||
assert (
|
||||
config.get_translation_calls_url("https://example.openai.azure.com", "translate-deployment")
|
||||
== "https://example.openai.azure.com/openai/v1/realtime/translations/calls"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_translation_client_secret_uses_custom_translation_path():
|
||||
client = MagicMock(spec=AsyncHTTPHandler)
|
||||
client.post = AsyncMock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={"value": "ek_test"},
|
||||
request=httpx.Request("POST", "https://api.openai.com/v1/realtime/translations/client_secrets"),
|
||||
)
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
handler = BaseLLMHTTPHandler()
|
||||
request_data = {"session": {"type": "translation", "model": "gpt-realtime-translate"}}
|
||||
|
||||
response = await handler.async_realtime_translation_client_secret_handler(
|
||||
api_base="https://api.openai.com",
|
||||
api_key="sk-test",
|
||||
request_data=request_data,
|
||||
logging_obj=logging_obj,
|
||||
timeout=10,
|
||||
provider_config=OpenAIRealtimeHTTPConfig(),
|
||||
model="gpt-realtime-translate",
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
call = client.post.call_args.kwargs
|
||||
assert call["url"] == "https://api.openai.com/v1/realtime/translations/client_secrets"
|
||||
assert call["json"] == request_data
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_azure_translation_client_secret_supports_entra_bearer_auth():
|
||||
client = MagicMock(spec=AsyncHTTPHandler)
|
||||
client.post = AsyncMock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={"value": "ek_test"},
|
||||
request=httpx.Request(
|
||||
"POST",
|
||||
"https://example.openai.azure.com/openai/v1/realtime/translations/client_secrets",
|
||||
),
|
||||
)
|
||||
)
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
await handler.async_realtime_translation_client_secret_handler(
|
||||
api_base="https://example.openai.azure.com",
|
||||
api_key="",
|
||||
request_data={"session": {"type": "translation", "model": "translate-deployment"}},
|
||||
logging_obj=MagicMock(),
|
||||
timeout=10,
|
||||
provider_config=AzureRealtimeHTTPConfig(),
|
||||
model="translate-deployment",
|
||||
extra_headers={"Authorization": "Bearer entra-token"},
|
||||
client=client,
|
||||
)
|
||||
|
||||
headers = client.post.call_args.kwargs["headers"]
|
||||
assert headers["Authorization"] == "Bearer entra-token"
|
||||
assert "api-key" not in headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_translation_calls_use_translation_session_and_path():
|
||||
client = MagicMock(spec=AsyncHTTPHandler)
|
||||
client.post = AsyncMock(
|
||||
return_value=httpx.Response(
|
||||
201,
|
||||
content=b"v=0\r\n",
|
||||
request=httpx.Request("POST", "https://api.openai.com/v1/realtime/translations/calls"),
|
||||
)
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
response = await handler.async_realtime_calls_handler(
|
||||
api_base="https://api.openai.com",
|
||||
openai_ephemeral_key="ek_test",
|
||||
sdp_body=b"v=0\r\n",
|
||||
logging_obj=logging_obj,
|
||||
timeout=10,
|
||||
provider_config=OpenAIRealtimeHTTPConfig(),
|
||||
model="gpt-realtime-translate",
|
||||
client=client,
|
||||
translation=True,
|
||||
)
|
||||
|
||||
assert response.status_code == 201
|
||||
call = client.post.call_args.kwargs
|
||||
assert call["url"] == "https://api.openai.com/v1/realtime/translations/calls"
|
||||
assert call["headers"]["Content-Type"] == "application/sdp"
|
||||
assert call["content"] == "v=0\r\n"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_standard_client_secret_uses_openai_sdk_resource():
|
||||
async def send_response(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url.path == "/v1/realtime/client_secrets"
|
||||
return httpx.Response(
|
||||
200,
|
||||
content=gzip.compress(b'{"value":"ek_test"}'),
|
||||
headers={"content-encoding": "gzip"},
|
||||
)
|
||||
|
||||
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
|
||||
openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client)
|
||||
handler = BaseLLMHTTPHandler()
|
||||
logging_obj = MagicMock()
|
||||
|
||||
response = await handler.async_realtime_client_secret_handler(
|
||||
api_base="https://example.com",
|
||||
api_key="sk-test",
|
||||
request_data={"session": {"type": "realtime", "model": "gpt-realtime-2.1"}},
|
||||
logging_obj=logging_obj,
|
||||
timeout=10,
|
||||
client=openai_client,
|
||||
use_openai_sdk=True,
|
||||
)
|
||||
await openai_client.close()
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"value": "ek_test"}
|
||||
assert "content-encoding" not in response.headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_translation_client_secret_uses_openai_sdk_custom_post():
|
||||
async def send_response(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url.path == "/v1/realtime/translations/client_secrets"
|
||||
assert json.loads(request.content)["session"]["audio"]["output"]["language"] == "es"
|
||||
return httpx.Response(200, json={"value": "ek_translation"})
|
||||
|
||||
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
|
||||
openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client)
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
response = await handler.async_realtime_translation_client_secret_handler(
|
||||
api_base="https://example.com",
|
||||
api_key="sk-test",
|
||||
request_data={
|
||||
"session": {
|
||||
"model": "gpt-realtime-translate",
|
||||
"audio": {"output": {"language": "es"}},
|
||||
}
|
||||
},
|
||||
logging_obj=MagicMock(),
|
||||
timeout=10,
|
||||
client=openai_client,
|
||||
use_openai_sdk=True,
|
||||
)
|
||||
await openai_client.close()
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"value": "ek_translation"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_translation_calls_use_openai_sdk_custom_post():
|
||||
async def send_response(request: httpx.Request) -> httpx.Response:
|
||||
assert request.url.path == "/v1/realtime/translations/calls"
|
||||
body = await request.aread()
|
||||
assert request.headers["content-type"] == "application/sdp"
|
||||
assert body == b"v=0\r\n"
|
||||
return httpx.Response(201, content=b"v=0\r\n")
|
||||
|
||||
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
|
||||
openai_client = AsyncOpenAI(api_key="ek_test", base_url="https://example.com/v1", http_client=http_client)
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
response = await handler.async_realtime_calls_handler(
|
||||
api_base="https://example.com",
|
||||
openai_ephemeral_key="ek_test",
|
||||
sdp_body=b"v=0\r\n",
|
||||
logging_obj=MagicMock(),
|
||||
timeout=10,
|
||||
model="gpt-realtime-translate",
|
||||
client=openai_client,
|
||||
translation=True,
|
||||
use_openai_sdk=True,
|
||||
)
|
||||
await openai_client.close()
|
||||
|
||||
assert response.status_code == 201
|
||||
assert response.text == "v=0\r\n"
|
||||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
297
tests/unit/llms/openai/transcriptions/test_gpt_transcribe.py
Normal file
297
tests/unit/llms/openai/transcriptions/test_gpt_transcribe.py
Normal file
|
|
@ -0,0 +1,297 @@
|
|||
import io
|
||||
import json
|
||||
import wave
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openai import AsyncOpenAI, AsyncStream, AzureOpenAI
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.audio_utils.transcription_streaming import wrap_transcription_stream
|
||||
from litellm.llms.azure.audio_transcriptions import AzureAudioTranscription
|
||||
from litellm.llms.openai.transcriptions.gpt_transformation import (
|
||||
OpenAIGPTTranscribeAudioTranscriptionConfig,
|
||||
)
|
||||
from litellm.llms.openai.transcriptions.handler import OpenAIAudioTranscription
|
||||
from litellm.main import _validate_gpt_transcription_request
|
||||
from litellm.types.utils import TranscriptionResponse
|
||||
from litellm.utils import get_optional_params_transcription
|
||||
|
||||
|
||||
def test_gpt_transcribe_config_uses_native_parameters_and_json():
|
||||
config = OpenAIGPTTranscribeAudioTranscriptionConfig()
|
||||
supported = config.get_supported_openai_params("gpt-transcribe")
|
||||
assert supported == ["prompt", "response_format", "keywords", "languages", "stream"]
|
||||
|
||||
audio_file = io.BytesIO(b"audio")
|
||||
request = config.transform_audio_transcription_request(
|
||||
model="gpt-transcribe",
|
||||
audio_file=audio_file,
|
||||
optional_params={"keywords": ["LiteLLM"], "languages": ["en", "fr"], "stream": True},
|
||||
litellm_params={},
|
||||
)
|
||||
assert request.data["response_format"] == "json"
|
||||
assert request.data["keywords"] == ["LiteLLM"]
|
||||
assert request.data["languages"] == ["en", "fr"]
|
||||
assert request.data["stream"] is True
|
||||
|
||||
|
||||
def test_gpt_transcribe_optional_params_are_preserved():
|
||||
params = get_optional_params_transcription(
|
||||
model="gpt-transcribe",
|
||||
custom_llm_provider="openai",
|
||||
keywords=["LiteLLM", "Realtime API"],
|
||||
languages=["en", "fr"],
|
||||
stream=True,
|
||||
)
|
||||
assert params == {
|
||||
"keywords": ["LiteLLM", "Realtime API"],
|
||||
"languages": ["en", "fr"],
|
||||
"stream": True,
|
||||
}
|
||||
|
||||
|
||||
def test_transcription_response_preserves_empty_languages():
|
||||
response = TranscriptionResponse(text="hello", languages=[])
|
||||
assert response.model_dump()["languages"] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_handler_returns_native_typed_stream():
|
||||
async def send_response(request: httpx.Request) -> httpx.Response:
|
||||
body = await request.aread()
|
||||
assert b'name="keywords[]"' in body
|
||||
assert b'name="languages[]"' in body
|
||||
assert b'name="stream"' in body
|
||||
events = (
|
||||
{"type": "transcript.text.delta", "delta": "hello "},
|
||||
{
|
||||
"type": "transcript.text.done",
|
||||
"text": "hello world",
|
||||
"languages": [],
|
||||
"usage": {
|
||||
"type": "tokens",
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 2,
|
||||
"total_tokens": 12,
|
||||
},
|
||||
},
|
||||
)
|
||||
content = "".join(f"data: {json.dumps(event)}\n\n" for event in events).encode()
|
||||
return httpx.Response(200, content=content, headers={"content-type": "text/event-stream"})
|
||||
|
||||
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
|
||||
openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
audio_file = io.BytesIO(b"audio")
|
||||
audio_file.name = "sample.wav"
|
||||
|
||||
handler = OpenAIAudioTranscription()
|
||||
result = handler.audio_transcriptions(
|
||||
model="gpt-transcribe",
|
||||
audio_file=audio_file,
|
||||
optional_params={"keywords": ["LiteLLM"], "languages": ["en"], "stream": True},
|
||||
litellm_params={},
|
||||
model_response=TranscriptionResponse(),
|
||||
timeout=10,
|
||||
max_retries=0,
|
||||
logging_obj=logging_obj,
|
||||
api_key="sk-test",
|
||||
api_base="https://example.com/v1",
|
||||
client=openai_client,
|
||||
atranscription=True,
|
||||
provider_config=OpenAIGPTTranscribeAudioTranscriptionConfig(),
|
||||
)
|
||||
stream = await result
|
||||
assert isinstance(stream, AsyncStream)
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.async_failure_handler = AsyncMock()
|
||||
wrapped_stream = wrap_transcription_stream(stream, logging_obj, datetime.now())
|
||||
received = [event async for event in wrapped_stream]
|
||||
await wrapped_stream.close()
|
||||
await openai_client.close()
|
||||
|
||||
assert [event.type for event in received] == ["transcript.text.delta", "transcript.text.done"]
|
||||
assert received[-1].languages == []
|
||||
logging_obj.async_success_handler.assert_awaited_once()
|
||||
logged_response = logging_obj.async_success_handler.await_args.kwargs["result"]
|
||||
assert logged_response.text == "hello world"
|
||||
assert logged_response.languages == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_atranscription_stream_preserves_duration_for_callback_cost():
|
||||
async def send_response(request: httpx.Request) -> httpx.Response:
|
||||
events = (
|
||||
{"type": "transcript.text.delta", "delta": "hello "},
|
||||
{
|
||||
"type": "transcript.text.done",
|
||||
"text": "hello world",
|
||||
"usage": {
|
||||
"type": "tokens",
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 2,
|
||||
"total_tokens": 12,
|
||||
},
|
||||
},
|
||||
)
|
||||
content = "".join(f"data: {json.dumps(event)}\n\n" for event in events).encode()
|
||||
return httpx.Response(200, content=content, headers={"content-type": "text/event-stream"})
|
||||
|
||||
http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response))
|
||||
openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.async_failure_handler = AsyncMock()
|
||||
audio_file = io.BytesIO()
|
||||
with wave.open(audio_file, "wb") as wav_file:
|
||||
wav_file.setnchannels(1)
|
||||
wav_file.setsampwidth(2)
|
||||
wav_file.setframerate(16000)
|
||||
wav_file.writeframes(b"\x00\x00" * 16000)
|
||||
audio_file.name = "sample.wav"
|
||||
|
||||
stream = await litellm.atranscription(
|
||||
model="openai/gpt-transcribe",
|
||||
file=audio_file,
|
||||
stream=True,
|
||||
client=openai_client,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
received = [event async for event in stream]
|
||||
await stream.close()
|
||||
await openai_client.close()
|
||||
|
||||
assert [event.type for event in received] == ["transcript.text.delta", "transcript.text.done"]
|
||||
logging_obj.async_success_handler.assert_awaited_once()
|
||||
logged_response = logging_obj.async_success_handler.await_args.kwargs["result"]
|
||||
assert logged_response._hidden_params["audio_transcription_duration"] == pytest.approx(1.0)
|
||||
|
||||
|
||||
def test_gpt_transcribe_rejects_conflicting_language_inputs():
|
||||
audio_file = io.BytesIO(b"audio")
|
||||
audio_file.name = "sample.wav"
|
||||
with pytest.raises(litellm.UnsupportedParamsError, match="cannot be used together"):
|
||||
litellm.transcription(
|
||||
model="gpt-transcribe",
|
||||
file=audio_file,
|
||||
language="en",
|
||||
languages=["fr"],
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
|
||||
def test_gpt_transcribe_rejects_whisper_response_formats():
|
||||
audio_file = io.BytesIO(b"audio")
|
||||
audio_file.name = "sample.wav"
|
||||
with pytest.raises(litellm.UnsupportedParamsError, match="only supports response_format='json'"):
|
||||
litellm.transcription(
|
||||
model="gpt-transcribe",
|
||||
file=audio_file,
|
||||
response_format="verbose_json",
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
|
||||
def test_gpt_live_transcribe_rejects_file_transcription():
|
||||
audio_file = io.BytesIO(b"audio")
|
||||
audio_file.name = "sample.wav"
|
||||
with pytest.raises(litellm.UnsupportedParamsError, match="Realtime API"):
|
||||
litellm.transcription(
|
||||
model="gpt-live-transcribe",
|
||||
file=audio_file,
|
||||
api_key="sk-test",
|
||||
)
|
||||
|
||||
|
||||
def test_azure_async_gpt_transcribe_forwards_v1_api_version():
|
||||
handler = AzureAudioTranscription()
|
||||
handler.async_audio_transcriptions = MagicMock(return_value=MagicMock())
|
||||
|
||||
handler.audio_transcriptions(
|
||||
model="gpt-transcribe",
|
||||
audio_file=io.BytesIO(b"audio"),
|
||||
optional_params={"stream": True},
|
||||
logging_obj=MagicMock(),
|
||||
model_response=TranscriptionResponse(),
|
||||
timeout=10,
|
||||
max_retries=0,
|
||||
api_key="sk-test",
|
||||
api_base="https://example.openai.azure.com",
|
||||
api_version="v1",
|
||||
atranscription=True,
|
||||
)
|
||||
|
||||
assert handler.async_audio_transcriptions.call_args.kwargs["api_version"] == "v1"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("api_version", [None, "v1", "latest", "preview"])
|
||||
def test_azure_gpt_transcribe_uses_deployment_scoped_api_version(api_version: str | None):
|
||||
resolved_api_version = _validate_gpt_transcription_request(
|
||||
model="gpt-transcribe",
|
||||
custom_llm_provider="azure",
|
||||
language=None,
|
||||
languages=None,
|
||||
response_format="json",
|
||||
api_version=api_version,
|
||||
)
|
||||
|
||||
assert resolved_api_version == litellm.AZURE_DEFAULT_API_VERSION
|
||||
|
||||
|
||||
def test_azure_gpt_transcribe_uses_deployment_scoped_route():
|
||||
def send_response(request: httpx.Request) -> httpx.Response:
|
||||
assert str(request.url) == (
|
||||
"https://example.openai.azure.com/openai/deployments/gpt-transcribe/audio/transcriptions"
|
||||
f"?api-version={litellm.AZURE_DEFAULT_API_VERSION}"
|
||||
)
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={"text": "hello", "languages": [{"code": "en"}], "usage": {"type": "duration", "seconds": 1}},
|
||||
)
|
||||
|
||||
http_client = httpx.Client(transport=httpx.MockTransport(send_response))
|
||||
client = AzureOpenAI(
|
||||
api_key="azure-test-key",
|
||||
azure_endpoint="https://example.openai.azure.com",
|
||||
api_version=litellm.AZURE_DEFAULT_API_VERSION,
|
||||
http_client=http_client,
|
||||
)
|
||||
audio_file = io.BytesIO(b"audio")
|
||||
audio_file.name = "sample.wav"
|
||||
|
||||
response = AzureAudioTranscription().audio_transcriptions(
|
||||
model="gpt-transcribe",
|
||||
audio_file=audio_file,
|
||||
optional_params={"response_format": "json"},
|
||||
logging_obj=MagicMock(),
|
||||
model_response=TranscriptionResponse(),
|
||||
timeout=10,
|
||||
max_retries=0,
|
||||
api_key="azure-test-key",
|
||||
api_base="https://example.openai.azure.com",
|
||||
api_version=litellm.AZURE_DEFAULT_API_VERSION,
|
||||
client=client,
|
||||
)
|
||||
|
||||
assert response.text == "hello"
|
||||
assert response.languages is not None
|
||||
assert [language.code for language in response.languages] == ["en"]
|
||||
client.close()
|
||||
|
||||
|
||||
def test_azure_gpt_transcribe_preserves_dated_api_version():
|
||||
resolved_api_version = _validate_gpt_transcription_request(
|
||||
model="gpt-transcribe",
|
||||
custom_llm_provider="azure",
|
||||
language=None,
|
||||
languages=None,
|
||||
response_format="json",
|
||||
api_version="2025-04-01-preview",
|
||||
)
|
||||
|
||||
assert resolved_api_version == "2025-04-01-preview"
|
||||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue