feat(tts): add Gemini audio tag rewrite
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
|
||||
import base64
|
||||
import struct
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -312,6 +313,112 @@ class TestGenerateGeminiTts:
|
||||
assert prompt_text == "Hi"
|
||||
assert "persona prompt file unavailable" in caplog.text
|
||||
|
||||
def test_audio_tags_disabled_does_not_call_rewriter(
|
||||
self, tmp_path, monkeypatch, mock_gemini_response
|
||||
):
|
||||
from tools.tts_tool import _generate_gemini_tts
|
||||
|
||||
config = {
|
||||
"gemini": {
|
||||
"model": "gemini-3.1-flash-tts-preview",
|
||||
"audio_tags": False,
|
||||
}
|
||||
}
|
||||
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
|
||||
|
||||
with patch("agent.auxiliary_client.call_llm") as mock_call_llm, \
|
||||
patch("requests.post", return_value=mock_gemini_response) as mock_post:
|
||||
_generate_gemini_tts("Hi there.", str(tmp_path / "test.wav"), config)
|
||||
|
||||
mock_call_llm.assert_not_called()
|
||||
prompt_text = mock_post.call_args[1]["json"]["contents"][0]["parts"][0]["text"]
|
||||
assert prompt_text == "Hi there."
|
||||
|
||||
def test_audio_tags_enabled_rewrites_hidden_tts_script(
|
||||
self, tmp_path, monkeypatch, mock_gemini_response
|
||||
):
|
||||
from tools.tts_tool import _generate_gemini_tts
|
||||
|
||||
persona_file = tmp_path / "voice-persona.md"
|
||||
persona_file.write_text(
|
||||
"### DIRECTOR'S NOTES\nStyle: Warm and amused.",
|
||||
encoding="utf-8",
|
||||
)
|
||||
response = SimpleNamespace(
|
||||
choices=[
|
||||
SimpleNamespace(
|
||||
message=SimpleNamespace(content="[warmly] Hi there. [soft laugh]")
|
||||
)
|
||||
]
|
||||
)
|
||||
config = {
|
||||
"gemini": {
|
||||
"model": "gemini-3.1-flash-tts-preview",
|
||||
"audio_tags": True,
|
||||
"persona_prompt_file": str(persona_file),
|
||||
}
|
||||
}
|
||||
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
|
||||
|
||||
with patch("agent.auxiliary_client.call_llm", return_value=response) as mock_call_llm, \
|
||||
patch("requests.post", return_value=mock_gemini_response) as mock_post:
|
||||
_generate_gemini_tts("Hi there.", str(tmp_path / "test.wav"), config)
|
||||
|
||||
mock_call_llm.assert_called_once()
|
||||
call_kwargs = mock_call_llm.call_args.kwargs
|
||||
assert call_kwargs["task"] == "tts_audio_tags"
|
||||
assert "Audio tags are inline square-bracket modifiers" in call_kwargs["messages"][0]["content"]
|
||||
assert "Style: Warm and amused." in call_kwargs["messages"][1]["content"]
|
||||
assert "Hi there." in call_kwargs["messages"][1]["content"]
|
||||
|
||||
prompt_text = mock_post.call_args[1]["json"]["contents"][0]["parts"][0]["text"]
|
||||
assert "Synthesize speech from the TRANSCRIPT only" in prompt_text
|
||||
assert "### DIRECTOR'S NOTES\nStyle: Warm and amused." in prompt_text
|
||||
assert "#### TRANSCRIPT\n[warmly] Hi there. [soft laugh]" in prompt_text
|
||||
|
||||
def test_audio_tags_enabled_skips_non_tag_capable_model(
|
||||
self, tmp_path, monkeypatch, mock_gemini_response, caplog
|
||||
):
|
||||
from tools.tts_tool import _generate_gemini_tts
|
||||
|
||||
config = {
|
||||
"gemini": {
|
||||
"model": "gemini-2.5-flash-preview-tts",
|
||||
"audio_tags": True,
|
||||
}
|
||||
}
|
||||
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
|
||||
|
||||
with patch("agent.auxiliary_client.call_llm") as mock_call_llm, \
|
||||
patch("requests.post", return_value=mock_gemini_response) as mock_post:
|
||||
_generate_gemini_tts("Hi there.", str(tmp_path / "test.wav"), config)
|
||||
|
||||
mock_call_llm.assert_not_called()
|
||||
prompt_text = mock_post.call_args[1]["json"]["contents"][0]["parts"][0]["text"]
|
||||
assert prompt_text == "Hi there."
|
||||
assert "not known to support Gemini audio tags" in caplog.text
|
||||
|
||||
def test_audio_tag_rewrite_failure_falls_back_to_original_text(
|
||||
self, tmp_path, monkeypatch, mock_gemini_response, caplog
|
||||
):
|
||||
from tools.tts_tool import _generate_gemini_tts
|
||||
|
||||
config = {
|
||||
"gemini": {
|
||||
"model": "gemini-3.1-flash-tts-preview",
|
||||
"audio_tags": True,
|
||||
}
|
||||
}
|
||||
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
|
||||
|
||||
with patch("agent.auxiliary_client.call_llm", side_effect=RuntimeError("boom")), \
|
||||
patch("requests.post", return_value=mock_gemini_response) as mock_post:
|
||||
_generate_gemini_tts("Hi there.", str(tmp_path / "test.wav"), config)
|
||||
|
||||
prompt_text = mock_post.call_args[1]["json"]["contents"][0]["parts"][0]["text"]
|
||||
assert prompt_text == "Hi there."
|
||||
assert "audio tag rewrite failed" in caplog.text
|
||||
|
||||
|
||||
class TestGeminiInCheckRequirements:
|
||||
def test_gemini_api_key_satisfies_requirements(self, monkeypatch):
|
||||
|
||||
Reference in New Issue
Block a user