mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(realtime): preserve input usage and provider routing
This commit is contained in:
parent
17dbd6ae92
commit
828707b9f2
12 changed files with 138 additions and 90 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue