fix(vertex_ai): apply finals before interims and refresh the token per stream

This commit is contained in:
mateo-berri 2026-09-18 16:53:06 -07:00
parent 4356fc58d8
commit 3d7a771ea7
6 changed files with 186 additions and 51 deletions

View file

@ -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(),

View file

@ -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,
)
)

View file

@ -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),
)

View file

@ -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)

View file

@ -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]

View file

@ -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(