feat(voice): add admin-configurable elevenlabs tts
This commit is contained in:
@@ -12,6 +12,7 @@ ai_module = importlib.import_module("services.ai_orchestrator_service.app")
|
||||
ai_app = ai_module.app
|
||||
voice_module = importlib.import_module("services.ai_orchestrator_service.voice")
|
||||
voice_config_module = importlib.import_module("services.ai_orchestrator_service.voice_name_config")
|
||||
voice_tts_config_module = importlib.import_module("services.shared.voice_tts_config")
|
||||
from services.interaction_service.app import app as interaction_app
|
||||
from services.shared.core import new_id, utc_now_iso
|
||||
from services.shared.db import get_session
|
||||
@@ -30,6 +31,7 @@ from services.shared.sql_models import (
|
||||
TelegramMessageRow,
|
||||
TelegramThreadRow,
|
||||
VoiceNameCollectionSettingsRow,
|
||||
VoiceTTSSettingsRow,
|
||||
VoiceAISessionRow,
|
||||
WhatsAppThreadRow,
|
||||
)
|
||||
@@ -66,6 +68,27 @@ def reset_voice_name_collection_settings():
|
||||
session.close()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_voice_tts_settings():
|
||||
session = get_session()
|
||||
try:
|
||||
row = session.execute(select(VoiceTTSSettingsRow)).scalar_one_or_none()
|
||||
if row is not None:
|
||||
session.delete(row)
|
||||
session.commit()
|
||||
finally:
|
||||
session.close()
|
||||
yield
|
||||
session = get_session()
|
||||
try:
|
||||
row = session.execute(select(VoiceTTSSettingsRow)).scalar_one_or_none()
|
||||
if row is not None:
|
||||
session.delete(row)
|
||||
session.commit()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _u(value: str) -> str:
|
||||
return value.encode("ascii").decode("unicode_escape")
|
||||
|
||||
@@ -2453,6 +2476,50 @@ def test_voice_name_collection_config_put_rejects_invalid_name_template():
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
def test_voice_tts_config_get_returns_defaults_when_not_persisted():
|
||||
ai_client = TestClient(ai_app)
|
||||
|
||||
response = ai_client.get("/ai/voice/config/tts", headers=admin_headers())
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = response.json()
|
||||
assert payload["source"] == "defaults"
|
||||
assert payload["updated_at"] is None
|
||||
assert payload["config"]["provider"] == "yandex"
|
||||
assert "elevenlabs" in payload["provider_options"]
|
||||
assert payload["voice_options"]["elevenlabs"]["ru"][0]["label"] == "Brian"
|
||||
|
||||
|
||||
def test_voice_tts_config_put_persists_custom_payload():
|
||||
ai_client = TestClient(ai_app)
|
||||
payload = voice_tts_config_module.voice_tts_default_config().model_dump()
|
||||
payload["provider"] = "elevenlabs"
|
||||
payload["elevenlabs"]["ru"]["voice"] = "nPczCjzI2devNBz1zQrb"
|
||||
payload["elevenlabs"]["ru"]["model_id"] = "eleven_v3"
|
||||
payload["elevenlabs"]["kz"]["language_code"] = "kk"
|
||||
|
||||
response = ai_client.put(
|
||||
"/ai/voice/config/tts",
|
||||
headers=admin_headers(),
|
||||
json=payload,
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
saved = response.json()
|
||||
assert saved["source"] == "database"
|
||||
assert saved["config"]["provider"] == "elevenlabs"
|
||||
assert saved["config"]["elevenlabs"]["ru"]["voice"] == "nPczCjzI2devNBz1zQrb"
|
||||
assert saved["config"]["elevenlabs"]["ru"]["model_id"] == "eleven_v3"
|
||||
|
||||
session = get_session()
|
||||
try:
|
||||
row = session.execute(select(VoiceTTSSettingsRow)).scalar_one()
|
||||
assert "elevenlabs" in row.config_json
|
||||
assert "nPczCjzI2devNBz1zQrb" in row.config_json
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_voice_start_disabled_hands_off_without_prompt():
|
||||
session = get_session()
|
||||
try:
|
||||
|
||||
@@ -9,13 +9,25 @@ import services.ai_voice_runtime_service.app as runtime_module
|
||||
from services.shared.core import utc_now_iso
|
||||
from services.shared.db import get_session
|
||||
from services.shared.models import VoiceAIStartIn
|
||||
from services.shared.sql_models import VoiceAISessionRow, VoiceTranscriptSegmentRow
|
||||
from services.shared.sql_models import VoiceAISessionRow, VoiceTranscriptSegmentRow, VoiceTTSSettingsRow
|
||||
from services.shared.voice_tts_config import save_voice_tts_config, voice_tts_default_config
|
||||
|
||||
|
||||
def _admin_headers() -> dict[str, str]:
|
||||
return {"X-User": "admin", "X-Role": "admin"}
|
||||
|
||||
|
||||
def _reset_voice_tts_settings() -> None:
|
||||
session = get_session()
|
||||
try:
|
||||
row = session.execute(select(VoiceTTSSettingsRow)).scalar_one_or_none()
|
||||
if row is not None:
|
||||
session.delete(row)
|
||||
session.commit()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_create_voice_ai_session_records_greeting_without_blocking_on_orchestrator(monkeypatch):
|
||||
allow_start = threading.Event()
|
||||
|
||||
@@ -194,6 +206,60 @@ def test_create_voice_ai_session_records_greeting_without_blocking_on_orchestrat
|
||||
raise AssertionError("background voice-session start did not finish in time")
|
||||
|
||||
|
||||
def test_create_voice_ai_session_uses_database_selected_tts_provider(monkeypatch):
|
||||
_reset_voice_tts_settings()
|
||||
|
||||
session = get_session()
|
||||
try:
|
||||
payload = voice_tts_default_config().model_dump()
|
||||
payload["provider"] = "elevenlabs"
|
||||
save_voice_tts_config(session, payload)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
def _fake_orchestrator_request(method: str, path: str, *, payload=None, timeout=10.0):
|
||||
assert method == "POST"
|
||||
assert path.endswith("/start")
|
||||
return {
|
||||
"session_id": "ais_voice_tts_runtime",
|
||||
"language": "ru",
|
||||
"greeting_text": "Здравствуйте. Чем помочь?",
|
||||
"disclosure_required": False,
|
||||
}
|
||||
|
||||
monkeypatch.setattr(runtime_module, "_orchestrator_request", _fake_orchestrator_request)
|
||||
|
||||
client = TestClient(runtime_module.app)
|
||||
response = client.post(
|
||||
"/internal/voice-ai/sessions",
|
||||
headers=_admin_headers(),
|
||||
json={
|
||||
"call_id": "call_voice_runtime_tts_provider",
|
||||
"linked_id": "linked_voice_runtime_tts_provider",
|
||||
"interaction_id": "int_voice_runtime_tts_provider",
|
||||
"queue_id": "que_voice_runtime_tts_provider",
|
||||
"caller_number": "+77010000032",
|
||||
"caller_name": "Runtime Caller",
|
||||
"agent_profile": "voice_support",
|
||||
"language_hint": "ru",
|
||||
"handoff_queue_id": "que_voice_runtime_tts_provider",
|
||||
"metadata": {"queue_code": "voice_lab_ai"},
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
session = get_session()
|
||||
try:
|
||||
voice_session = session.execute(
|
||||
select(VoiceAISessionRow).where(VoiceAISessionRow.call_id == "call_voice_runtime_tts_provider")
|
||||
).scalar_one()
|
||||
assert voice_session.tts_provider == "elevenlabs"
|
||||
finally:
|
||||
session.close()
|
||||
_reset_voice_tts_settings()
|
||||
|
||||
|
||||
def test_mark_reply_delivered_finalizes_only_matching_pending_assistant_segment():
|
||||
now = utc_now_iso()
|
||||
session = get_session()
|
||||
|
||||
@@ -3,7 +3,38 @@ from __future__ import annotations
|
||||
import base64
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
|
||||
from services.ai_voice_runtime_service.runtime_tts_provider import RuntimeConfiguredTTSProvider
|
||||
from services.ai_voice_runtime_service.providers import tts as tts_module
|
||||
from services.shared.db import get_session
|
||||
from services.shared.sql_init import init_sql_schema
|
||||
from services.shared.sql_models import VoiceTTSSettingsRow
|
||||
from services.shared.voice_tts_config import save_voice_tts_config, voice_tts_default_config
|
||||
|
||||
init_sql_schema()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_voice_tts_settings():
|
||||
session = get_session()
|
||||
try:
|
||||
row = session.execute(select(VoiceTTSSettingsRow)).scalar_one_or_none()
|
||||
if row is not None:
|
||||
session.delete(row)
|
||||
session.commit()
|
||||
finally:
|
||||
session.close()
|
||||
yield
|
||||
session = get_session()
|
||||
try:
|
||||
row = session.execute(select(VoiceTTSSettingsRow)).scalar_one_or_none()
|
||||
if row is not None:
|
||||
session.delete(row)
|
||||
session.commit()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
class _DummyResponse:
|
||||
@@ -252,3 +283,88 @@ def test_yandex_tts_provider_uses_cached_audio_without_credentials(tmp_path, mon
|
||||
|
||||
assert replay.audio_bytes == cached.audio_bytes
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
def test_elevenlabs_tts_provider_posts_voice_id_model_and_language_code(tmp_path, monkeypatch):
|
||||
calls: list[dict] = []
|
||||
|
||||
class _DummyClient:
|
||||
def __init__(self, *, timeout: float) -> None:
|
||||
self.timeout = timeout
|
||||
|
||||
def __enter__(self) -> _DummyClient:
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb) -> None:
|
||||
return None
|
||||
|
||||
def post(self, url: str, *, headers: dict[str, str], params: dict[str, str], json: dict) -> _DummyResponse:
|
||||
calls.append({"url": url, "headers": headers, "params": params, "json": json, "timeout": self.timeout})
|
||||
return _DummyResponse(b"\x01\x00\x02\x00")
|
||||
|
||||
monkeypatch.setenv("AI_VOICE_TTS_ELEVENLABS_API_KEY", "elevenlabs-key")
|
||||
monkeypatch.setenv("AI_VOICE_TTS_CACHE_ENABLED", "0")
|
||||
monkeypatch.setattr(tts_module.httpx, "Client", _DummyClient)
|
||||
|
||||
provider = tts_module.ElevenLabsTTSProvider(
|
||||
api_base="https://api.elevenlabs.example",
|
||||
ru_voice="nPczCjzI2devNBz1zQrb",
|
||||
kz_voice="nPczCjzI2devNBz1zQrb",
|
||||
ru_model="eleven_v3",
|
||||
kz_model="eleven_v3",
|
||||
ru_language_code="ru",
|
||||
kz_language_code="kk",
|
||||
output_format="pcm_16000",
|
||||
)
|
||||
synthesis = provider.synthesize("Привет из ElevenLabs", language="kz")
|
||||
|
||||
assert synthesis.audio_bytes == b"\x01\x00\x02\x00"
|
||||
assert synthesis.sample_rate_hz == 16000
|
||||
assert len(calls) == 1
|
||||
assert calls[0]["url"] == "https://api.elevenlabs.example/v1/text-to-speech/nPczCjzI2devNBz1zQrb"
|
||||
assert calls[0]["headers"]["xi-api-key"] == "elevenlabs-key"
|
||||
assert calls[0]["params"]["output_format"] == "pcm_16000"
|
||||
assert calls[0]["json"]["model_id"] == "eleven_v3"
|
||||
assert calls[0]["json"]["language_code"] == "kk"
|
||||
|
||||
|
||||
def test_runtime_configured_tts_provider_uses_database_selected_provider(monkeypatch):
|
||||
calls: list[dict] = []
|
||||
|
||||
class _DummyClient:
|
||||
def __init__(self, *, timeout: float) -> None:
|
||||
self.timeout = timeout
|
||||
|
||||
def __enter__(self) -> _DummyClient:
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb) -> None:
|
||||
return None
|
||||
|
||||
def post(self, url: str, *, headers: dict[str, str], params: dict[str, str], json: dict) -> _DummyResponse:
|
||||
calls.append({"url": url, "headers": headers, "params": params, "json": json, "timeout": self.timeout})
|
||||
return _DummyResponse(b"\x10\x00\x20\x00")
|
||||
|
||||
monkeypatch.setenv("AI_VOICE_TTS_ELEVENLABS_API_KEY", "elevenlabs-key")
|
||||
monkeypatch.setenv("AI_VOICE_TTS_CACHE_ENABLED", "0")
|
||||
monkeypatch.setattr(tts_module.httpx, "Client", _DummyClient)
|
||||
|
||||
session = get_session()
|
||||
try:
|
||||
payload = voice_tts_default_config().model_dump()
|
||||
payload["provider"] = "elevenlabs"
|
||||
payload["elevenlabs"]["ru"]["voice"] = "nPczCjzI2devNBz1zQrb"
|
||||
payload["elevenlabs"]["kz"]["voice"] = "nPczCjzI2devNBz1zQrb"
|
||||
payload["elevenlabs"]["ru"]["model_id"] = "eleven_v3"
|
||||
payload["elevenlabs"]["kz"]["model_id"] = "eleven_v3"
|
||||
save_voice_tts_config(session, payload)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
provider = RuntimeConfiguredTTSProvider(default_provider_name="yandex")
|
||||
synthesis = provider.synthesize("Сәлем", language="kz")
|
||||
|
||||
assert provider.current_provider_name() == "elevenlabs"
|
||||
assert synthesis.audio_bytes == b"\x10\x00\x20\x00"
|
||||
assert len(calls) == 1
|
||||
assert calls[0]["json"]["language_code"] == "kk"
|
||||
|
||||
Reference in New Issue
Block a user