mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-23 00:41:40 +00:00
fix(vertex_ai): apply finals before interims and refresh the token per stream
This commit is contained in:
parent
4356fc58d8
commit
3d7a771ea7
6 changed files with 186 additions and 51 deletions
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue