fix(vertex-ai): recognize provider-stripped TTS model names

This commit is contained in:
Emerson Gomes 2026-09-15 13:13:20 -05:00
parent c3baeed801
commit dde62d74ba
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
2 changed files with 8 additions and 3 deletions

View file

@ -736,7 +736,11 @@ def is_gemini_tts_model(model: str, custom_llm_provider: str | None = None) -> b
else model
)
bundled_model_cost: Final = _get_bundled_model_cost_map()
bundled_model_info: Final = bundled_model_cost.get(provider_model) or bundled_model_cost.get(model)
bundled_model_info: Final = (
bundled_model_cost.get(provider_model)
or bundled_model_cost.get(model)
or (bundled_model_cost.get(f"vertex_ai/{model}") if custom_llm_provider is None else None)
)
return (
bundled_model_info is not None
and bundled_model_info.get("mode") == "audio_speech"

View file

@ -531,12 +531,13 @@ if __name__ == "__main__":
@pytest.mark.parametrize("voice_key", ["name", "voiceName", "voice_name", "voice"])
@pytest.mark.parametrize("wrapper", [None, "speechConfig", "speech_config"])
def test_speech_bridge_preserves_single_speaker_voice_mapping(voice_key: str, wrapper: str | None):
@pytest.mark.parametrize("model", ["gemini-2.5-flash-tts", "gemini-2.5-pro-tts", GEMINI_3_1_FLASH_TTS_MODEL])
def test_speech_bridge_preserves_single_speaker_voice_mapping(voice_key: str, wrapper: str | None, model: str):
handler = SpeechToCompletionBridgeTransformationHandler()
config = {voice_key: "Umbriel", "language_code": "en-US"}
voice = {wrapper: config} if wrapper else config
result = handler.transform_request(
model=f"vertex_ai/{GEMINI_3_1_FLASH_TTS_MODEL}",
model=model,
input="Hello",
voice=voice,
optional_params={},