fix(realtime): preserve input usage and provider routing

This commit is contained in:
Emerson Gomes 2026-09-23 03:09:01 -05:00
parent 17dbd6ae92
commit 828707b9f2
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
12 changed files with 138 additions and 90 deletions

View file

@ -440,18 +440,30 @@ class RealTimeStreaming:
if event_obj.get("type") == "session.closed":
usage: Final = event_obj.get("usage")
output_seconds: Final = usage.get("output_seconds") if isinstance(usage, dict) else None
if isinstance(output_seconds, (int, float)):
input_seconds: Final = usage.get("input_seconds") if isinstance(usage, dict) else None
input_seconds: Final = usage.get("input_seconds") if isinstance(usage, dict) else None
synthetic_output_seconds: Final = (
self._translation_output_audio_bytes / self._translation_output_bytes_per_second
if self._translation_output_audio_bytes > 0
else None
)
resolved_output_seconds: Final = (
output_seconds if isinstance(output_seconds, (int, float)) else synthetic_output_seconds
)
if isinstance(input_seconds, (int, float)) or resolved_output_seconds is not None:
if not self._should_store_message(event_obj):
self.messages.append(
OpenAIRealtimeTranslationClosedEvent(
type="session.closed",
usage=OpenAIRealtimeTranslationDurationUsage(
type="duration",
output_seconds=output_seconds,
**({"input_seconds": input_seconds} if isinstance(input_seconds, (int, float)) else {}),
),
normalized_usage: Final = (
OpenAIRealtimeTranslationDurationUsage(
type="duration",
input_seconds=float(input_seconds),
output_seconds=float(resolved_output_seconds or 0.0),
)
if isinstance(input_seconds, (int, float))
else OpenAIRealtimeTranslationDurationUsage(
type="duration", output_seconds=float(resolved_output_seconds or 0.0)
)
)
self.messages.append(
OpenAIRealtimeTranslationClosedEvent(type="session.closed", usage=normalized_usage)
)
self._translation_usage_finalized = True
return
@ -500,7 +512,10 @@ class RealTimeStreaming:
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)):
if isinstance(event_usage, dict) and (
isinstance(event_usage.get("input_seconds"), (int, float))
or isinstance(event_usage.get("output_seconds"), (int, float))
):
self._translation_usage_finalized = True
return
if self._translation_output_audio_bytes == 0:

View file

@ -4,6 +4,7 @@ from typing import TYPE_CHECKING, Any, Final
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from pydantic import BaseModel
import litellm
from litellm._uuid import uuid
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name
from litellm.llms.base_llm.audio_transcription.transformation import sdk_compatible_transcription_request_data
@ -42,6 +43,19 @@ class AzureAudioTranscription(AzureChatCompletion):
) -> TranscriptionResponse | Coroutine[Any, Any, TranscriptionResponse]:
data: Final = {"model": model, "file": audio_file, **optional_params}
sdk_data: Final = sdk_compatible_transcription_request_data(data)
model_info: Final = (
litellm.get_model_info(model=model, custom_llm_provider="azure")
if f"azure/{model}" in litellm.model_cost
else None
)
provider_specific_entry: Final = model_info.get("provider_specific_entry") if model_info is not None else None
resolved_api_version: Final = (
litellm.AZURE_DEFAULT_API_VERSION
if provider_specific_entry is not None
and provider_specific_entry.get("transcription_deployment_api") == 1
and api_version in ("v1", "latest", "preview")
else api_version
)
if atranscription is True:
return self.async_audio_transcriptions(
@ -51,7 +65,7 @@ class AzureAudioTranscription(AzureChatCompletion):
timeout=timeout,
api_key=api_key,
api_base=api_base,
api_version=api_version,
api_version=resolved_api_version,
client=client,
max_retries=max_retries,
logging_obj=logging_obj,
@ -61,7 +75,7 @@ class AzureAudioTranscription(AzureChatCompletion):
)
azure_client: Final = self.get_azure_openai_client(
api_version=api_version,
api_version=resolved_api_version,
api_base=api_base,
api_key=api_key,
model=model,

View file

@ -42,10 +42,16 @@ async def forward_messages(client_ws: Any, backend_ws: Any):
def azure_realtime_requires_ga(model: str) -> bool:
try:
model_info: Final = litellm.get_model_info(model=model, custom_llm_provider="azure")
azure_model_info: Final = litellm.get_model_info(model=model, custom_llm_provider="azure")
except Exception: # noqa: BLE001 # unmapped deployments can select a protocol explicitly
return False
return (model_info.get("provider_specific_entry") or {}).get("realtime_ga_only") == 1
try:
openai_model_info: Final = litellm.get_model_info(model=model, custom_llm_provider="openai")
except Exception: # noqa: BLE001 # unmapped deployments can select a protocol explicitly
return False
openai_entry: Final = openai_model_info.get("provider_specific_entry")
return openai_entry is not None and openai_entry.get("realtime_ga_only") == 1
azure_entry: Final = azure_model_info.get("provider_specific_entry")
return azure_entry is not None and azure_entry.get("realtime_ga_only") == 1
def azure_realtime_protocol_for_client(

View file

@ -7877,16 +7877,15 @@ def _validate_gpt_transcription_request(
language: str | None,
languages: Sequence[str] | None,
response_format: str | None,
api_version: str | None,
) -> str | None:
model_cost_key: Final = f"{custom_llm_provider}/{model}" if custom_llm_provider == "azure" else model
model_info: Final = (
get_model_info(model=model, custom_llm_provider=custom_llm_provider)
if model_cost_key in litellm.model_cost
else {}
) -> None:
model_cost_key: Final = next(
(key for key in (f"{custom_llm_provider}/{model}", model) if key in litellm.model_cost), None
)
supported_endpoints: Final = model_info.get("supported_endpoints")
provider_specific_entry: Final = model_info.get("provider_specific_entry") or {}
model_info: Final = (
get_model_info(model=model, custom_llm_provider=custom_llm_provider) if model_cost_key is not None else None
)
supported_endpoints: Final = model_info.get("supported_endpoints") if model_info is not None else None
provider_specific_entry: Final = model_info.get("provider_specific_entry") if model_info is not None else None
if language is not None and languages is not None:
raise litellm.UnsupportedParamsError(
message="language and languages cannot be used together",
@ -7899,17 +7898,16 @@ def _validate_gpt_transcription_request(
model=model,
llm_provider=custom_llm_provider,
)
if provider_specific_entry.get("transcription_json_only") == 1 and response_format not in (None, "json"):
if (
provider_specific_entry is not None
and provider_specific_entry.get("transcription_json_only") == 1
and response_format not in (None, "json")
):
raise litellm.UnsupportedParamsError(
message=f"{model} only supports response_format='json'",
model=model,
llm_provider=custom_llm_provider,
)
if provider_specific_entry.get("transcription_deployment_api") == 1:
if api_version in ("v1", "latest", "preview"):
return litellm.AZURE_DEFAULT_API_VERSION
return api_version
return api_version
@client
@ -7976,13 +7974,12 @@ 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(
_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(
@ -8038,7 +8035,7 @@ def transcription(
# azure configs
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
azure_api_version: Final = validated_api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
azure_api_version: Final = 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")

View file

@ -33942,6 +33942,9 @@
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "realtime",
"provider_specific_entry": {
"realtime_ga_only": 1
},
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 2.4e-05,
"source": "https://developers.openai.com/api/docs/pricing",

View file

@ -2313,7 +2313,7 @@ class OpenAIRealtimeResponseUsage(TypedDict):
class OpenAIRealtimeTranslationDurationUsage(TypedDict):
type: ReadOnly[Literal["duration"]]
output_seconds: ReadOnly[float]
output_seconds: NotRequired[ReadOnly[float]]
input_seconds: NotRequired[ReadOnly[float]]

View file

@ -33942,6 +33942,9 @@
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "realtime",
"provider_specific_entry": {
"realtime_ga_only": 1
},
"output_cost_per_audio_token": 6.4e-05,
"output_cost_per_token": 2.4e-05,
"source": "https://developers.openai.com/api/docs/pricing",

View file

@ -1,6 +1,6 @@
import asyncio
import json
from collections.abc import Coroutine
from collections.abc import Coroutine, Mapping
from dataclasses import dataclass
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
@ -2999,6 +2999,34 @@ def test_translation_prefers_provider_duration_over_audio_byte_estimate():
assert closed_events[0]["usage"] == {"type": "duration", "input_seconds": 0.25, "output_seconds": 0.5}
@pytest.mark.parametrize(
("output_audio_bytes", "expected_usage"),
[
(0, {"type": "duration", "input_seconds": 0.25, "output_seconds": 0.0}),
(48000, {"type": "duration", "input_seconds": 0.25, "output_seconds": 1.0}),
],
)
def test_translation_preserves_input_only_provider_usage(
output_audio_bytes: int, expected_usage: Mapping[str, str | float]
) -> None:
streaming = RealTimeStreaming(
websocket=MagicMock(),
backend_ws=MagicMock(),
logging_obj=MagicMock(),
model="gpt-realtime-translate",
translation_session=True,
)
streaming._translation_output_audio_bytes = output_audio_bytes
streaming._capture_translation_output_audio(
{"type": "session.closed", "usage": {"type": "duration", "input_seconds": 0.25}}
)
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"] == expected_usage
@pytest.mark.asyncio
async def test_audio_delta_frame_parsed_at_most_once():
client_ws = _beta_client_ws()

View file

@ -47,14 +47,15 @@ def test_azure_transcription_session_url_uses_deployment_and_api_version():
@pytest.mark.parametrize("api_version", [None, "v1", "latest", "preview"])
def test_azure_ga_realtime_http_urls(api_version, monkeypatch: pytest.MonkeyPatch):
@pytest.mark.parametrize("model", ["gpt-realtime-2.1", "gpt-realtime-2", "gpt-realtime-2-2026-05-06"])
def test_azure_ga_realtime_http_urls(api_version, model: str, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
litellm.get_model_info.cache_clear()
cfg = AzureRealtimeHTTPConfig()
base = "https://my.openai.azure.com"
assert cfg.get_complete_url(base, "gpt-realtime-2.1", api_version) == (f"{base}/openai/v1/realtime/client_secrets")
assert cfg.get_realtime_calls_url(base, "gpt-realtime-2.1", api_version) == (f"{base}/openai/v1/realtime/calls")
assert cfg.get_complete_url(base, model, api_version) == (f"{base}/openai/v1/realtime/client_secrets")
assert cfg.get_realtime_calls_url(base, model, 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"
)

View file

@ -4687,12 +4687,22 @@ def test_realtime_translation_duration_cost(_local_model_cost_map):
assert actual == pytest.approx(2 * litellm.model_cost[model]["output_cost_per_second"])
def test_realtime_translation_duration_cost_includes_provider_input_usage(_local_model_cost_map):
@pytest.mark.parametrize("output_seconds", [None, 2.0])
def test_realtime_translation_duration_cost_includes_provider_input_usage(
_local_model_cost_map, output_seconds: float | None
):
from litellm.cost_calculator import handle_realtime_translation_cost_calculation
model: Final = "gpt-realtime-translate"
events: Final[OpenAIRealtimeStreamList] = [
{"type": "session.closed", "usage": {"type": "duration", "input_seconds": 3.0, "output_seconds": 2.0}}
{
"type": "session.closed",
"usage": {
"type": "duration",
"input_seconds": 3.0,
**({"output_seconds": output_seconds} if output_seconds is not None else {}),
},
}
]
actual: Final = handle_realtime_translation_cost_calculation(
results=events,
@ -4702,6 +4712,6 @@ def test_realtime_translation_duration_cost_includes_provider_input_usage(_local
expected: Final = (
3 * litellm.model_cost[model]["input_cost_per_second"]
+ 2 * litellm.model_cost[model]["output_cost_per_second"]
+ (output_seconds or 0) * litellm.model_cost[model]["output_cost_per_second"]
)
assert actual == pytest.approx(expected)

View file

@ -16,7 +16,6 @@ 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
@ -252,7 +251,19 @@ def test_gpt_live_transcribe_rejects_file_transcription(local_model_cost_map: No
)
def test_azure_async_gpt_transcribe_forwards_v1_api_version():
@pytest.mark.parametrize(
("api_version", "expected_api_version"),
[
("v1", litellm.AZURE_DEFAULT_API_VERSION),
("latest", litellm.AZURE_DEFAULT_API_VERSION),
("preview", litellm.AZURE_DEFAULT_API_VERSION),
(None, None),
("2025-04-01-preview", "2025-04-01-preview"),
],
)
def test_azure_gpt_transcribe_resolves_api_version_in_provider(
local_model_cost_map: None, api_version: str | None, expected_api_version: str | None
) -> None:
handler = AzureAudioTranscription()
handler.async_audio_transcriptions = MagicMock(return_value=MagicMock())
@ -266,39 +277,11 @@ def test_azure_async_gpt_transcribe_forwards_v1_api_version():
max_retries=0,
api_key="sk-test",
api_base="https://example.openai.azure.com",
api_version="v1",
api_version=api_version,
atranscription=True,
)
assert handler.async_audio_transcriptions.call_args.kwargs["api_version"] == "v1"
@pytest.mark.parametrize("api_version", ["v1", "latest", "preview"])
def test_azure_gpt_transcribe_uses_deployment_scoped_api_version(local_model_cost_map: None, 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_keeps_unset_api_version_for_configured_default(local_model_cost_map: None) -> None:
assert (
_validate_gpt_transcription_request(
model="gpt-transcribe",
custom_llm_provider="azure",
language=None,
languages=None,
response_format="json",
api_version=None,
)
is None
)
assert handler.async_audio_transcriptions.call_args.kwargs["api_version"] == expected_api_version
def test_azure_gpt_transcribe_uses_deployment_scoped_route():
@ -342,19 +325,6 @@ def test_azure_gpt_transcribe_uses_deployment_scoped_route():
client.close()
def test_azure_gpt_transcribe_preserves_dated_api_version(local_model_cost_map: None) -> None:
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"
@pytest.mark.asyncio
async def test_azure_gpt_transcribe_sends_language_hints_in_sdk_extra_body():
async def send_response(request: httpx.Request) -> httpx.Response:

View file

@ -463,13 +463,14 @@ _GA_CLIENT: Final = _ClientWebSocketWithHeaders(headers=())
_BETA_CLIENT: Final = _ClientWebSocketWithHeaders(headers=((b"openai-beta", b"realtime=v1"),))
def test_azure_ga_only_protocol_comes_from_model_metadata(local_model_cost_map) -> None:
@pytest.mark.parametrize("model", ["gpt-realtime-2.1", "gpt-realtime-2", "gpt-realtime-2-2026-05-06"])
def test_azure_ga_only_protocol_comes_from_model_metadata(local_model_cost_map, model: str) -> None:
from litellm.llms.azure.realtime.handler import azure_realtime_protocol_for_client
assert (
azure_realtime_protocol_for_client(
None,
model="gpt-realtime-2.1",
model=model,
realtime_mode="realtime",
query_params=None,
websocket=_BETA_CLIENT,
@ -479,7 +480,7 @@ def test_azure_ga_only_protocol_comes_from_model_metadata(local_model_cost_map)
with pytest.raises(ValueError, match="requires the Azure OpenAI v1 Realtime API"):
azure_realtime_protocol_for_client(
"beta",
model="gpt-realtime-2.1",
model=model,
realtime_mode="realtime",
query_params=None,
websocket=_BETA_CLIENT,