mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(gemini): map OpenAI voice names to Gemini prebuilt TTS voices (#45303)
* fix(gemini): map OpenAI voice names to Gemini prebuilt TTS voices * test(gemini): type the voice mapping test helpers * fix(gemini): leave nova out of the voice mapping since Gemini accepts it as given * test(gemini): add integration cells for the OpenAI voice mapping on Gemini and Vertex TTS --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
fa7ad80a62
commit
d6db8e8744
5 changed files with 1529 additions and 2 deletions
|
|
@ -7,6 +7,7 @@ import time
|
|||
from collections.abc import Callable, Mapping, Sequence
|
||||
from copy import deepcopy
|
||||
from functools import partial
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args
|
||||
|
||||
import httpx
|
||||
|
|
@ -115,6 +116,25 @@ else:
|
|||
|
||||
SUPPORTED_REASONING_EFFORTS: Final = ("minimal", "low", "medium", "high", "none", "disable")
|
||||
|
||||
# Gemini prebuilt voices per https://ai.google.dev/gemini-api/docs/speech-generation (2026-10-01), by nearest character;
|
||||
# nova is left out since Gemini 3.x TTS models accept it as given (live check 2026-10-08)
|
||||
OPENAI_TO_GEMINI_TTS_VOICES: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"alloy": "Kore",
|
||||
"ash": "Iapetus",
|
||||
"ballad": "Algieba",
|
||||
"cedar": "Achird",
|
||||
"coral": "Sulafat",
|
||||
"echo": "Charon",
|
||||
"fable": "Umbriel",
|
||||
"marin": "Despina",
|
||||
"onyx": "Orus",
|
||||
"sage": "Vindemiatrix",
|
||||
"shimmer": "Achernar",
|
||||
"verse": "Puck",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _unsupported_reasoning_effort(reasoning_effort: str) -> UnsupportedParamsError:
|
||||
return UnsupportedParamsError(
|
||||
|
|
@ -1071,7 +1091,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
speechConfig = {
|
||||
voiceConfig: {
|
||||
prebuiltVoiceConfig: {
|
||||
voiceName: "alloy",
|
||||
voiceName: "Kore",
|
||||
}
|
||||
},
|
||||
languageCode: "en-US",
|
||||
|
|
@ -1096,7 +1116,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
speech_config: Final[SpeechConfig] = {}
|
||||
|
||||
if "voice" in value:
|
||||
prebuilt_voice_config: Final[PrebuiltVoiceConfig] = {"voiceName": value["voice"]}
|
||||
voice: Final = value["voice"]
|
||||
voice_name: Final = (
|
||||
OPENAI_TO_GEMINI_TTS_VOICES.get(voice.lower(), voice) if isinstance(voice, str) else voice
|
||||
)
|
||||
prebuilt_voice_config: Final[PrebuiltVoiceConfig] = {"voiceName": voice_name}
|
||||
voice_config: Final[VoiceConfig] = {"prebuiltVoiceConfig": prebuilt_voice_config}
|
||||
speech_config["voiceConfig"] = voice_config
|
||||
|
||||
|
|
|
|||
451
tests/integration/providers/_gemini_tts_voice_support.py
Normal file
451
tests/integration/providers/_gemini_tts_voice_support.py
Normal file
|
|
@ -0,0 +1,451 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Callable, Generator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from hashlib import sha256
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from threading import Thread
|
||||
from typing import Final
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import websockets
|
||||
import yaml
|
||||
from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.upstream import GEMINI_LIVE_PATH, ScenarioHandle, delete_scenario, register_scenario
|
||||
from integration._support.vertex import service_account_json
|
||||
from integration._support.wire import Reply, Request
|
||||
from integration.cost_calculation.cost_tracking_case import RealtimeResponse
|
||||
from pydantic import JsonValue
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
|
||||
BACKEND: Final = "gemini-3.8-flash-tts"
|
||||
GEMINI_API_KEY: Final = "synthetic-gemini-key"
|
||||
VERTEX_PROJECT: Final = "scripted-tts-project"
|
||||
VERTEX_LOCATION: Final = "us-central1"
|
||||
VERTEX_MODEL_PATH: Final = (
|
||||
f"/v1/projects/{VERTEX_PROJECT}/locations/{VERTEX_LOCATION}/publishers/google/models/{BACKEND}"
|
||||
)
|
||||
GEMINI_PREBUILT_VOICES: Final = frozenset(
|
||||
{
|
||||
"Zephyr",
|
||||
"Puck",
|
||||
"Charon",
|
||||
"Kore",
|
||||
"Fenrir",
|
||||
"Leda",
|
||||
"Orus",
|
||||
"Aoede",
|
||||
"Callirrhoe",
|
||||
"Autonoe",
|
||||
"Enceladus",
|
||||
"Iapetus",
|
||||
"Umbriel",
|
||||
"Algieba",
|
||||
"Despina",
|
||||
"Erinome",
|
||||
"Algenib",
|
||||
"Rasalgethi",
|
||||
"Laomedeia",
|
||||
"Achernar",
|
||||
"Alnilam",
|
||||
"Schedar",
|
||||
"Gacrux",
|
||||
"Pulcherrima",
|
||||
"Achird",
|
||||
"Zubenelgenubi",
|
||||
"Vindemiatrix",
|
||||
"Sadachbia",
|
||||
"Sadaltager",
|
||||
"Sulafat",
|
||||
}
|
||||
)
|
||||
ACCEPTED_VOICES: Final = GEMINI_PREBUILT_VOICES | {"nova"}
|
||||
REJECTED_PREFIX: Final = "No matching speaker voice found for name: "
|
||||
NO_VOICE: Final = "<no voiceConfig>"
|
||||
PCM: Final = bytes(range(256)) * 4
|
||||
PCM_B64: Final = base64.b64encode(PCM).decode()
|
||||
AUDIO_MIME: Final = "audio/L16;codec=pcm;rate=24000"
|
||||
PROMPT_TOKENS: Final = 5
|
||||
AUDIO_TOKENS: Final = 50
|
||||
SPEND_BY_ID_SQL: Final = 'SELECT call_type, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
|
||||
SPEND_BY_KEY_SQL: Final = 'SELECT call_type, request_id FROM "LiteLLM_SpendLogs" WHERE api_key = %s'
|
||||
|
||||
|
||||
def marker() -> str:
|
||||
return f"tts-{uuid.uuid4().hex[:12]}"
|
||||
|
||||
|
||||
def audio_param(voice: JsonValue) -> dict[str, JsonValue]:
|
||||
return {"voice": voice, "format": "pcm16"}
|
||||
|
||||
|
||||
def chat_body(model: str, voice: JsonValue, text: str, *, stream: bool = False) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": text}],
|
||||
"modalities": ["audio"],
|
||||
"audio": audio_param(voice),
|
||||
**({"stream": True} if stream else {}),
|
||||
}
|
||||
|
||||
|
||||
def received_voice(request: Request) -> JsonValue:
|
||||
return voice_in(JSON_OBJECT.validate_json(request.body))
|
||||
|
||||
|
||||
def text_in(body: Mapping[str, JsonValue]) -> str:
|
||||
contents: Final = body["contents"]
|
||||
assert isinstance(contents, list) and contents, body
|
||||
parts: Final = object_value(contents[0])["parts"]
|
||||
assert isinstance(parts, list) and parts, body
|
||||
return string_value(object_value(parts[0])["text"])
|
||||
|
||||
|
||||
def received_text(request: Request) -> str:
|
||||
return text_in(JSON_OBJECT.validate_json(request.body))
|
||||
|
||||
|
||||
def _candidate(text: str) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"content": {
|
||||
"role": "model",
|
||||
"parts": [{"text": text}, {"inlineData": {"mimeType": AUDIO_MIME, "data": PCM_B64}}],
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
}
|
||||
|
||||
|
||||
def _usage() -> dict[str, JsonValue]:
|
||||
return {
|
||||
"promptTokenCount": PROMPT_TOKENS,
|
||||
"candidatesTokenCount": AUDIO_TOKENS,
|
||||
"totalTokenCount": PROMPT_TOKENS + AUDIO_TOKENS,
|
||||
"candidatesTokensDetails": [{"modality": "AUDIO", "tokenCount": AUDIO_TOKENS}],
|
||||
}
|
||||
|
||||
|
||||
def _rejection(voice: JsonValue) -> Reply:
|
||||
error: Final = {"code": 400, "message": f"{REJECTED_PREFIX}{voice} and language: ", "status": "INVALID_ARGUMENT"}
|
||||
return Reply(status=400, body=json.dumps({"error": error}).encode())
|
||||
|
||||
|
||||
def _accepted(voice: JsonValue) -> bool:
|
||||
return voice == NO_VOICE or (isinstance(voice, str) and voice in ACCEPTED_VOICES)
|
||||
|
||||
|
||||
def gemini_peer(request: Request) -> Reply:
|
||||
assert request.method == "POST", request.method
|
||||
voice: Final = received_voice(request)
|
||||
if not _accepted(voice):
|
||||
return _rejection(voice)
|
||||
text: Final = received_text(request)
|
||||
if "streamGenerateContent" in request.target:
|
||||
first: Final = json.dumps({"candidates": [{"content": {"role": "model", "parts": [{"text": text}]}}]})
|
||||
last: Final = json.dumps({"candidates": [_candidate("")], "usageMetadata": _usage(), "modelVersion": BACKEND})
|
||||
return Reply(
|
||||
content_type="text/event-stream", chunks=(f"data: {first}\n\n".encode(), f"data: {last}\n\n".encode())
|
||||
)
|
||||
body: Final = {"candidates": [_candidate(text)], "usageMetadata": _usage(), "modelVersion": BACKEND}
|
||||
return Reply(body=json.dumps(body).encode())
|
||||
|
||||
|
||||
def gemini_deployment(scenario: Scenario, wire_url: str, **extra: JsonValue) -> str:
|
||||
return scenario.model(model=f"gemini/{BACKEND}", api_base=wire_url, api_key=GEMINI_API_KEY, **extra)
|
||||
|
||||
|
||||
def vertex_deployment(
|
||||
scenario: Scenario, wire_url: str, token_url: str, *, api_base: str | None = None, **extra: JsonValue
|
||||
) -> str:
|
||||
return scenario.model(
|
||||
model=f"vertex_ai/{BACKEND}",
|
||||
api_base=api_base or f"{wire_url}{VERTEX_MODEL_PATH}",
|
||||
api_key=None,
|
||||
vertex_project=VERTEX_PROJECT,
|
||||
vertex_location=VERTEX_LOCATION,
|
||||
vertex_credentials=service_account_json(VERTEX_PROJECT, token_url.rstrip("/")),
|
||||
**extra,
|
||||
)
|
||||
|
||||
|
||||
def openai_client(gateway: Gateway) -> openai.OpenAI:
|
||||
return openai.OpenAI(
|
||||
base_url=f"{gateway.client.base_url}/v1",
|
||||
api_key=gateway.key,
|
||||
max_retries=0,
|
||||
http_client=httpx.Client(trust_env=False, timeout=60),
|
||||
)
|
||||
|
||||
|
||||
def async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI:
|
||||
return openai.AsyncOpenAI(
|
||||
base_url=f"{gateway.client.base_url}/v1",
|
||||
api_key=gateway.key,
|
||||
max_retries=0,
|
||||
http_client=httpx.AsyncClient(trust_env=False, timeout=60),
|
||||
)
|
||||
|
||||
|
||||
def anthropic_client(gateway: Gateway) -> anthropic.Anthropic:
|
||||
return anthropic.Anthropic(
|
||||
base_url=str(gateway.client.base_url),
|
||||
api_key=gateway.key,
|
||||
max_retries=0,
|
||||
http_client=httpx.Client(trust_env=False, timeout=60),
|
||||
)
|
||||
|
||||
|
||||
def spend_row(response_id: str) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(SPEND_BY_ID_SQL, (response_id,)), lambda values: len(values) == 1, seconds=70
|
||||
)
|
||||
return rows[0]
|
||||
|
||||
|
||||
def spend_rows_for_key(key: str, count: int) -> list[dict[str, JsonValue]]:
|
||||
return eventually(
|
||||
lambda: read_rows(SPEND_BY_KEY_SQL, (sha256(key.encode()).hexdigest(),)),
|
||||
lambda rows: len(rows) == count,
|
||||
seconds=70,
|
||||
)
|
||||
|
||||
|
||||
def error_message(response: httpx.Response) -> str:
|
||||
error: Final = JSON_OBJECT.validate_json(response.content)["error"]
|
||||
return string_value(object_value(error)["message"]) if isinstance(error, dict) else string_value(error)
|
||||
|
||||
|
||||
def audio_data(payload: Mapping[str, JsonValue]) -> str:
|
||||
choices: Final = payload["choices"]
|
||||
assert isinstance(choices, list) and len(choices) == 1, payload
|
||||
message: Final = object_value(object_value(choices[0])["message"])
|
||||
return string_value(object_value(message["audio"])["data"])
|
||||
|
||||
|
||||
LIVE_BACKEND: Final = "gemini-3.8-live"
|
||||
LIVE_INPUT_TOKENS: Final = 7
|
||||
LIVE_OUTPUT_TOKENS: Final = 5
|
||||
SPEND_BY_CALL_SQL: Final = 'SELECT call_type, completion_tokens FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = %s'
|
||||
|
||||
|
||||
def voice_in(body: Mapping[str, JsonValue]) -> JsonValue:
|
||||
generation: Final = body.get("generationConfig")
|
||||
speech: Final = object_value(generation).get("speechConfig") if isinstance(generation, dict) else None
|
||||
voice_config: Final = object_value(speech).get("voiceConfig") if isinstance(speech, dict) else None
|
||||
if not isinstance(voice_config, dict):
|
||||
return NO_VOICE
|
||||
return object_value(voice_config["prebuiltVoiceConfig"])["voiceName"]
|
||||
|
||||
|
||||
def spend_row_by_call(call_id: str) -> dict[str, JsonValue]:
|
||||
assert call_id, "the response carried no x-litellm-call-id"
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(SPEND_BY_CALL_SQL, (call_id,)), lambda values: len(values) == 1, seconds=70
|
||||
)
|
||||
return rows[0]
|
||||
|
||||
|
||||
def live_turn_script() -> RealtimeResponse:
|
||||
return RealtimeResponse(
|
||||
content_type="application/x-realtime",
|
||||
events=(
|
||||
{"serverContent": {"modelTurn": {"parts": [{"text": "scripted $REQUEST_ID"}]}}},
|
||||
{
|
||||
"serverContent": {"turnComplete": True},
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": LIVE_INPUT_TOKENS,
|
||||
"responseTokenCount": LIVE_OUTPUT_TOKENS,
|
||||
"totalTokenCount": LIVE_INPUT_TOKENS + LIVE_OUTPUT_TOKENS,
|
||||
},
|
||||
},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def live_scenario(scenario: Scenario) -> ScenarioHandle:
|
||||
handle: Final = register_scenario(f"tts-live-{uuid.uuid4().hex[:12]}", live_turn_script())
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
return handle
|
||||
|
||||
|
||||
def live_deployment(scenario: Scenario, project: str, upstream_url: str) -> str:
|
||||
return scenario.model(
|
||||
model=f"vertex_ai/{LIVE_BACKEND}",
|
||||
api_base=upstream_url.rstrip("/"),
|
||||
api_key=None,
|
||||
vertex_project=project,
|
||||
vertex_location=VERTEX_LOCATION,
|
||||
vertex_credentials=service_account_json(project, upstream_url.rstrip("/")),
|
||||
model_info={"mode": "realtime"},
|
||||
)
|
||||
|
||||
|
||||
def _user_turn() -> str:
|
||||
item: Final = {"type": "message", "role": "user", "content": [{"type": "input_text", "text": "say the line"}]}
|
||||
return json.dumps({"type": "conversation.item.create", "item": item})
|
||||
|
||||
|
||||
async def _live_frames(socket: ClientConnection, voice: str) -> AsyncIterator[dict[str, JsonValue]]:
|
||||
try:
|
||||
first: Final = JSON_OBJECT.validate_json(await socket.recv())
|
||||
yield first
|
||||
if first.get("type") != "session.created":
|
||||
return
|
||||
await socket.send(json.dumps({"type": "session.update", "session": {"voice": voice}}))
|
||||
await socket.send(_user_turn())
|
||||
async for message in socket:
|
||||
event: Final = JSON_OBJECT.validate_json(message)
|
||||
yield event
|
||||
if event.get("type") == "response.done":
|
||||
return
|
||||
except ConnectionClosed as closed:
|
||||
yield {"type": "closed", "code": 1006 if closed.rcvd is None else closed.rcvd.code}
|
||||
|
||||
|
||||
async def _live_session(ws_base: str, model: str, key: str, voice: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
headers: Final = {"Authorization": f"Bearer {key}"}
|
||||
async with websockets.connect(f"{ws_base}/v1/realtime?model={model}", additional_headers=headers) as socket:
|
||||
return tuple([frame async for frame in _live_frames(socket, voice)])
|
||||
|
||||
|
||||
def live_turn(gateway: Gateway, model: str, key: str, voice: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
ws_base: Final = str(gateway.client.base_url).rstrip("/").replace("http://", "ws://")
|
||||
return asyncio.run(asyncio.wait_for(_live_session(ws_base, model, key, voice), 60))
|
||||
|
||||
|
||||
def frame_types(frames: tuple[dict[str, JsonValue], ...]) -> tuple[str, ...]:
|
||||
return tuple(string_value(frame["type"]) for frame in frames)
|
||||
|
||||
|
||||
def live_setups(gateway: Gateway, project: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
|
||||
observed: Final = tuple(map(object_value, upstream.get("/__observations").json()["requests"]))
|
||||
frames: Final = tuple(
|
||||
object_value(request["body"])
|
||||
for request in observed
|
||||
if request["api_key"] == project and request["method"] == "WEBSOCKET_FRAME"
|
||||
)
|
||||
assert all(request["path"] == GEMINI_LIVE_PATH for request in observed if request["api_key"] == project)
|
||||
return tuple(object_value(frame["setup"]) for frame in frames if "setup" in frame)
|
||||
|
||||
|
||||
HEALTH_DEFAULT: Final = "tts-health-default-voice"
|
||||
HEALTH_NOVA: Final = "tts-health-nova"
|
||||
HEALTH_FABLE: Final = "tts-health-fable"
|
||||
|
||||
|
||||
def _health_model(name: str, **model_info: JsonValue) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"model_name": name,
|
||||
"litellm_params": {"model": f"gemini/{BACKEND}", "api_key": GEMINI_API_KEY},
|
||||
"model_info": {"mode": "audio_speech", **model_info},
|
||||
}
|
||||
|
||||
|
||||
def owned_config(directory: Path, model_list: Sequence[JsonValue]) -> Path:
|
||||
config: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
|
||||
settings: Final = object_value(config["litellm_settings"])
|
||||
cache_params: Final = {**object_value(settings["cache_params"]), "namespace": f"tts-{uuid.uuid4().hex}"}
|
||||
merged: Final = {
|
||||
**config,
|
||||
"model_list": list(model_list),
|
||||
"litellm_settings": {**settings, "cache_params": cache_params},
|
||||
"router_settings": {**object_value(config["router_settings"]), "num_retries": 0},
|
||||
}
|
||||
path: Final = directory / f"gemini-tts-{uuid.uuid4().hex}.yaml"
|
||||
path.write_text(yaml.safe_dump(merged))
|
||||
return path
|
||||
|
||||
|
||||
def health_config(directory: Path) -> Path:
|
||||
return owned_config(
|
||||
directory,
|
||||
(
|
||||
_health_model(HEALTH_DEFAULT),
|
||||
_health_model(HEALTH_NOVA, health_check_voice="nova"),
|
||||
_health_model(HEALTH_FABLE, health_check_voice="fable"),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def speech_deployment(scenario: Scenario, wire_url: str) -> str:
|
||||
return scenario.model(
|
||||
model=f"gemini/{BACKEND}", api_base=wire_url, api_key=GEMINI_API_KEY, model_info={"mode": "audio_speech"}
|
||||
)
|
||||
|
||||
|
||||
def wav_payload(body: bytes) -> bytes:
|
||||
assert body[:4] == b"RIFF" and body[8:12] == b"WAVE", body[:12]
|
||||
return body[44:]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ChunkedPeer:
|
||||
url: str
|
||||
received: SimpleQueue[Request]
|
||||
|
||||
def drain(self) -> tuple[Request, ...]:
|
||||
return tuple(self.received.get_nowait() for _ in range(self.received.qsize()))
|
||||
|
||||
|
||||
def _read_chunked(handler: BaseHTTPRequestHandler) -> bytes:
|
||||
def chunks() -> Generator[bytes, None, None]:
|
||||
while True:
|
||||
size: Final = int(handler.rfile.readline().split(b";", 1)[0].strip(), 16)
|
||||
if size == 0:
|
||||
handler.rfile.readline()
|
||||
return
|
||||
yield handler.rfile.read(size)
|
||||
handler.rfile.readline()
|
||||
|
||||
return b"".join(chunks())
|
||||
|
||||
|
||||
@contextmanager
|
||||
def chunked_peer(respond: Callable[[Request], Reply]) -> Generator[ChunkedPeer, None, None]:
|
||||
"""Owned TCP peer like ``wire_server`` that also reads a ``transfer-encoding: chunked`` upload body,
|
||||
the shape the Vertex files route streams a batch file to GCS in."""
|
||||
received: Final[SimpleQueue[Request]] = SimpleQueue()
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_POST(self) -> None:
|
||||
chunked: Final = self.headers.get("transfer-encoding", "").lower() == "chunked"
|
||||
body: Final = (
|
||||
_read_chunked(self) if chunked else self.rfile.read(int(self.headers.get("content-length", "0")))
|
||||
)
|
||||
request: Final = Request(
|
||||
self.command, self.path, {name.lower(): value for name, value in self.headers.items()}, body
|
||||
)
|
||||
received.put(request)
|
||||
reply: Final = respond(request)
|
||||
self.send_response(reply.status)
|
||||
self.send_header("content-type", reply.content_type)
|
||||
self.send_header("content-length", str(len(reply.body)))
|
||||
self.send_header("connection", "close")
|
||||
self.end_headers()
|
||||
self.wfile.write(reply.body)
|
||||
|
||||
def log_message(self, format: str, *args: object) -> None:
|
||||
return
|
||||
|
||||
server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
|
||||
thread: Final = Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
yield ChunkedPeer(f"http://127.0.0.1:{server.server_address[1]}", received)
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join(timeout=5)
|
||||
|
|
@ -0,0 +1,354 @@
|
|||
"""The speech route, the audio_speech health check and deferred Vertex Live setups map OpenAI voices.
|
||||
|
||||
The speech bridge calls the Gemini chat path without the deployment's ``api_base`` (tracked separately),
|
||||
so these rows boot their own proxy with ``GEMINI_API_BASE`` pointed at the scripted peer, and with
|
||||
``LITELLM_GEMINI_LIVE_DEFER_SETUP`` so a client ``session.update`` voice builds the Vertex Live setup.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
from collections.abc import Callable, Iterator
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
from integration._support.client import (
|
||||
JSON_OBJECT,
|
||||
Gateway,
|
||||
eventually,
|
||||
gateway_from_environment,
|
||||
object_value,
|
||||
string_value,
|
||||
)
|
||||
from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from integration.providers._gemini_tts_voice_support import (
|
||||
BACKEND,
|
||||
GEMINI_API_KEY,
|
||||
HEALTH_DEFAULT,
|
||||
HEALTH_FABLE,
|
||||
HEALTH_NOVA,
|
||||
NO_VOICE,
|
||||
PCM,
|
||||
REJECTED_PREFIX,
|
||||
async_openai_client,
|
||||
frame_types,
|
||||
gemini_peer,
|
||||
health_config,
|
||||
live_deployment,
|
||||
live_scenario,
|
||||
live_setups,
|
||||
live_turn,
|
||||
marker,
|
||||
openai_client,
|
||||
received_text,
|
||||
received_voice,
|
||||
speech_deployment,
|
||||
spend_row_by_call,
|
||||
voice_in,
|
||||
wav_payload,
|
||||
)
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
|
||||
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
BURST: Final = 12
|
||||
VOICES: Final = ("alloy", "fable", "Kore")
|
||||
MAPPED: Final = MappingProxyType({"alloy": "Kore", "fable": "Umbriel", "Kore": "Kore"})
|
||||
|
||||
|
||||
def _overrides(wire_url: str) -> dict[str, str]:
|
||||
return {"GEMINI_API_BASE": wire_url, "LITELLM_GEMINI_LIVE_DEFER_SETUP": "true"}
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def shared_wire() -> Iterator[Wire]:
|
||||
with wire_server(gemini_peer) as wire:
|
||||
yield wire
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def owned(shared_wire: Wire, tmp_path_factory: pytest.TempPathFactory) -> Iterator[OwnedProxy]:
|
||||
directory: Final = tmp_path_factory.mktemp("gemini-tts")
|
||||
with gateway_from_environment() as gateway:
|
||||
config: Final = health_config(directory)
|
||||
with owned_proxy_process(gateway, directory, _overrides(shared_wire.url), config=config, workers=2) as proxy:
|
||||
yield proxy
|
||||
|
||||
|
||||
def _speech_body(model: str, text: str, voice: str, **extra: JsonValue) -> dict[str, JsonValue]:
|
||||
return {"model": model, "input": text, "voice": voice, **extra}
|
||||
|
||||
|
||||
def _received(wire: Wire, text: str) -> Request:
|
||||
matching: Final = tuple(request for request in wire.drain() if received_text(request) == text)
|
||||
assert len(matching) == 1, [request.target for request in matching]
|
||||
return matching[0]
|
||||
|
||||
|
||||
def _assert_speech_logged(call_id: str) -> None:
|
||||
row: Final = spend_row_by_call(call_id)
|
||||
assert row["call_type"] == "aspeech", row
|
||||
|
||||
|
||||
def test_speech_route_maps_alloy_and_answers_wav(owned: OwnedProxy, shared_wire: Wire) -> None:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = speech_deployment(scenario, shared_wire.url)
|
||||
text: Final = marker()
|
||||
response: Final = owned.gateway.request("POST", "/v1/audio/speech", _speech_body(model, text, "alloy"))
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.headers["content-type"].startswith("audio/wav"), response.headers
|
||||
assert wav_payload(response.content) == PCM
|
||||
assert received_voice(_received(shared_wire, text)) == "Kore"
|
||||
_assert_speech_logged(response.headers.get("x-litellm-call-id", ""))
|
||||
|
||||
|
||||
def test_speech_sdk_maps_alloy_and_answers_wav(owned: OwnedProxy, shared_wire: Wire) -> None:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = speech_deployment(scenario, shared_wire.url)
|
||||
text: Final = marker()
|
||||
raw: Final = openai_client(owned.gateway).audio.speech.with_raw_response.create(
|
||||
model=model, input=text, voice="alloy"
|
||||
)
|
||||
assert wav_payload(raw.content) == PCM
|
||||
assert received_voice(_received(shared_wire, text)) == "Kore"
|
||||
_assert_speech_logged(raw.headers.get("x-litellm-call-id", ""))
|
||||
|
||||
|
||||
async def _async_speech(gateway: Gateway, model: str, text: str) -> tuple[bytes, str, str]:
|
||||
client: Final = async_openai_client(gateway)
|
||||
try:
|
||||
raw: Final = await client.audio.speech.with_raw_response.create(
|
||||
model=model, input=text, voice="Alloy", response_format="pcm"
|
||||
)
|
||||
return raw.content, raw.headers.get("content-type", ""), raw.headers.get("x-litellm-call-id", "")
|
||||
finally:
|
||||
await client.close()
|
||||
|
||||
|
||||
def test_async_speech_sdk_maps_a_capitalised_alloy_and_answers_raw_pcm(owned: OwnedProxy, shared_wire: Wire) -> None:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = speech_deployment(scenario, shared_wire.url)
|
||||
text: Final = marker()
|
||||
content, content_type, call_id = asyncio.run(_async_speech(owned.gateway, model, text))
|
||||
assert content == PCM, content[:16]
|
||||
assert content_type.startswith("audio/pcm"), content_type
|
||||
assert received_voice(_received(shared_wire, text)) == "Kore"
|
||||
_assert_speech_logged(call_id)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("voice", ["nova", "Kore"])
|
||||
def test_speech_route_forwards_gemini_voice_names_verbatim(owned: OwnedProxy, shared_wire: Wire, voice: str) -> None:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = speech_deployment(scenario, shared_wire.url)
|
||||
text: Final = marker()
|
||||
response: Final = owned.gateway.request("POST", "/v1/audio/speech", _speech_body(model, text, voice))
|
||||
assert response.status_code == 200, response.text
|
||||
assert wav_payload(response.content) == PCM
|
||||
assert received_voice(_received(shared_wire, text)) == voice
|
||||
_assert_speech_logged(response.headers.get("x-litellm-call-id", ""))
|
||||
|
||||
|
||||
def _health(gateway: Gateway, model: str) -> dict[str, JsonValue]:
|
||||
response: Final = gateway.request("GET", "/health", params={"model": model})
|
||||
return JSON_OBJECT.validate_json(response.content)
|
||||
|
||||
|
||||
def _counts(report: dict[str, JsonValue]) -> tuple[JsonValue, JsonValue]:
|
||||
return report["healthy_count"], report["unhealthy_count"]
|
||||
|
||||
|
||||
def _only_voice(wire: Wire) -> JsonValue:
|
||||
requests: Final = wire.drain()
|
||||
assert len(requests) == 1, [request.target for request in requests]
|
||||
return received_voice(requests[0])
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "mapped"), [(HEALTH_DEFAULT, "Kore"), (HEALTH_NOVA, "nova"), (HEALTH_FABLE, "Umbriel")]
|
||||
)
|
||||
def test_audio_speech_health_check_reaches_the_peer_with_a_gemini_voice(
|
||||
owned: OwnedProxy, shared_wire: Wire, model: str, mapped: str
|
||||
) -> None:
|
||||
shared_wire.drain()
|
||||
report: Final = _health(owned.gateway, model)
|
||||
assert _counts(report) == (1, 0), report
|
||||
assert _only_voice(shared_wire) == mapped
|
||||
|
||||
|
||||
def test_unknown_health_check_voice_reports_the_vendor_rejection(owned: OwnedProxy, shared_wire: Wire) -> None:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model=f"gemini/{BACKEND}",
|
||||
api_base=shared_wire.url,
|
||||
api_key=GEMINI_API_KEY,
|
||||
model_info={"mode": "audio_speech", "health_check_voice": "Custom-Voice"},
|
||||
)
|
||||
shared_wire.drain()
|
||||
report: Final = _health(owned.gateway, model)
|
||||
assert _counts(report) == (0, 1), report
|
||||
unhealthy: Final = report["unhealthy_endpoints"]
|
||||
assert isinstance(unhealthy, list) and len(unhealthy) == 1, report
|
||||
assert f"{REJECTED_PREFIX}Custom-Voice" in string_value(object_value(unhealthy[0])["error"]), report
|
||||
assert _only_voice(shared_wire) == "Custom-Voice"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("voice", "expected"), [("fable", "Umbriel"), ("alloy", NO_VOICE), ("Kore", "Kore")], ids=["fable", "alloy", "Kore"]
|
||||
)
|
||||
def test_deferred_vertex_live_setup_carries_the_session_voice(owned: OwnedProxy, voice: str, expected: str) -> None:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
handle: Final = live_scenario(scenario)
|
||||
key: Final = scenario.key()
|
||||
model: Final = live_deployment(scenario, handle.scenario_id, owned.gateway.upstream_url)
|
||||
frames: Final = live_turn(owned.gateway, model, key, voice)
|
||||
types: Final = frame_types(frames)
|
||||
assert types[0] == "session.created" and types[-1] == "response.done", frames
|
||||
setups: Final = live_setups(owned.gateway, handle.scenario_id)
|
||||
assert len(setups) == 1, setups
|
||||
assert voice_in(setups[0]) == expected, setups[0]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Sent:
|
||||
voice: str
|
||||
text: str
|
||||
status: int
|
||||
body: bytes
|
||||
call_id: str
|
||||
|
||||
|
||||
async def _fire(url: str, key: str, model: str, *, tolerate_transport_errors: bool = False) -> tuple[Sent, ...]:
|
||||
async def one(client: httpx.AsyncClient, index: int) -> Sent:
|
||||
voice: Final = VOICES[index % len(VOICES)]
|
||||
text: Final = marker()
|
||||
try:
|
||||
response: Final = await client.post(
|
||||
"/v1/audio/speech", json=_speech_body(model, text, voice), headers={"Authorization": f"Bearer {key}"}
|
||||
)
|
||||
except httpx.TransportError as error:
|
||||
if not tolerate_transport_errors:
|
||||
raise
|
||||
return Sent(voice, text, 0, repr(error).encode(), "")
|
||||
return Sent(voice, text, response.status_code, response.content, response.headers.get("x-litellm-call-id", ""))
|
||||
|
||||
async with httpx.AsyncClient(base_url=url, timeout=60, trust_env=False) as client:
|
||||
return tuple(await asyncio.gather(*(one(client, index) for index in range(BURST))))
|
||||
|
||||
|
||||
def _held_peer(release: threading.Event) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert release.wait(timeout=120), "Held peer was never released"
|
||||
return gemini_peer(request)
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _held_connections(pid: int, wire_url: str) -> int:
|
||||
port: Final = urlsplit(wire_url).port
|
||||
return sum(
|
||||
1
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
|
||||
)
|
||||
|
||||
|
||||
def _by_text(received: tuple[Request, ...]) -> dict[str, tuple[Request, ...]]:
|
||||
texts: Final = {received_text(request) for request in received}
|
||||
return {text: tuple(request for request in received if received_text(request) == text) for text in texts}
|
||||
|
||||
|
||||
def _assert_served(items: tuple[Sent, ...], received: tuple[Request, ...]) -> None:
|
||||
by_text: Final = _by_text(received)
|
||||
for item in items:
|
||||
assert item.status == 200, (item.voice, item.body[:200])
|
||||
assert wav_payload(item.body) == PCM, item.text
|
||||
assert len(by_text.get(item.text, ())) == 1, item.text
|
||||
assert received_voice(by_text[item.text][0]) == MAPPED[item.voice], item.text
|
||||
_assert_speech_logged(item.call_id)
|
||||
|
||||
|
||||
def _worker_pids(proxy: OwnedProxy) -> tuple[int, ...]:
|
||||
return eventually(
|
||||
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(proxy.log.read_text())),
|
||||
lambda pids: len(pids) == 2,
|
||||
seconds=graceful_stop_seconds(),
|
||||
)
|
||||
|
||||
|
||||
def test_worker_sigkill_mid_speech_burst_leaves_the_sibling_serving(gateway: Gateway, tmp_path: Path) -> None:
|
||||
release: Final = threading.Event()
|
||||
with wire_server(_held_peer(release)) as wire:
|
||||
config: Final = health_config(tmp_path)
|
||||
with owned_proxy_process(gateway, tmp_path, _overrides(wire.url), config=config, workers=2) as proxy:
|
||||
workers: Final = _worker_pids(proxy)
|
||||
url: Final = str(proxy.gateway.client.base_url)
|
||||
with proxy.gateway.scenario() as scenario:
|
||||
model: Final = speech_deployment(scenario, wire.url)
|
||||
loop: Final = asyncio.new_event_loop()
|
||||
burst: Final = loop.run_in_executor(
|
||||
None, lambda: asyncio.run(_fire(url, proxy.gateway.key, model, tolerate_transport_errors=True))
|
||||
)
|
||||
try:
|
||||
eventually(lambda: wire.received.qsize(), lambda size: size >= BURST, 90)
|
||||
held_by: Final = MappingProxyType({pid: _held_connections(pid, wire.url) for pid in workers})
|
||||
victim: Final = max(workers, key=held_by.__getitem__)
|
||||
os.kill(victim, signal.SIGKILL)
|
||||
finally:
|
||||
release.set()
|
||||
served: Final = loop.run_until_complete(burst)
|
||||
loop.close()
|
||||
during: Final = wire.drain()
|
||||
after: Final = asyncio.run(_fire(url, proxy.gateway.key, model))
|
||||
after_received: Final = wire.drain()
|
||||
assert sum(held_by.values()) == BURST, held_by
|
||||
assert held_by[victim] > 0, held_by
|
||||
completed: Final = tuple(item for item in served if item.status == 200)
|
||||
assert len(completed) == BURST - held_by[victim], (len(completed), held_by)
|
||||
assert all(len(requests) == 1 for requests in _by_text(during).values())
|
||||
_assert_served(completed, during)
|
||||
_assert_served(after, after_received)
|
||||
|
||||
|
||||
def test_proxy_restart_mid_speech_burst_serves_the_next_burst_after_the_reboot(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
release: Final = threading.Event()
|
||||
with wire_server(_held_peer(release)) as wire, gateway.scenario() as scenario:
|
||||
config: Final = health_config(tmp_path)
|
||||
model: Final = speech_deployment(scenario, wire.url)
|
||||
with owned_proxy_process(gateway, tmp_path, _overrides(wire.url), config=config, workers=2) as first:
|
||||
url: Final = str(first.gateway.client.base_url)
|
||||
loop: Final = asyncio.new_event_loop()
|
||||
burst: Final = loop.run_in_executor(
|
||||
None, lambda: asyncio.run(_fire(url, first.gateway.key, model, tolerate_transport_errors=True))
|
||||
)
|
||||
try:
|
||||
eventually(lambda: wire.received.qsize(), lambda size: size >= BURST, 90)
|
||||
first.process.terminate()
|
||||
finally:
|
||||
release.set()
|
||||
served: Final = loop.run_until_complete(burst)
|
||||
loop.close()
|
||||
during: Final = wire.drain()
|
||||
with owned_proxy_process(gateway, tmp_path, _overrides(wire.url), config=config, workers=2) as second:
|
||||
after: Final = asyncio.run(_fire(str(second.gateway.client.base_url), second.gateway.key, model))
|
||||
after_received: Final = wire.drain()
|
||||
completed: Final = tuple(item for item in served if item.status == 200)
|
||||
assert all(len(requests) == 1 for requests in _by_text(during).values())
|
||||
_assert_served(completed, during)
|
||||
for item in served:
|
||||
if item.status != 200:
|
||||
assert item.status == 0 or b'"error"' in item.body or item.body == b"", (item.status, item.body[:200])
|
||||
_assert_served(after, after_received)
|
||||
|
|
@ -0,0 +1,652 @@
|
|||
"""Gemini and Vertex TTS deployments map OpenAI voice names onto Gemini prebuilt voices.
|
||||
|
||||
The scripted peer on 127.0.0.1 answers an audio part when the ``voiceName`` it receives is one of
|
||||
Gemini's prebuilt voices and the vendor's ``No matching speaker voice found`` 400 otherwise, so a row
|
||||
is green only when the voice the peer received is the mapped one. Every row runs through the rig proxy.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import dataclasses
|
||||
import json
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from integration.providers._gemini_tts_voice_support import (
|
||||
AUDIO_TOKENS,
|
||||
BACKEND,
|
||||
GEMINI_API_KEY,
|
||||
NO_VOICE,
|
||||
PCM_B64,
|
||||
REJECTED_PREFIX,
|
||||
VERTEX_MODEL_PATH,
|
||||
anthropic_client,
|
||||
async_openai_client,
|
||||
audio_data,
|
||||
audio_param,
|
||||
chat_body,
|
||||
chunked_peer,
|
||||
error_message,
|
||||
frame_types,
|
||||
gemini_deployment,
|
||||
gemini_peer,
|
||||
live_deployment,
|
||||
live_scenario,
|
||||
live_setups,
|
||||
live_turn,
|
||||
marker,
|
||||
openai_client,
|
||||
received_text,
|
||||
received_voice,
|
||||
spend_row_by_call,
|
||||
text_in,
|
||||
vertex_deployment,
|
||||
voice_in,
|
||||
)
|
||||
from pydantic import JsonValue
|
||||
|
||||
pytestmark: Final = pytest.mark.timeout(180)
|
||||
|
||||
BUCKET: Final = "integration-tts-bucket"
|
||||
FIVE_KB: Final = "v" * 5120
|
||||
BURST: Final = 24
|
||||
SLOW_BURST: Final = 12
|
||||
OUTAGE_AFTER: Final = 6
|
||||
SLOW_SECONDS: Final = 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Answer:
|
||||
status: int
|
||||
text: str
|
||||
response_id: str
|
||||
audio: str | None
|
||||
content: str
|
||||
call_id: str
|
||||
|
||||
|
||||
ChatClient = Callable[..., Answer]
|
||||
|
||||
|
||||
def _failed(status: int, text: str) -> Answer:
|
||||
return Answer(status, text, "", None, "", "")
|
||||
|
||||
|
||||
def _delta_text(chunk: Mapping[str, JsonValue]) -> str:
|
||||
choices: Final = chunk["choices"]
|
||||
if not isinstance(choices, list) or not choices:
|
||||
return ""
|
||||
content: Final = object_value(object_value(choices[0])["delta"]).get("content")
|
||||
return content if isinstance(content, str) else ""
|
||||
|
||||
|
||||
def _message_text(body: Mapping[str, JsonValue]) -> str:
|
||||
choices: Final = body["choices"]
|
||||
assert isinstance(choices, list) and len(choices) == 1, body
|
||||
content: Final = object_value(object_value(choices[0])["message"]).get("content")
|
||||
return content if isinstance(content, str) else ""
|
||||
|
||||
|
||||
def _sse_chunks(text: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
lines: Final = tuple(line[6:] for line in text.splitlines() if line.startswith("data: "))
|
||||
return tuple(JSON_OBJECT.validate_json(line) for line in lines if line != "[DONE]")
|
||||
|
||||
|
||||
def _error_text(response: httpx.Response) -> str:
|
||||
response.read()
|
||||
return response.text
|
||||
|
||||
|
||||
async def _async_error_text(response: httpx.Response) -> str:
|
||||
await response.aread()
|
||||
return response.text
|
||||
|
||||
|
||||
def _httpx_chat(gateway: Gateway, model: str, voice: JsonValue, text: str, *, stream: bool) -> Answer:
|
||||
response: Final = gateway.request("POST", "/v1/chat/completions", chat_body(model, voice, text, stream=stream))
|
||||
call_id: Final = response.headers.get("x-litellm-call-id", "")
|
||||
if response.status_code != 200:
|
||||
return _failed(response.status_code, response.text)
|
||||
if stream:
|
||||
chunks: Final = _sse_chunks(response.text)
|
||||
content: Final = "".join(_delta_text(chunk) for chunk in chunks)
|
||||
return Answer(200, response.text, "", None, content, call_id)
|
||||
body: Final = JSON_OBJECT.validate_json(response.content)
|
||||
return Answer(200, response.text, string_value(body["id"]), audio_data(body), _message_text(body), call_id)
|
||||
|
||||
|
||||
def _sdk_chat(gateway: Gateway, model: str, voice: str, text: str, *, stream: bool) -> Answer:
|
||||
client: Final = openai_client(gateway)
|
||||
messages: Final = [{"role": "user", "content": text}]
|
||||
try:
|
||||
if stream:
|
||||
raw: Final = client.chat.completions.with_raw_response.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
modalities=["audio"],
|
||||
audio={"voice": voice, "format": "pcm16"},
|
||||
stream=True,
|
||||
)
|
||||
chunks: Final = list(raw.parse())
|
||||
content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
|
||||
return Answer(200, "", "", None, content, raw.headers.get("x-litellm-call-id", ""))
|
||||
completed: Final = client.chat.completions.with_raw_response.create(
|
||||
model=model, messages=messages, modalities=["audio"], audio={"voice": voice, "format": "pcm16"}
|
||||
)
|
||||
completion: Final = completed.parse()
|
||||
message: Final = completion.choices[0].message
|
||||
audio: Final = message.audio.data if message.audio is not None else None
|
||||
return Answer(
|
||||
200,
|
||||
completion.model_dump_json(),
|
||||
completion.id,
|
||||
audio,
|
||||
message.content or "",
|
||||
completed.headers.get("x-litellm-call-id", ""),
|
||||
)
|
||||
except openai.APIStatusError as error:
|
||||
return _failed(error.status_code, _error_text(error.response))
|
||||
|
||||
|
||||
async def _async_sdk_chat(gateway: Gateway, model: str, voice: str, text: str, *, stream: bool) -> Answer:
|
||||
client: Final = async_openai_client(gateway)
|
||||
messages: Final = [{"role": "user", "content": text}]
|
||||
try:
|
||||
if stream:
|
||||
raw: Final = await client.chat.completions.with_raw_response.create(
|
||||
model=model,
|
||||
messages=messages,
|
||||
modalities=["audio"],
|
||||
audio={"voice": voice, "format": "pcm16"},
|
||||
stream=True,
|
||||
)
|
||||
chunks: Final = [chunk async for chunk in raw.parse()]
|
||||
content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices)
|
||||
return Answer(200, "", "", None, content, raw.headers.get("x-litellm-call-id", ""))
|
||||
completed: Final = await client.chat.completions.with_raw_response.create(
|
||||
model=model, messages=messages, modalities=["audio"], audio={"voice": voice, "format": "pcm16"}
|
||||
)
|
||||
completion: Final = completed.parse()
|
||||
message: Final = completion.choices[0].message
|
||||
audio: Final = message.audio.data if message.audio is not None else None
|
||||
return Answer(
|
||||
200,
|
||||
completion.model_dump_json(),
|
||||
completion.id,
|
||||
audio,
|
||||
message.content or "",
|
||||
completed.headers.get("x-litellm-call-id", ""),
|
||||
)
|
||||
except openai.APIStatusError as error:
|
||||
return _failed(error.status_code, await _async_error_text(error.response))
|
||||
finally:
|
||||
await client.close()
|
||||
|
||||
|
||||
def _sdk_async_chat(gateway: Gateway, model: str, voice: str, text: str, *, stream: bool) -> Answer:
|
||||
return asyncio.run(_async_sdk_chat(gateway, model, voice, text, stream=stream))
|
||||
|
||||
|
||||
CLIENTS: Final[Mapping[str, ChatClient]] = {"httpx": _httpx_chat, "sdk": _sdk_chat, "sdk-async": _sdk_async_chat}
|
||||
|
||||
|
||||
def _deployment(scenario: Scenario, provider: str, wire_url: str, upstream_url: str, **extra: JsonValue) -> str:
|
||||
if provider == "gemini":
|
||||
return gemini_deployment(scenario, wire_url, **extra)
|
||||
return vertex_deployment(scenario, wire_url, upstream_url, **extra)
|
||||
|
||||
|
||||
def _received(wire: Wire, text: str) -> Request:
|
||||
matching: Final = tuple(request for request in wire.drain() if received_text(request) == text)
|
||||
assert len(matching) == 1, [request.target for request in matching]
|
||||
return matching[0]
|
||||
|
||||
|
||||
def _assert_target(request: Request, provider: str, *, stream: bool) -> None:
|
||||
suffix: Final = ":streamGenerateContent?alt=sse" if stream else ":generateContent"
|
||||
if provider == "gemini":
|
||||
assert request.target == f"/models/{BACKEND}{suffix}", request.target
|
||||
assert request.headers["x-goog-api-key"] == GEMINI_API_KEY, request.headers
|
||||
return
|
||||
assert request.target == f"{VERTEX_MODEL_PATH}{suffix}", request.target
|
||||
assert request.headers["authorization"] == "Bearer scripted-token", request.headers
|
||||
|
||||
|
||||
def _assert_logged(answer: Answer) -> None:
|
||||
row: Final = spend_row_by_call(answer.call_id)
|
||||
assert (row["call_type"], row["completion_tokens"]) == ("acompletion", AUDIO_TOKENS), row
|
||||
|
||||
|
||||
MAPPED: Final = (
|
||||
pytest.param("gemini", "httpx", False, "alloy", "Kore", id="gemini-httpx-alloy"),
|
||||
pytest.param("vertex", "httpx", False, "alloy", "Kore", id="vertex-httpx-alloy"),
|
||||
pytest.param("gemini", "sdk", False, "ALLOY", "Kore", id="gemini-sdk-ALLOY"),
|
||||
pytest.param("gemini", "sdk-async", False, "fable", "Umbriel", id="gemini-sdk-async-fable"),
|
||||
pytest.param("vertex", "sdk", False, "onyx", "Orus", id="vertex-sdk-onyx"),
|
||||
pytest.param("gemini", "httpx", True, "echo", "Charon", id="gemini-httpx-stream-echo"),
|
||||
pytest.param("vertex", "sdk", True, "onyx", "Orus", id="vertex-sdk-stream-onyx"),
|
||||
pytest.param("vertex", "sdk-async", True, "alloy", "Kore", id="vertex-sdk-async-stream-alloy"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("provider", "client", "stream", "voice", "mapped"), MAPPED)
|
||||
def test_openai_voice_reaches_the_peer_as_the_gemini_voice(
|
||||
gateway: Gateway, provider: str, client: str, stream: bool, voice: str, mapped: str
|
||||
) -> None:
|
||||
with wire_server(gemini_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, provider, wire.url, gateway.upstream_url)
|
||||
text: Final = marker()
|
||||
answer: Final = CLIENTS[client](gateway, model, voice, text, stream=stream)
|
||||
assert answer.status == 200, answer.text
|
||||
request: Final = _received(wire, text)
|
||||
_assert_target(request, provider, stream=stream)
|
||||
assert received_voice(request) == mapped, request.body
|
||||
assert answer.content == text, answer.text
|
||||
assert answer.audio == (None if stream else PCM_B64), answer.text
|
||||
_assert_logged(answer)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("provider", "voice"), [("gemini", "Kore"), ("vertex", "nova")])
|
||||
def test_gemini_voice_names_pass_through_verbatim(gateway: Gateway, provider: str, voice: str) -> None:
|
||||
with wire_server(gemini_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _deployment(scenario, provider, wire.url, gateway.upstream_url)
|
||||
text: Final = marker()
|
||||
answer: Final = _httpx_chat(gateway, model, voice, text, stream=False)
|
||||
assert answer.status == 200, answer.text
|
||||
assert received_voice(_received(wire, text)) == voice
|
||||
assert answer.audio == PCM_B64, answer.text
|
||||
_assert_logged(answer)
|
||||
|
||||
|
||||
def _raw_chat(gateway: Gateway, body: str) -> httpx.Response:
|
||||
headers: Final = {"Authorization": f"Bearer {gateway.key}", "content-type": "application/json"}
|
||||
return gateway.client.post("/v1/chat/completions", content=body.encode(), headers=headers)
|
||||
|
||||
|
||||
def _raw_body(model: str, text: str, audio_literal: str) -> str:
|
||||
return (
|
||||
f'{{"model": {json.dumps(model)}, "messages": [{{"role": "user", "content": {json.dumps(text)}}}], '
|
||||
f'"modalities": ["audio"], "audio": {audio_literal}}}'
|
||||
)
|
||||
|
||||
|
||||
def test_vendor_rejection_of_an_unknown_voice_reaches_the_caller_and_the_proxy_keeps_serving(gateway: Gateway) -> None:
|
||||
with wire_server(gemini_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = gemini_deployment(scenario, wire.url)
|
||||
text: Final = marker()
|
||||
refused: Final = gateway.request("POST", "/v1/chat/completions", chat_body(model, "Custom-Voice", text))
|
||||
assert refused.status_code == 400, refused.text
|
||||
assert f"{REJECTED_PREFIX}Custom-Voice" in error_message(refused), refused.text
|
||||
assert received_voice(_received(wire, text)) == "Custom-Voice"
|
||||
follow_up: Final = marker()
|
||||
served: Final = _httpx_chat(gateway, model, "Kore", follow_up, stream=False)
|
||||
assert served.status == 200, served.text
|
||||
assert received_voice(_received(wire, follow_up)) == "Kore"
|
||||
|
||||
|
||||
MALFORMED: Final = (
|
||||
pytest.param('{"voice": 7, "format": "pcm16"}', "7", id="int"),
|
||||
pytest.param('{"voice": null, "format": "pcm16"}', "None", id="null"),
|
||||
pytest.param('{"voice": "", "format": "pcm16"}', "", id="empty"),
|
||||
pytest.param(f'{{"voice": "{FIVE_KB}", "format": "pcm16"}}', FIVE_KB, id="5kb"),
|
||||
pytest.param('{"voice": ["alloy"], "format": "pcm16"}', "['alloy']", id="list"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("audio_literal", "rejected_as"), MALFORMED)
|
||||
def test_malformed_voice_values_are_refused_without_taking_the_proxy_down(
|
||||
gateway: Gateway, audio_literal: str, rejected_as: str
|
||||
) -> None:
|
||||
with wire_server(gemini_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = gemini_deployment(scenario, wire.url)
|
||||
text: Final = marker()
|
||||
refused: Final = _raw_chat(gateway, _raw_body(model, text, audio_literal))
|
||||
assert refused.status_code == 400, refused.text
|
||||
assert f"{REJECTED_PREFIX}{rejected_as} and language: " in error_message(refused), refused.text
|
||||
received: Final = _received(wire, text)
|
||||
assert received_voice(received) == json.loads(audio_literal)["voice"], received.body
|
||||
follow_up: Final = marker()
|
||||
served: Final = _httpx_chat(gateway, model, "alloy", follow_up, stream=False)
|
||||
assert served.status == 200, served.text
|
||||
assert received_voice(_received(wire, follow_up)) == "Kore"
|
||||
|
||||
|
||||
def test_duplicated_voice_key_keeps_the_last_value(gateway: Gateway) -> None:
|
||||
with wire_server(gemini_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = gemini_deployment(scenario, wire.url)
|
||||
text: Final = marker()
|
||||
response: Final = _raw_chat(
|
||||
gateway, _raw_body(model, text, '{"voice": "alloy", "voice": "fable", "format": "pcm16"}')
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert received_voice(_received(wire, text)) == "Umbriel"
|
||||
assert audio_data(JSON_OBJECT.validate_json(response.content)) == PCM_B64
|
||||
|
||||
|
||||
def test_audio_without_a_voice_sends_no_voice_config(gateway: Gateway) -> None:
|
||||
with wire_server(gemini_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = gemini_deployment(scenario, wire.url)
|
||||
text: Final = marker()
|
||||
body: Final = {**chat_body(model, "alloy", text), "audio": {"format": "pcm16"}}
|
||||
response: Final = gateway.request("POST", "/v1/chat/completions", body)
|
||||
assert response.status_code == 200, response.text
|
||||
assert received_voice(_received(wire, text)) == NO_VOICE
|
||||
assert audio_data(JSON_OBJECT.validate_json(response.content)) == PCM_B64
|
||||
|
||||
|
||||
def _messages_bridge(gateway: Gateway, model: str, text: str, *, stream: bool) -> Answer:
|
||||
client: Final = anthropic_client(gateway)
|
||||
try:
|
||||
if stream:
|
||||
raw: Final = client.messages.with_raw_response.create(
|
||||
model=model, max_tokens=64, messages=[{"role": "user", "content": text}], stream=True
|
||||
)
|
||||
collected: Final = "".join(
|
||||
event.delta.text
|
||||
for event in raw.parse()
|
||||
if event.type == "content_block_delta" and event.delta.type == "text_delta"
|
||||
)
|
||||
return Answer(200, "", "", None, collected, raw.headers.get("x-litellm-call-id", ""))
|
||||
completed: Final = client.messages.with_raw_response.create(
|
||||
model=model, max_tokens=64, messages=[{"role": "user", "content": text}]
|
||||
)
|
||||
message: Final = completed.parse()
|
||||
content: Final = "".join(block.text for block in message.content if block.type == "text")
|
||||
return Answer(
|
||||
200, message.model_dump_json(), message.id, None, content, completed.headers.get("x-litellm-call-id", "")
|
||||
)
|
||||
except anthropic.APIStatusError as error:
|
||||
return _failed(error.status_code, _error_text(error.response))
|
||||
|
||||
|
||||
def _responses_bridge(gateway: Gateway, model: str, text: str, *, stream: bool) -> Answer:
|
||||
client: Final = openai_client(gateway)
|
||||
try:
|
||||
if stream:
|
||||
raw: Final = client.responses.with_raw_response.create(model=model, input=text, stream=True)
|
||||
events: Final = list(raw.parse())
|
||||
collected: Final = "".join(event.delta for event in events if event.type == "response.output_text.delta")
|
||||
assert events[-1].type == "response.completed", [event.type for event in events]
|
||||
return Answer(200, "", "", None, collected, raw.headers.get("x-litellm-call-id", ""))
|
||||
completed: Final = client.responses.with_raw_response.create(model=model, input=text)
|
||||
response: Final = completed.parse()
|
||||
return Answer(
|
||||
200,
|
||||
response.model_dump_json(),
|
||||
response.id,
|
||||
None,
|
||||
response.output_text,
|
||||
completed.headers.get("x-litellm-call-id", ""),
|
||||
)
|
||||
except openai.APIStatusError as error:
|
||||
return _failed(error.status_code, _error_text(error.response))
|
||||
|
||||
|
||||
def _completions_bridge(gateway: Gateway, model: str, text: str, *, stream: bool) -> Answer:
|
||||
response: Final = gateway.request("POST", "/v1/completions", {"model": model, "prompt": text, "max_tokens": 64})
|
||||
if response.status_code != 200:
|
||||
return _failed(response.status_code, response.text)
|
||||
body: Final = JSON_OBJECT.validate_json(response.content)
|
||||
choices: Final = body["choices"]
|
||||
assert isinstance(choices, list) and len(choices) == 1, response.text
|
||||
return Answer(
|
||||
200,
|
||||
response.text,
|
||||
string_value(body["id"]),
|
||||
None,
|
||||
string_value(object_value(choices[0])["text"]),
|
||||
response.headers.get("x-litellm-call-id", ""),
|
||||
)
|
||||
|
||||
|
||||
BRIDGES: Final[Mapping[str, Callable[..., Answer]]] = {
|
||||
"messages": _messages_bridge,
|
||||
"responses": _responses_bridge,
|
||||
"completions": _completions_bridge,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("bridge", "stream"),
|
||||
[("messages", False), ("messages", True), ("responses", False), ("responses", True), ("completions", False)],
|
||||
ids=["messages", "messages-stream", "responses", "responses-stream", "completions"],
|
||||
)
|
||||
def test_bridged_endpoints_carry_the_deployment_voice(gateway: Gateway, bridge: str, stream: bool) -> None:
|
||||
with wire_server(gemini_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = gemini_deployment(scenario, wire.url, audio=audio_param("alloy"), modalities=["audio"])
|
||||
text: Final = marker()
|
||||
answer: Final = BRIDGES[bridge](gateway, model, text, stream=stream)
|
||||
assert answer.status == 200, answer.text
|
||||
request: Final = _received(wire, text)
|
||||
_assert_target(request, "gemini", stream=stream)
|
||||
assert received_voice(request) == "Kore", request.body
|
||||
assert answer.content == text, answer.text
|
||||
row: Final = spend_row_by_call(answer.call_id)
|
||||
assert row["completion_tokens"] == AUDIO_TOKENS, row
|
||||
|
||||
|
||||
def _gcs_peer(request: Request) -> Reply:
|
||||
assert request.method == "POST", request.method
|
||||
assert request.target.startswith(f"/upload/storage/v1/b/{BUCKET}/o?uploadType=media&name="), request.target
|
||||
assert request.headers["authorization"] == "Bearer scripted-token", request.headers
|
||||
name: Final = request.target.rsplit("name=", 1)[1]
|
||||
body: Final = {
|
||||
"kind": "storage#object",
|
||||
"id": f"{BUCKET}/{name}/1",
|
||||
"name": name,
|
||||
"bucket": BUCKET,
|
||||
"size": str(len(request.body)),
|
||||
"timeCreated": "2026-10-08T00:00:00.000Z",
|
||||
"contentType": "application/json",
|
||||
}
|
||||
return Reply(body=json.dumps(body).encode())
|
||||
|
||||
|
||||
def _batch_line(custom_id: str, text: str, voice: str) -> str:
|
||||
body: Final = {
|
||||
"model": BACKEND,
|
||||
"messages": [{"role": "user", "content": text}],
|
||||
"modalities": ["audio"],
|
||||
"audio": audio_param(voice),
|
||||
}
|
||||
return json.dumps({"custom_id": custom_id, "method": "POST", "url": "/v1/chat/completions", "body": body})
|
||||
|
||||
|
||||
def _uploaded_rows(request: Request) -> tuple[dict[str, JsonValue], ...]:
|
||||
lines: Final = tuple(line for line in request.body.decode().splitlines() if line.strip())
|
||||
return tuple(object_value(JSON_OBJECT.validate_json(line)["request"]) for line in lines)
|
||||
|
||||
|
||||
def _decoded_file_id(file_id: str) -> tuple[str, str]:
|
||||
assert file_id.startswith("file-"), file_id
|
||||
encoded: Final = file_id[len("file-") :]
|
||||
decoded: Final = base64.urlsafe_b64decode(encoded + "=" * (-len(encoded) % 4)).decode()
|
||||
assert decoded.startswith("litellm:"), decoded
|
||||
raw, model = decoded[len("litellm:") :].rsplit(";model,", 1)
|
||||
return raw, model
|
||||
|
||||
|
||||
def test_vertex_batch_file_rows_carry_the_mapped_voice(gateway: Gateway) -> None:
|
||||
with chunked_peer(_gcs_peer) as wire, gateway.scenario() as scenario:
|
||||
model: Final = vertex_deployment(
|
||||
scenario, wire.url, gateway.upstream_url, api_base=wire.url, gcs_bucket_name=BUCKET
|
||||
)
|
||||
first: Final = marker()
|
||||
second: Final = marker()
|
||||
content: Final = f"{_batch_line('row-1', first, 'alloy')}\n{_batch_line('row-2', second, 'Kore')}\n".encode()
|
||||
response: Final = gateway.request_multipart(
|
||||
"/v1/files", {"purpose": "batch", "model": model}, {"file": ("rows.jsonl", content, "application/jsonl")}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
uploaded: Final = JSON_OBJECT.validate_json(response.content)
|
||||
requests: Final = wire.drain()
|
||||
assert len(requests) == 1, [request.target for request in requests]
|
||||
assert (uploaded["object"], uploaded["purpose"], uploaded["bytes"]) == ("file", "batch", len(requests[0].body))
|
||||
raw, file_model = _decoded_file_id(string_value(uploaded["id"]))
|
||||
assert raw.startswith(f"gs://{BUCKET}/") and file_model == model, response.text
|
||||
rows: Final = _uploaded_rows(requests[0])
|
||||
assert tuple(voice_in(row) for row in rows) == ("Kore", "Kore"), requests[0].body
|
||||
assert tuple(text_in(row) for row in rows) == (first, second)
|
||||
|
||||
|
||||
def test_vertex_live_default_mode_sends_one_setup_without_the_client_voice(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = live_scenario(scenario)
|
||||
key: Final = scenario.key()
|
||||
model: Final = live_deployment(scenario, handle.scenario_id, gateway.upstream_url)
|
||||
frames: Final = live_turn(gateway, model, key, "fable")
|
||||
types: Final = frame_types(frames)
|
||||
assert types[0] == "session.created" and types[-1] == "response.done", frames
|
||||
setups: Final = live_setups(gateway, handle.scenario_id)
|
||||
assert len(setups) == 1, setups
|
||||
assert voice_in(setups[0]) == NO_VOICE, setups[0]
|
||||
|
||||
|
||||
def _free_port() -> int:
|
||||
with socket.socket() as reserve:
|
||||
reserve.bind(("127.0.0.1", 0))
|
||||
return int(reserve.getsockname()[1])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Fired:
|
||||
provider: str
|
||||
stream: bool
|
||||
voice: str
|
||||
text: str
|
||||
status: int
|
||||
text_body: str
|
||||
call_id: str
|
||||
|
||||
|
||||
VOICES: Final = ("alloy", "fable", "Kore")
|
||||
|
||||
|
||||
async def _fire(
|
||||
gateway: Gateway, models: Mapping[str, str], count: int, *, tolerate_transport_errors: bool = False
|
||||
) -> tuple[Fired, ...]:
|
||||
async def one(client: httpx.AsyncClient, index: int) -> Fired:
|
||||
provider: Final = "gemini" if index % 2 == 0 else "vertex"
|
||||
stream: Final = index % 4 >= 2
|
||||
voice: Final = VOICES[index % len(VOICES)]
|
||||
text: Final = marker()
|
||||
body: Final = chat_body(models[provider], voice, text, stream=stream)
|
||||
try:
|
||||
response: Final = await client.post(
|
||||
"/v1/chat/completions", json=body, headers={"Authorization": f"Bearer {gateway.key}"}
|
||||
)
|
||||
except httpx.TransportError as error:
|
||||
if not tolerate_transport_errors:
|
||||
raise
|
||||
return Fired(provider, stream, voice, text, 0, repr(error), "")
|
||||
call_id: Final = response.headers.get("x-litellm-call-id", "")
|
||||
if response.status_code != 200:
|
||||
return Fired(provider, stream, voice, text, response.status_code, response.text, call_id)
|
||||
content: Final = (
|
||||
"".join(_delta_text(chunk) for chunk in _sse_chunks(response.text))
|
||||
if stream
|
||||
else _message_text(JSON_OBJECT.validate_json(response.content))
|
||||
)
|
||||
assert content == text, response.text
|
||||
return Fired(provider, stream, voice, text, 200, response.text, call_id)
|
||||
|
||||
async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client:
|
||||
return tuple(await asyncio.gather(*(one(client, index) for index in range(count))))
|
||||
|
||||
|
||||
def _by_text(received: tuple[Request, ...]) -> dict[str, tuple[Request, ...]]:
|
||||
texts: Final = {received_text(request) for request in received}
|
||||
return {text: tuple(request for request in received if received_text(request) == text) for text in texts}
|
||||
|
||||
|
||||
def _expected_voice(voice: str) -> str:
|
||||
return {"alloy": "Kore", "fable": "Umbriel"}.get(voice, voice)
|
||||
|
||||
|
||||
def _assert_served(items: tuple[Fired, ...], received: tuple[Request, ...]) -> None:
|
||||
by_text: Final = _by_text(received)
|
||||
for item in items:
|
||||
assert item.status == 200, (item.provider, item.voice, item.text_body)
|
||||
assert len(by_text.get(item.text, ())) == 1, item.text
|
||||
assert received_voice(by_text[item.text][0]) == _expected_voice(item.voice), item.text
|
||||
row: Final = spend_row_by_call(item.call_id)
|
||||
assert row["call_type"] == "acompletion", row
|
||||
|
||||
|
||||
def test_peer_outage_mid_burst_ends_each_request_once_and_recovers_on_the_same_port(gateway: Gateway) -> None:
|
||||
port: Final = _free_port()
|
||||
with gateway.scenario() as scenario:
|
||||
url: Final = f"http://127.0.0.1:{port}"
|
||||
models: Final = {
|
||||
"gemini": gemini_deployment(scenario, url),
|
||||
"vertex": vertex_deployment(scenario, url, gateway.upstream_url),
|
||||
}
|
||||
with wire_server(gemini_peer, port=port) as wire:
|
||||
healthy: Final = asyncio.run(_fire(gateway, models, BURST))
|
||||
_assert_served(healthy, wire.drain())
|
||||
racing: Final = _start_in_thread(lambda: asyncio.run(_fire(gateway, models, BURST)))
|
||||
eventually(lambda: wire.received.qsize(), lambda size: size >= OUTAGE_AFTER, 30)
|
||||
down: Final = asyncio.run(_fire(gateway, models, BURST))
|
||||
raced: Final = racing()
|
||||
raced_received: Final = wire.drain()
|
||||
with wire_server(gemini_peer, port=port) as restarted:
|
||||
recovered: Final = asyncio.run(_fire(gateway, models, BURST))
|
||||
recovered_received: Final = restarted.drain()
|
||||
assert all(len(requests) == 1 for requests in _by_text(raced_received).values())
|
||||
_assert_served(tuple(item for item in raced if item.status == 200), raced_received)
|
||||
lost: Final = tuple(item for item in raced if item.status != 200)
|
||||
for item in (*lost, *down):
|
||||
assert item.status >= 500, (item.status, item.text_body)
|
||||
assert '"error"' in item.text_body, item.text_body
|
||||
assert all(item.status != 200 for item in down), [item.status for item in down]
|
||||
_assert_served(recovered, recovered_received)
|
||||
assert {received_text(request) for request in recovered_received} == {item.text for item in recovered}
|
||||
|
||||
|
||||
def _start_in_thread(work: Callable[[], tuple[Fired, ...]]) -> Callable[[], tuple[Fired, ...]]:
|
||||
results: Final[list[tuple[Fired, ...]]] = []
|
||||
thread: Final = threading.Thread(target=lambda: results.append(work()))
|
||||
thread.start()
|
||||
|
||||
def finish() -> tuple[Fired, ...]:
|
||||
thread.join(timeout=120)
|
||||
assert not thread.is_alive(), "the racing burst never finished"
|
||||
return results[0]
|
||||
|
||||
return finish
|
||||
|
||||
|
||||
def _slow_peer(request: Request) -> Reply:
|
||||
reply: Final = gemini_peer(request)
|
||||
if reply.chunks is not None:
|
||||
return dataclasses.replace(reply, pause_between_chunks=SLOW_SECONDS)
|
||||
time.sleep(SLOW_SECONDS)
|
||||
return reply
|
||||
|
||||
|
||||
def test_slow_peer_under_a_burst_answers_every_request_once_while_liveliness_stays_up(gateway: Gateway) -> None:
|
||||
with wire_server(_slow_peer) as wire, gateway.scenario() as scenario:
|
||||
models: Final = {
|
||||
"gemini": gemini_deployment(scenario, wire.url),
|
||||
"vertex": vertex_deployment(scenario, wire.url, gateway.upstream_url),
|
||||
}
|
||||
burst: Final = _start_in_thread(lambda: asyncio.run(_fire(gateway, models, SLOW_BURST)))
|
||||
eventually(lambda: wire.received.qsize(), lambda size: size >= SLOW_BURST, 30)
|
||||
liveliness: Final = gateway.request("GET", "/health/liveliness")
|
||||
assert liveliness.status_code == 200, liveliness.text
|
||||
served: Final = burst()
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == SLOW_BURST, [received_text(request) for request in received]
|
||||
_assert_served(served, received)
|
||||
|
|
@ -8899,3 +8899,49 @@ def load_vertex_ai_credentials():
|
|||
|
||||
# Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS
|
||||
os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name)
|
||||
|
||||
|
||||
# Gemini prebuilt TTS voices per https://ai.google.dev/gemini-api/docs/speech-generation (page dated 2026-10-01, read 2026-10-07)
|
||||
GEMINI_TTS_PREBUILT_VOICES: Final[frozenset[str]] = frozenset(
|
||||
{
|
||||
"Zephyr", "Puck", "Charon", "Kore", "Fenrir", "Leda", "Orus", "Aoede", "Callirrhoe", "Autonoe",
|
||||
"Enceladus", "Iapetus", "Umbriel", "Algieba", "Despina", "Erinome", "Algenib", "Rasalgethi",
|
||||
"Laomedeia", "Achernar", "Alnilam", "Schedar", "Gacrux", "Pulcherrima", "Achird", "Zubenelgenubi",
|
||||
"Vindemiatrix", "Sadachbia", "Sadaltager", "Sulafat",
|
||||
}
|
||||
)
|
||||
OPENAI_TTS_VOICES: Final[tuple[str, ...]] = (
|
||||
"alloy", "ash", "ballad", "cedar", "coral", "echo", "fable", "marin", "onyx", "sage", "shimmer", "verse",
|
||||
)
|
||||
|
||||
|
||||
def _mapped_voice_name(audio: dict[str, str]) -> str:
|
||||
return VertexGeminiConfig()._map_audio_params(audio)["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("openai_voice", OPENAI_TTS_VOICES)
|
||||
def test_map_audio_params_maps_openai_voice_names_to_gemini_prebuilt_voices(openai_voice: str) -> None:
|
||||
"""Regression for the audio_speech health check default voice (alloy) 400ing on Gemini TTS models."""
|
||||
mapped: Final = _mapped_voice_name({"voice": openai_voice, "format": "pcm16"})
|
||||
assert mapped in GEMINI_TTS_PREBUILT_VOICES
|
||||
assert _mapped_voice_name({"voice": openai_voice.capitalize()}) == mapped
|
||||
|
||||
|
||||
def test_map_audio_params_maps_the_health_check_default_voice_to_kore() -> None:
|
||||
assert _mapped_voice_name({"voice": "alloy", "format": "pcm16"}) == "Kore"
|
||||
|
||||
|
||||
# nova is accepted by Gemini 3.x TTS models as given (live check 2026-10-08), so it must not be remapped
|
||||
@pytest.mark.parametrize("gemini_voice", ("Kore", "Puck", "Sulafat", "nova", "custom-voice"))
|
||||
def test_map_audio_params_passes_non_openai_voice_names_through(gemini_voice: str) -> None:
|
||||
assert _mapped_voice_name({"voice": gemini_voice, "format": "pcm16"}) == gemini_voice
|
||||
|
||||
|
||||
def test_map_openai_params_audio_voice_reaches_speech_config_mapped() -> None:
|
||||
optional_params: Final = VertexGeminiConfig().map_openai_params(
|
||||
non_default_params={"modalities": ["audio"], "audio": {"voice": "alloy", "format": "pcm16"}},
|
||||
optional_params={},
|
||||
model="gemini-3.8-flash-tts",
|
||||
drop_params=False,
|
||||
)
|
||||
assert optional_params["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue