From 3d7a771ea7d2a327a4cf09b5425fdfe4e3fbc69d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:53:06 -0700 Subject: [PATCH] fix(vertex_ai): apply finals before interims and refresh the token per stream --- .../audio_transcription/realtime_backend.py | 33 ++++++--- .../realtime_transformation.py | 30 ++++---- .../llms/vertex_ai/realtime/transformation.py | 23 +++++++ litellm/realtime_api/main.py | 35 ++++------ .../test_vertex_ai_realtime_backend.py | 68 ++++++++++++++++++- .../test_vertex_ai_realtime_transformation.py | 48 ++++++++++++- 6 files changed, 186 insertions(+), 51 deletions(-) diff --git a/litellm/llms/vertex_ai/audio_transcription/realtime_backend.py b/litellm/llms/vertex_ai/audio_transcription/realtime_backend.py index 87c72c193bc..0c16fea9e9f 100644 --- a/litellm/llms/vertex_ai/audio_transcription/realtime_backend.py +++ b/litellm/llms/vertex_ai/audio_transcription/realtime_backend.py @@ -80,7 +80,7 @@ class _Closed: pass -def open_speech_client(target: SpeechStreamingTarget) -> SpeechStreamingClient: +def open_speech_client(target: SpeechStreamingTarget, access_token: str) -> SpeechStreamingClient: try: from google.api_core.client_options import ClientOptions from google.cloud.speech_v2 import SpeechAsyncClient @@ -88,7 +88,7 @@ def open_speech_client(target: SpeechStreamingTarget) -> SpeechStreamingClient: except ImportError as e: raise ImportError(SPEECH_SDK_INSTALL_HINT) from e return SpeechAsyncClient( - credentials=Credentials(token=target.access_token), + credentials=Credentials(token=access_token), transport="grpc_asyncio", client_options=ClientOptions(api_endpoint=target.api_endpoint), ) @@ -157,6 +157,7 @@ class _RecognizeStream: self.speech_active: bool = False self.billed_seconds: float = 0.0 self._cancelled: bool = False + self._closed: bool = False self._task: asyncio.Task[None] | None = None async def send_audio(self, audio: bytes) -> None: @@ -170,8 +171,15 @@ class _RecognizeStream: if self._task is not None: self._task.cancel() + async def close(self) -> None: + if self._closed: + return + self._closed = True + await self._client.transport.close() + async def relay(self, outbox: "asyncio.Queue[str | _StreamFailure | _Closed]", billed_before: float) -> float: if self._cancelled: + await self.close() return 0.0 task: Final = asyncio.create_task(self._forward(outbox, billed_before)) self._task = task @@ -181,6 +189,8 @@ class _RecognizeStream: task.cancel() await asyncio.wait((task,)) raise + finally: + await self.close() return self.billed_seconds async def _forward(self, outbox: "asyncio.Queue[str | _StreamFailure | _Closed]", billed_before: float) -> None: @@ -209,7 +219,7 @@ class SpeechStreamingBackend: self, target: SpeechStreamingTarget, *, - client_factory: Callable[[SpeechStreamingTarget], SpeechStreamingClient] = open_speech_client, + client_factory: Callable[[SpeechStreamingTarget, str], SpeechStreamingClient] = open_speech_client, clock: Callable[[], float] = time.monotonic, rotation_seconds: float = STREAM_ROTATION_SECONDS, rotation_deadline_seconds: float = STREAM_ROTATION_DEADLINE_SECONDS, @@ -222,7 +232,6 @@ class SpeechStreamingBackend: self._outbox: Final[asyncio.Queue[str | _StreamFailure | _Closed]] = asyncio.Queue(maxsize=OUTBOX_SIZE) self._links: Final[asyncio.Queue[_RecognizeStream | str]] = asyncio.Queue(maxsize=_LINK_QUEUE_SIZE) self._pump: asyncio.Task[None] | None = None - self._client: SpeechStreamingClient | None = None self._config: StreamingRecognitionConfig | None = None self._turn: tuple[_RecognizeStream, ...] = () self._billed_before: float = 0.0 @@ -282,13 +291,16 @@ class SpeechStreamingBackend: if pump is not None: pump.cancel() await asyncio.wait((pump,)) - client: Final = self._client - self._client = None - if client is not None: - await client.transport.close() + await self._close_unrelayed_streams() if not self._outbox.full(): self._outbox.put_nowait(_Closed()) + async def _close_unrelayed_streams(self) -> None: + unrelayed: Final = tuple(self._links.get_nowait() for _ in range(self._links.qsize())) + for link in unrelayed: + if isinstance(link, _RecognizeStream): + await link.close() + async def _link(self, item: _RecognizeStream | str) -> None: if self._pump is None: self._pump = asyncio.create_task(self._pump_links()) @@ -333,10 +345,9 @@ class SpeechStreamingBackend: config: Final = self._config if config is None: raise RuntimeError("audio was sent before the Speech-to-Text stream was configured") - if self._client is None: - self._client = self._client_factory(self._target) + access_token: Final = await self._target.resolve_access_token() stream: Final = _RecognizeStream( - client=self._client, + client=self._client_factory(self._target, access_token), request_type=StreamingRecognizeRequest, first_request=StreamingRecognizeRequest(recognizer=self._target.recognizer, streaming_config=config), opened_at=self._clock(), diff --git a/litellm/llms/vertex_ai/audio_transcription/realtime_transformation.py b/litellm/llms/vertex_ai/audio_transcription/realtime_transformation.py index 6ec7a21a134..dab2e980fd0 100644 --- a/litellm/llms/vertex_ai/audio_transcription/realtime_transformation.py +++ b/litellm/llms/vertex_ai/audio_transcription/realtime_transformation.py @@ -1,4 +1,4 @@ -from collections.abc import Callable, Mapping +from collections.abc import Awaitable, Callable, Mapping from dataclasses import dataclass, replace from typing import Final @@ -6,7 +6,7 @@ from pydantic import JsonValue, TypeAdapter from typing_extensions import assert_never import litellm -from litellm import verbose_logger +from litellm._logging import verbose_logger from litellm._uuid import uuid from litellm.litellm_core_utils.audio_utils.utils import normalize_transcription_language_to_bcp47 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -73,7 +73,7 @@ class ChirpProtocolError(RealtimeTranscriptionProtocolError): class SpeechStreamingTarget: api_endpoint: str recognizer: str - access_token: str + resolve_access_token: Callable[[], Awaitable[str]] @dataclass(frozen=True, slots=True) @@ -249,16 +249,14 @@ class ChirpEventTransformer: finals: Final = tuple( result.transcript.strip() for result in frame.results if result.is_final and result.transcript.strip() ) - begin_events: Final = self._begin() if frame.speech_event == "begin" or interim or finals else () - interim_events: Final = self._hypothesis(interim) if interim else () + begin_events: Final = self._begin() if frame.speech_event == "begin" else () final_events: Final = tuple(event for final in finals for event in self._final(final)) + interim_events: Final = self._hypothesis(interim) if interim else () end_events: Final = self._stop() if frame.speech_event == "end" else () - return (*begin_events, *interim_events, *final_events, *end_events) + return (*begin_events, *final_events, *interim_events, *end_events) def _begin(self) -> tuple[OpenAIRealtimeEvents, ...]: - if self._turn is None: - self._turn = _Turn(item_id=self._new_item_id()) - turn: Final = self._turn + turn: Final = self._require_turn() if turn.started_emitted or not self._require_config().server_vad: return () self._turn = replace(turn, started_emitted=True) @@ -272,21 +270,23 @@ class ChirpEventTransformer: return (speech_event("input_audio_buffer.speech_stopped", turn.item_id),) def _hypothesis(self, text: str) -> tuple[OpenAIRealtimeEvents, ...]: + begin_events: Final = self._begin() turn: Final = self._require_turn() hypothesis: Final = _join_transcript(turn.committed, text) delta: Final = new_words(turn.preview, hypothesis) self._turn = replace(turn, preview=hypothesis) - return (delta_event(turn.item_id, delta),) if delta else () + return (*begin_events, delta_event(turn.item_id, delta)) if delta else begin_events def _final(self, text: str) -> tuple[OpenAIRealtimeEvents, ...]: + begin_events: Final = self._begin() turn: Final = self._require_turn() committed: Final = _join_transcript(turn.committed, text) delta: Final = new_words(turn.preview, committed) self._turn = replace(turn, committed=committed, preview=committed) delta_events: Final[tuple[OpenAIRealtimeEvents, ...]] = (delta_event(turn.item_id, delta),) if delta else () if not self._require_config().server_vad: - return delta_events - return (*delta_events, *self._complete()) + return (*begin_events, *delta_events) + return (*begin_events, *delta_events, *self._complete()) def _finish_turn(self) -> tuple[OpenAIRealtimeEvents, ...]: if self._turn is None: @@ -326,12 +326,12 @@ class VertexChirpRealtimeConfig(BaseRealtimeConfig): def __init__( self, *, - access_token: str, + resolve_access_token: Callable[[], Awaitable[str]], project: str, location: str | None, backend_factory: Callable[[SpeechStreamingTarget], RealtimeBackend] = _default_backend_factory, ) -> None: - self._access_token: Final = access_token + self._resolve_access_token: Final = resolve_access_token self._project: Final = validate_vertex_transcription_project_id(project) self._location: Final = validate_vertex_transcription_location(location, DEFAULT_SPEECH_TO_TEXT_LOCATION) self._backend_factory: Final = backend_factory @@ -357,7 +357,7 @@ class VertexChirpRealtimeConfig(BaseRealtimeConfig): SpeechStreamingTarget( api_endpoint=url, recognizer=f"projects/{self._project}/locations/{self._location}/recognizers/_", - access_token=self._access_token, + resolve_access_token=self._resolve_access_token, ) ) diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index fe59034c27b..9fed6d52f0e 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -12,10 +12,16 @@ Auth: OAuth2 Bearer token (not an API key). """ import json +from collections.abc import Awaitable, Callable from typing import Final from litellm import verbose_logger from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig +from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import ( + VertexChirpRealtimeConfig, + is_vertex_speech_to_text_model, +) +from litellm.llms.vertex_ai.vertex_llm_base import VertexBase class VertexAIRealtimeConfig(GeminiRealtimeConfig): @@ -232,3 +238,20 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): return [] return super().transform_realtime_request(message, model, session_configuration_request) + + +def vertex_realtime_config( + model: str, + *, + access_token: str, + resolve_access_token: Callable[[], Awaitable[str]], + project: str, + location: str | None, +) -> VertexAIRealtimeConfig | VertexChirpRealtimeConfig: + if is_vertex_speech_to_text_model(model): + return VertexChirpRealtimeConfig(resolve_access_token=resolve_access_token, project=project, location=location) + return VertexAIRealtimeConfig( + access_token=access_token, + project=project, + location=VertexBase.get_vertex_region(vertex_region=location, model=model), + ) diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index aed7bf15bdc..0e83edab5e1 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -38,11 +38,8 @@ from ..llms.azure.realtime.handler import AzureOpenAIRealtime, azure_realtime_pr from ..llms.bedrock.realtime.handler import BedrockRealtime from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context from ..llms.openai.realtime.handler import OpenAIRealtime -from ..llms.vertex_ai.audio_transcription.realtime_transformation import ( - VertexChirpRealtimeConfig, - is_vertex_speech_to_text_model, -) -from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig +from ..llms.vertex_ai.audio_transcription.realtime_transformation import is_vertex_speech_to_text_model +from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig, vertex_realtime_config from ..llms.vertex_ai.vertex_llm_base import VertexBase from ..llms.xai.realtime.handler import XAIRealtime from ..utils import client as wrapper_client @@ -555,9 +552,19 @@ async def _arealtime( timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS, ) - vertex_realtime_config: Final = _vertex_realtime_config( - model=model, + async def resolve_vertex_access_token() -> str: + refreshed_token, _ = await _resolve_vertex_access_token_bounded( + credentials=vertex_credentials, + project_id=resolved_project, + resolver=vertex_access_token_resolver, + timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS, + ) + return refreshed_token + + vertex_provider_config: Final = vertex_realtime_config( + model, access_token=access_token, + resolve_access_token=resolve_vertex_access_token, project=resolved_project, location=vertex_location, ) @@ -566,7 +573,7 @@ async def _arealtime( model=model, websocket=websocket, logging_obj=litellm_logging_obj, - provider_config=vertex_realtime_config, + provider_config=vertex_provider_config, api_base=dynamic_api_base or litellm_params.api_base, api_key=None, client=client, @@ -580,18 +587,6 @@ async def _arealtime( raise ValueError(f"Unsupported model: {model}") -def _vertex_realtime_config( - model: str, access_token: str, project: str, location: str | None -) -> VertexAIRealtimeConfig | VertexChirpRealtimeConfig: - if is_vertex_speech_to_text_model(model): - return VertexChirpRealtimeConfig(access_token=access_token, project=project, location=location) - return VertexAIRealtimeConfig( - access_token=access_token, - project=project, - location=vertex_llm_base.get_vertex_region(vertex_region=location, model=model), - ) - - def _is_transcription_only_realtime_model(model: str, custom_llm_provider: str) -> bool: try: model_info: Final = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) diff --git a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py b/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py index 49758e62415..d6f65c90806 100644 --- a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py +++ b/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_backend.py @@ -1,6 +1,7 @@ import asyncio import json from collections.abc import AsyncIterator, Sequence +from dataclasses import replace from datetime import timedelta from typing import Final @@ -17,10 +18,15 @@ from websockets.exceptions import ConnectionClosedError, ConnectionClosedOK from litellm.llms.vertex_ai.audio_transcription.realtime_backend import REQUEST_QUEUE_SIZE, SpeechStreamingBackend from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import SpeechStreamingTarget + +async def _static_token() -> str: + return "token" + + TARGET: Final = SpeechStreamingTarget( api_endpoint="us-speech.googleapis.com", recognizer="projects/proj-1/locations/us/recognizers/_", - access_token="token", + resolve_access_token=_static_token, ) CONFIGURE: Final = json.dumps( {"kind": "configure", "model": "chirp_3", "language_codes": ["en-US"], "sample_rate_hertz": 16_000} @@ -101,7 +107,7 @@ class _FakeSpeechClient: def _backend(client: _FakeSpeechClient, **kwargs: object) -> SpeechStreamingBackend: - return SpeechStreamingBackend(TARGET, client_factory=lambda target: client, **kwargs) + return SpeechStreamingBackend(TARGET, client_factory=lambda target, access_token: client, **kwargs) async def _recv(backend: SpeechStreamingBackend) -> dict[str, object]: @@ -346,6 +352,64 @@ async def test_rotation_is_forced_at_the_deadline_during_continuous_speech(): assert [_audio(stream) for stream in client.streams] == [[b"\x01\x01", b"\x02\x02"], [b"\x03\x03"]] +@pytest.mark.asyncio +async def test_every_stream_opens_its_own_client_with_a_freshly_resolved_token(): + now = [0.0] + tokens = iter(("token-1", "token-2")) + seen_tokens: list[str] = [] + clients = [_FakeSpeechClient([_response("first")]), _FakeSpeechClient([_response("second")])] + unopened = iter(clients) + + async def resolve_access_token() -> str: + return next(tokens) + + def open_client(target: SpeechStreamingTarget, access_token: str) -> _FakeSpeechClient: + seen_tokens.append(access_token) + return next(unopened) + + backend = SpeechStreamingBackend( + replace(TARGET, resolve_access_token=resolve_access_token), + client_factory=open_client, + clock=lambda: now[0], + rotation_seconds=240.0, + ) + async with backend: + await _configure(backend) + await backend.send(b"\x01\x01") + assert await _transcript(backend) == "first" + now[0] = 240.0 + await backend.send(b"\x02\x02") + assert await _transcript(backend) == "second" + assert clients[0].transport.closed + assert not clients[1].transport.closed + assert seen_tokens == ["token-1", "token-2"] + assert [len(client.streams) for client in clients] == [1, 1] + assert clients[1].transport.closed + + +@pytest.mark.asyncio +async def test_close_releases_a_rotated_stream_that_never_started_relaying(): + now = [0.0] + hold = asyncio.Event() + clients = [_FakeSpeechClient([_response("first"), hold]), _FakeSpeechClient([_response("never")])] + unopened = iter(clients) + backend = SpeechStreamingBackend( + TARGET, + client_factory=lambda target, access_token: next(unopened), + clock=lambda: now[0], + rotation_seconds=240.0, + ) + await _configure(backend) + await backend.send(b"\x01\x01") + assert await _transcript(backend) == "first" + now[0] = 240.0 + await backend.send(b"\x02\x02") + await asyncio.sleep(0) + assert clients[1].streams == [] + await backend.close() + assert [client.transport.closed for client in clients] == [True, True] + + @pytest.mark.asyncio async def test_discard_turn_cancels_every_stream_of_the_turn(): now = [0.0] diff --git a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_transformation.py b/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_transformation.py index fcbfad8bf12..719dd621c82 100644 --- a/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/audio_transcription/test_vertex_ai_realtime_transformation.py @@ -63,8 +63,12 @@ def _ga_session_update( ) +async def _token() -> str: + return "token" + + def _config(location: str | None = "us") -> VertexChirpRealtimeConfig: - return VertexChirpRealtimeConfig(access_token="token", project="proj-1", location=location) + return VertexChirpRealtimeConfig(resolve_access_token=_token, project="proj-1", location=location) def _configured( @@ -261,6 +265,41 @@ def test_server_vad_turn_streams_new_words_then_completes_with_usage(): assert _backend_events(config, _response(speech_event="end")) == [] +def test_server_vad_final_result_completes_before_the_interim_that_follows_it(): + config = _configured() + _backend_events(config, _response(speech_event="begin")) + events = _backend_events(config, _response(("four score", True), ("and seven", False))) + assert _types(events) == [ + DELTA, + "input_audio_buffer.speech_stopped", + COMPLETED, + "input_audio_buffer.speech_started", + DELTA, + ] + assert events[2]["transcript"] == "four score" + assert events[4]["delta"] == "and seven" + assert events[4]["item_id"] != events[2]["item_id"] + assert events[4]["item_id"] == events[3]["item_id"] + finished = _backend_events(config, _response(("and seven years", True))) + assert [(event["type"], event.get("delta", event.get("transcript"))) for event in finished] == [ + (DELTA, " years"), + ("input_audio_buffer.speech_stopped", None), + (COMPLETED, "and seven years"), + ] + assert {event["item_id"] for event in finished} == {events[4]["item_id"]} + + +def test_manual_turn_keeps_the_interim_that_follows_a_final_in_the_same_frame(): + config = _configured(turn_detection=None) + first = _backend_events(config, _response(("four score", True), ("and seven", False))) + assert [(event["type"], event["delta"]) for event in first] == [(DELTA, "four score"), (DELTA, " and seven")] + second = _backend_events(config, _response(("and seven years", True))) + assert [event["delta"] for event in second] == [" years"] + completed = _backend_events(config, VertexSpeechStreamingTurnFinished()) + assert [(event["type"], event["transcript"]) for event in completed] == [(COMPLETED, "four score and seven years")] + assert {event["item_id"] for event in (*first, *second, *completed)} == {first[0]["item_id"]} + + def test_manual_turns_complete_on_commit_without_speech_events(): config = _configured(turn_detection=None) assert _backend_events(config, _response(speech_event="begin")) == [] @@ -322,7 +361,9 @@ async def test_open_backend_targets_the_regional_speech_endpoint(): targets.append(target) return _NullBackend() - config = VertexChirpRealtimeConfig(access_token="token", project="proj-1", location=None, backend_factory=factory) + config = VertexChirpRealtimeConfig( + resolve_access_token=_token, project="proj-1", location=None, backend_factory=factory + ) url = config.get_complete_url(None, "vertex_ai/chirp_3") assert url == "us-speech.googleapis.com" assert config.validate_environment({}, MODEL, "https://" + url) == {} @@ -332,9 +373,10 @@ async def test_open_backend_targets_the_regional_speech_endpoint(): SpeechStreamingTarget( api_endpoint="us-speech.googleapis.com", recognizer="projects/proj-1/locations/us/recognizers/_", - access_token="token", + resolve_access_token=_token, ) ] + assert await targets[0].resolve_access_token() == "token" @pytest.mark.parametrize(