diff --git a/callback_client.py b/callback_client.py new file mode 100644 index 0000000..215ba33 --- /dev/null +++ b/callback_client.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +import logging +import os +from typing import Any + +import httpx + +from realtime_voice_service.crm_client import _issue_service_token + + +LOGGER = logging.getLogger("uvicorn.error") + + +def _enabled() -> bool: + return os.getenv("CALLBACK_DISPATCH_ENABLED", "1").strip().lower() in {"1", "true", "yes"} + + +def _bridge_base_url() -> str: + return str( + os.getenv("ASTERISK_BRIDGE_SERVICE_URL", "http://asterisk-bridge-service:8000") + ).rstrip("/") + + +def _secret() -> str: + return str(os.getenv("CRM_APP_TOKEN_SECRET", "")).strip() + + +def _request_headers() -> dict[str, str] | None: + secret = _secret() + if not secret: + return None + return {"Authorization": f"Bearer {_issue_service_token(secret)}"} + + +async def schedule_callback( + *, + phone: str, + in_minutes: int | None = None, + scheduled_for: str | None = None, + reason: str | None = None, + customer_name: str | None = None, + source_call_id: str | None = None, + source_interaction_id: str | None = None, + idempotency_key: str | None = None, +) -> dict[str, Any] | None: + if not _enabled(): + LOGGER.info("callback dispatch disabled by env") + return None + + headers = _request_headers() + if headers is None: + LOGGER.warning("callback client: CRM_APP_TOKEN_SECRET is not configured") + return None + + payload: dict[str, Any] = {"phone": phone} + if in_minutes is not None: + payload["in_minutes"] = int(in_minutes) + if scheduled_for: + payload["scheduled_for"] = scheduled_for + if reason: + payload["reason"] = reason + if customer_name: + payload["customer_name"] = customer_name + if source_call_id: + payload["source_call_id"] = source_call_id + if source_interaction_id: + payload["source_interaction_id"] = source_interaction_id + if idempotency_key: + payload["idempotency_key"] = idempotency_key + + url = f"{_bridge_base_url()}/asterisk/callbacks" + try: + async with httpx.AsyncClient(timeout=8.0) as client: + resp = await client.post(url, json=payload, headers=headers) + if resp.status_code in (200, 201): + data = resp.json() + LOGGER.info( + "callback scheduled: job_id=%s phone=%s scheduled_for=%s", + data.get("job_id"), + data.get("phone_e164"), + data.get("scheduled_for"), + ) + return data + if resp.status_code == 409: + LOGGER.info("callback already scheduled for phone=%s: %s", phone, resp.text[:200]) + return {"status": "conflict", "detail": resp.text} + LOGGER.warning( + "callback schedule failed: status=%s body=%s", + resp.status_code, + resp.text[:300], + ) + return None + except Exception: + LOGGER.exception("callback client: request failed phone=%s", phone) + return None + + +async def cancel_callback(*, job_id: str) -> dict[str, Any] | None: + if not _enabled(): + return None + headers = _request_headers() + if headers is None: + return None + url = f"{_bridge_base_url()}/asterisk/callbacks/{job_id}" + try: + async with httpx.AsyncClient(timeout=8.0) as client: + resp = await client.delete(url, headers=headers) + if resp.status_code == 200: + return resp.json() + LOGGER.warning( + "callback cancel failed: job_id=%s status=%s body=%s", + job_id, + resp.status_code, + resp.text[:300], + ) + return None + except Exception: + LOGGER.exception("callback client: cancel failed job_id=%s", job_id) + return None diff --git a/core/session.py b/core/session.py index 891d85a..09ab7f9 100644 --- a/core/session.py +++ b/core/session.py @@ -245,6 +245,54 @@ def _preview_text(text: str, *, limit: int = 160) -> str: return f"{normalized[:limit]}..." +_KAZAKH_SPECIFIC_LETTERS = set("әғқңөұүһі") +_RUSSIAN_SPECIFIC_LETTERS = set("ыэъё") + + +def _detect_session_language(text: str) -> str | None: + normalized = str(text or "").strip().lower() + if not normalized: + return None + has_kazakh = any(ch in _KAZAKH_SPECIFIC_LETTERS for ch in normalized) + if has_kazakh: + return "kk" + has_russian = any(ch in _RUSSIAN_SPECIFIC_LETTERS for ch in normalized) + if has_russian: + return "ru" + cyrillic_chars = sum(1 for ch in normalized if "а" <= ch <= "я" or ch == "ё") + if cyrillic_chars >= 2: + return "ru" + return None + + +_TTS_LANGUAGE_CODE_MAP = { + "ru": "ru", + "kk": "kk", +} + +_STT_YANDEX_LANGUAGE_MAP = { + "ru": "ru-RU", + "kk": "kk-KZ", +} + +_STT_ELEVENLABS_LANGUAGE_MAP = { + "ru": "rus", + "kk": "kaz", +} + + +def _tts_language_for(session_language: str | None) -> str | None: + if not session_language: + return None + return _TTS_LANGUAGE_CODE_MAP.get(session_language.lower()) + + +def _stt_language_for(session_language: str | None) -> str | None: + if not session_language: + return None + return session_language.lower() if session_language.lower() in {"ru", "kk"} else None + + def _audio_duration_ms(audio_bytes: bytes, *, sample_rate_hz: int) -> int: if not audio_bytes or sample_rate_hz <= 0: return 0 @@ -993,6 +1041,7 @@ class CallSession: self._personalized_greeting_template = DEFAULT_PERSONALIZED_GREETING_TEMPLATE self._live_stt_stream: BaseSTTStream | None = None self._latest_partial_transcript = "" + self._session_language: str | None = None self._default_vad_silence_timeout_ms = getattr(self._vad, "default_speech_end_silence_ms", 550) self._semantic_hold_silence_timeout_ms = max(self._default_vad_silence_timeout_ms, SEMANTIC_ENDPOINTING_HOLD_MS) LOGGER.info( @@ -1235,6 +1284,18 @@ class CallSession: ) self._active_user_transcript = transcript + if self._session_language is None: + detected_language = _detect_session_language(transcript) + if detected_language: + self._session_language = detected_language + LOGGER.info( + "realtime session %s language locked: epoch=%s language=%s transcript=%r", + self.session_id, + epoch, + detected_language, + _preview_text(transcript), + ) + LOGGER.info( "realtime session %s transcript accepted: epoch=%s chars=%s text=%r", self.session_id, @@ -1334,7 +1395,11 @@ class CallSession: ), name=f"{self.session_id}-filler-delay-{epoch}", ) - async for event in self._llm.generate_stream(transcript, self._build_llm_context()): + async for event in self._llm.generate_stream( + transcript, + self._build_llm_context(), + language_code=self._session_language, + ): self._ensure_generation(epoch) if event.type == "tool_call_start": LOGGER.info( @@ -1546,7 +1611,11 @@ class CallSession: await self._stt.transcribe( audio_bytes, keyterms=self._name_capture_keyterms() if self._awaiting_customer_name else None, - language_code=self._name_capture_language_code() if self._awaiting_customer_name else None, + language_code=( + self._name_capture_language_code() + if self._awaiting_customer_name + else _stt_language_for(self._session_language) + ), force_batch=self._awaiting_customer_name, ) ).strip() @@ -1825,7 +1894,8 @@ class CallSession: sentence_queue=sentence_queue, started_holder=tts_started_monotonic, actually_spoken_chunks=actually_spoken_chunks, - ) + ), + language_code=_tts_language_for(self._session_language), ): self._ensure_generation(epoch) if not audio_chunk: @@ -2061,6 +2131,7 @@ class CallSession: live_stt_stream = await self._stt.start_stream( partial_callback=self._handle_partial_transcript, keyterms=None, + language_code=_stt_language_for(self._session_language), ) except Exception: LOGGER.exception("realtime session %s failed to start live STT stream", self.session_id) @@ -2137,7 +2208,11 @@ class CallSession: return await self._stt.transcribe( utterance_audio, keyterms=self._name_capture_keyterms() if self._awaiting_customer_name else None, - language_code=self._name_capture_language_code() if self._awaiting_customer_name else None, + language_code=( + self._name_capture_language_code() + if self._awaiting_customer_name + else _stt_language_for(self._session_language) + ), force_batch=self._awaiting_customer_name, ) diff --git a/providers/base.py b/providers/base.py index 878fed9..fa129cf 100644 --- a/providers/base.py +++ b/providers/base.py @@ -47,20 +47,32 @@ class BaseSTT(ABC): *, partial_callback: PartialTranscriptCallback | None = None, keyterms: Sequence[str] | None = None, + language_code: str | None = None, ) -> BaseSTTStream | None: - del partial_callback, keyterms + del partial_callback, keyterms, language_code return None class BaseLLM(ABC): @abstractmethod - async def generate_stream(self, text: str, context: list) -> AsyncGenerator[LLMStreamEvent, None]: + async def generate_stream( + self, + text: str, + context: list, + *, + language_code: str | None = None, + ) -> AsyncGenerator[LLMStreamEvent, None]: raise NotImplementedError class BaseTTS(ABC): @abstractmethod - async def synthesize_stream(self, text_stream: AsyncIterable[str]) -> AsyncGenerator[bytes, None]: + async def synthesize_stream( + self, + text_stream: AsyncIterable[str], + *, + language_code: str | None = None, + ) -> AsyncGenerator[bytes, None]: raise NotImplementedError @@ -106,8 +118,9 @@ class MockSTT(BaseSTT): *, partial_callback: PartialTranscriptCallback | None = None, keyterms: Sequence[str] | None = None, + language_code: str | None = None, ) -> BaseSTTStream | None: - del keyterms + del keyterms, language_code return _MockSTTStream(parent=self, partial_callback=partial_callback) @@ -121,8 +134,14 @@ class MockLLM(BaseLLM): self._token_delay_ms = max(token_delay_ms, 0) self._scripted_responses = list(scripted_responses or []) - async def generate_stream(self, text: str, context: list) -> AsyncGenerator[LLMStreamEvent, None]: - del context + async def generate_stream( + self, + text: str, + context: list, + *, + language_code: str | None = None, + ) -> AsyncGenerator[LLMStreamEvent, None]: + del context, language_code response_text = ( self._scripted_responses.pop(0) if self._scripted_responses @@ -152,7 +171,13 @@ class MockTTS(BaseTTS): self._amplitude = amplitude self._milliseconds_per_word = max(milliseconds_per_word, 40) - async def synthesize_stream(self, text_stream: AsyncIterable[str]) -> AsyncGenerator[bytes, None]: + async def synthesize_stream( + self, + text_stream: AsyncIterable[str], + *, + language_code: str | None = None, + ) -> AsyncGenerator[bytes, None]: + del language_code async for text in text_stream: normalized = text.strip() if not normalized: diff --git a/providers/llm.py b/providers/llm.py index 4a6e7e4..e552d5c 100644 --- a/providers/llm.py +++ b/providers/llm.py @@ -31,6 +31,45 @@ _WEB_SEARCH_POLICY = ( "Если для поиска не хватает города, даты или объекта, задай короткий уточняющий вопрос." ) _WEB_SEARCH_POLICY_MARKER = "если клиент спрашивает про погоду" +_CALLBACK_POLICY = ( + "Правило перезвона: если клиент просит перезвонить позже (через N минут, в определённое время или просит " + "связаться с ним позже), вызови инструмент schedule_callback с номером телефона клиента и временем. " + "Если номер не назван — попроси клиента продиктовать его. После успешного вызова инструмента подтверди " + "клиенту коротко на русском, в какое время мы перезвоним. Не вызывай инструмент несколько раз подряд " + "и не обещай перезвон до того, как инструмент вернул успех." +) +_CALLBACK_POLICY_MARKER = "если клиент просит перезвонить позже" +_LANGUAGE_DIRECTIVES = { + "ru": ( + "Язык диалога: русский. Отвечай только на русском языке, естественно и кратко, в одном спокойном " + "стиле для голоса. Не переключайся на другие языки, даже если в запросе клиента есть отдельные " + "слова на другом языке." + ), + "kk": ( + "Сөйлесу тілі: қазақ тілі. Тек қазақ тілінде, табиғи және қысқа жауап бер, дауыс үшін бір қалыпты " + "стильде. Клиенттің сөзінде бөтен тілден бөлек сөздер болса да, басқа тілге ауыспа." + ), +} + + +def _normalize_language_code(value: str | None) -> str | None: + if value is None: + return None + raw = str(value).strip().lower().replace("_", "-") + if not raw: + return None + if raw in {"ru", "rus", "ru-ru"}: + return "ru" + if raw in {"kk", "kaz", "kz", "kk-kz", "kz-kz"}: + return "kk" + return raw[:2] or None + + +def _language_directive(language_code: str | None) -> str | None: + normalized = _normalize_language_code(language_code) + if normalized is None: + return None + return _LANGUAGE_DIRECTIVES.get(normalized) _LIVE_SEARCH_PATTERN = re.compile( r"(" @@ -72,9 +111,11 @@ def _with_conversation_close_policy(system_prompt: str) -> str: def _with_runtime_policies(system_prompt: str) -> str: prompt = _with_conversation_close_policy(system_prompt) - if _WEB_SEARCH_POLICY_MARKER in prompt: - return prompt - return f"{prompt}\n\n{_WEB_SEARCH_POLICY}" if prompt else _WEB_SEARCH_POLICY + if _WEB_SEARCH_POLICY_MARKER not in prompt: + prompt = f"{prompt}\n\n{_WEB_SEARCH_POLICY}" if prompt else _WEB_SEARCH_POLICY + if _CALLBACK_POLICY_MARKER not in prompt: + prompt = f"{prompt}\n\n{_CALLBACK_POLICY}" if prompt else _CALLBACK_POLICY + return prompt def _should_force_live_search(text: str) -> bool: @@ -198,11 +239,17 @@ class OpenAILLM(BaseLLM): self._max_tool_roundtrips, ) - async def generate_stream(self, text: str, context: list) -> AsyncGenerator[LLMStreamEvent, None]: + async def generate_stream( + self, + text: str, + context: list, + *, + language_code: str | None = None, + ) -> AsyncGenerator[LLMStreamEvent, None]: if not self._api_key: raise RuntimeError("OPENAI_API_KEY is required for OpenAI LLM") - messages = self._build_messages(text, context) + messages = self._build_messages(text, context, language_code=language_code) force_live_search = _should_force_live_search(text) turn_started_monotonic = time.perf_counter() LOGGER.info( @@ -478,6 +525,8 @@ class OpenAILLM(BaseLLM): ) if name == "serper": result = await self._run_serper_tool(arguments) + elif name == "schedule_callback": + result = await self._run_schedule_callback_tool(arguments, tool_call_id=str(tool_call.get("id") or "")) else: result = f"Tool `{name}` is not supported by this runtime." LOGGER.info( @@ -547,6 +596,64 @@ class OpenAILLM(BaseLLM): LOGGER.info("Serper summary built: chars=%s preview=%r", len(summary), _preview_text(summary)) return summary + async def _run_schedule_callback_tool(self, arguments: dict[str, Any], *, tool_call_id: str) -> str: + from realtime_voice_service import callback_client + + phone = str(arguments.get("phone") or "").strip() + if not phone: + return "Не указан номер телефона. Попроси клиента продиктовать номер." + + in_minutes = arguments.get("in_minutes") + scheduled_for = arguments.get("scheduled_for") + if in_minutes is None and not scheduled_for: + return "Не указано время: передай in_minutes или scheduled_for." + + try: + in_minutes_int = int(in_minutes) if in_minutes is not None else None + except (TypeError, ValueError): + in_minutes_int = None + + reason = str(arguments.get("reason") or "").strip() or None + customer_name = str(arguments.get("customer_name") or "").strip() or None + idempotency_key = (tool_call_id or "")[:128] or None + + LOGGER.info( + "schedule_callback tool start: phone=%s in_minutes=%s scheduled_for=%s reason=%r", + phone, + in_minutes_int, + scheduled_for, + _preview_text(reason or ""), + ) + + result = await callback_client.schedule_callback( + phone=phone, + in_minutes=in_minutes_int, + scheduled_for=str(scheduled_for) if scheduled_for else None, + reason=reason, + customer_name=customer_name, + idempotency_key=idempotency_key, + ) + + if result is None: + return ( + "Не удалось запланировать перезвон: сервис недоступен. Извинись и предложи " + "клиенту перезвонить самостоятельно." + ) + if result.get("status") == "conflict": + return ( + "На этот номер уже запланирован активный перезвон. Сообщи клиенту, что " + "перезвон уже в очереди." + ) + + scheduled = result.get("scheduled_for") or scheduled_for or "soon" + phone_e164 = result.get("phone_e164") or phone + job_id = result.get("job_id") or "" + return ( + f"Перезвон запланирован: job_id={job_id}, номер {phone_e164}, время {scheduled} (UTC). " + "Подтверди клиенту коротко на русском, в какое примерно время мы перезвоним " + "(переведи UTC во время клиента, если знаешь его часовой пояс — иначе скажи в минутах от сейчас)." + ) + async def _get_serper_session(self): try: import aiohttp @@ -564,46 +671,96 @@ class OpenAILLM(BaseLLM): return self._serper_session def _build_tools(self) -> list[dict[str, Any]] | None: - if not self._enable_tools or not self._serper_api_key: + if not self._enable_tools: return None - return [ - { - "type": "function", - "function": { - "name": "serper", - "description": ( - "Search the public web for recent or external information when the user asks " - "about weather, current facts, news, prices, currency rates, websites, company data, " - "schedules, addresses, phone numbers, or anything requiring live search." - ), - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": "Precise search query to send to Serper.", - }, - "num": { - "type": "integer", - "description": "How many results to fetch, usually 3 to 5.", - "minimum": 1, - "maximum": 10, - }, - "hl": { - "type": "string", - "description": "UI language code, for example ru or en.", - }, - "gl": { - "type": "string", - "description": "Country code for result localization, for example kz or us.", + tools: list[dict[str, Any]] = [] + if self._serper_api_key: + tools.append( + { + "type": "function", + "function": { + "name": "serper", + "description": ( + "Search the public web for recent or external information when the user asks " + "about weather, current facts, news, prices, currency rates, websites, company data, " + "schedules, addresses, phone numbers, or anything requiring live search." + ), + "parameters": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "Precise search query to send to Serper.", + }, + "num": { + "type": "integer", + "description": "How many results to fetch, usually 3 to 5.", + "minimum": 1, + "maximum": 10, + }, + "hl": { + "type": "string", + "description": "UI language code, for example ru or en.", + }, + "gl": { + "type": "string", + "description": "Country code for result localization, for example kz or us.", + }, }, + "required": ["query"], + "additionalProperties": False, }, - "required": ["query"], - "additionalProperties": False, }, - }, - } - ] + } + ) + if _read_bool_env("CALLBACK_DISPATCH_ENABLED", True): + tools.append( + { + "type": "function", + "function": { + "name": "schedule_callback", + "description": ( + "Schedule an outbound callback to the customer when they ask to be called back later. " + "The system will dial the customer at the requested time and reconnect them with the AI. " + "Use only when the customer explicitly requests a callback. After the tool succeeds, " + "confirm the callback time briefly in Russian." + ), + "parameters": { + "type": "object", + "properties": { + "phone": { + "type": "string", + "description": ( + "Customer phone number. Accepts +7XXXXXXXXXX, 8XXXXXXXXXX, or any " + "format with digits — the backend normalizes it." + ), + }, + "in_minutes": { + "type": "integer", + "description": "How many minutes from now to call back. Use either in_minutes or scheduled_for.", + "minimum": 0, + "maximum": 10080, + }, + "scheduled_for": { + "type": "string", + "description": "ISO-8601 UTC datetime to dial at, e.g. 2026-05-02T14:30:00Z. Use either this or in_minutes.", + }, + "reason": { + "type": "string", + "description": "One short sentence in Russian: why the callback is needed (the customer's topic).", + }, + "customer_name": { + "type": "string", + "description": "Customer's name if already known from the conversation.", + }, + }, + "required": ["phone"], + "additionalProperties": False, + }, + }, + } + ) + return tools or None def _get_client(self): if self._client is not None: @@ -633,10 +790,19 @@ class OpenAILLM(BaseLLM): def _supports_custom_temperature(self) -> bool: return not self._model.lower().startswith("gpt-5") - def _build_messages(self, text: str, context: list) -> list[dict[str, Any]]: + def _build_messages( + self, + text: str, + context: list, + *, + language_code: str | None = None, + ) -> list[dict[str, Any]]: messages: list[dict[str, Any]] = [] if self._system_prompt: messages.append({"role": "system", "content": self._system_prompt}) + language_directive = _language_directive(language_code) + if language_directive: + messages.append({"role": "system", "content": language_directive}) context_messages: list[dict[str, Any]] = [] for entry in context: @@ -780,8 +946,14 @@ class OllamaLLM(BaseLLM): self._num_ctx or "default", ) - async def generate_stream(self, text: str, context: list) -> AsyncGenerator[LLMStreamEvent, None]: - messages = self._build_messages(text, context) + async def generate_stream( + self, + text: str, + context: list, + *, + language_code: str | None = None, + ) -> AsyncGenerator[LLMStreamEvent, None]: + messages = self._build_messages(text, context, language_code=language_code) options: dict[str, Any] = {"temperature": self._temperature} if self._num_predict > 0: options["num_predict"] = self._num_predict @@ -874,10 +1046,19 @@ class OllamaLLM(BaseLLM): self._client = httpx.AsyncClient(timeout=timeout) return self._client - def _build_messages(self, text: str, context: list) -> list[dict[str, Any]]: + def _build_messages( + self, + text: str, + context: list, + *, + language_code: str | None = None, + ) -> list[dict[str, Any]]: messages: list[dict[str, Any]] = [] if self._system_prompt: messages.append({"role": "system", "content": self._system_prompt}) + language_directive = _language_directive(language_code) + if language_directive: + messages.append({"role": "system", "content": language_directive}) context_messages: list[dict[str, Any]] = [] for entry in context: diff --git a/providers/stt.py b/providers/stt.py index ba72e78..8e455c0 100644 --- a/providers/stt.py +++ b/providers/stt.py @@ -663,15 +663,18 @@ class ElevenLabsSTT(BaseSTT): *, partial_callback: PartialTranscriptCallback | None = None, keyterms: Sequence[str] | None = None, + language_code: str | None = None, ) -> BaseSTTStream | None: if not self._use_realtime: return None if not self._api_key: raise RuntimeError("ELEVENLABS_API_KEY is required for ElevenLabs STT") + effective_language_code = _normalize_elevenlabs_language_code(language_code) or self._language_code websocket_url = self._build_realtime_websocket_url( sample_rate_hz=self._target_sample_rate_hz, keyterms=keyterms, + language_code=effective_language_code, ) normalized_keyterm_count = len( _normalize_keyterms( @@ -684,7 +687,7 @@ class ElevenLabsSTT(BaseSTT): "ElevenLabs STT live stream starting: realtime_model=%s sample_rate=%s language=%s keyterms=%s", self._realtime_model_id, self._target_sample_rate_hz, - self._language_code, + effective_language_code, normalized_keyterm_count, ) stream = ElevenLabsRealtimeSTTStream( @@ -1211,6 +1214,33 @@ class FallbackSTT(BaseSTT): force_batch=force_batch, ) + async def start_stream( + self, + *, + partial_callback: PartialTranscriptCallback | None = None, + keyterms: Sequence[str] | None = None, + language_code: str | None = None, + ) -> BaseSTTStream | None: + try: + stream = await self._primary.start_stream( + partial_callback=partial_callback, + keyterms=keyterms, + language_code=language_code, + ) + if stream is not None: + return stream + except Exception: + LOGGER.exception( + "Primary STT start_stream failed; falling back: primary=%s fallback=%s", + type(self._primary).__name__, + type(self._fallback).__name__, + ) + return await self._fallback.start_stream( + partial_callback=partial_callback, + keyterms=keyterms, + language_code=language_code, + ) + async def close(self) -> None: for provider in (self._primary, self._fallback): close = getattr(provider, "close", None) diff --git a/providers/tts.py b/providers/tts.py index ebe75c7..df3bd7c 100644 --- a/providers/tts.py +++ b/providers/tts.py @@ -166,7 +166,12 @@ class ElevenLabsTTS(BaseTTS): self._voice_settings, ) - async def synthesize_stream(self, text_stream: AsyncIterable[str]) -> AsyncGenerator[bytes, None]: + async def synthesize_stream( + self, + text_stream: AsyncIterable[str], + *, + language_code: str | None = None, + ) -> AsyncGenerator[bytes, None]: if not self._api_key: raise RuntimeError("ELEVENLABS_API_KEY is required for ElevenLabs TTS") if not self._voice_id: @@ -178,7 +183,8 @@ class ElevenLabsTTS(BaseTTS): except Exception as exc: # noqa: BLE001 raise RuntimeError("The `websockets` package is required for ElevenLabs TTS WebSocket streaming") from exc - websocket_url = self._build_websocket_url() + effective_language_code = (str(language_code).strip() if language_code else "") or self._language_code + websocket_url = self._build_websocket_url(language_code=effective_language_code) started_monotonic = time.perf_counter() audio_chunk_count = 0 audio_byte_count = 0 @@ -193,7 +199,7 @@ class ElevenLabsTTS(BaseTTS): self._output_format, self._provider_sample_rate_hz, self._target_sample_rate_hz, - self._language_code, + effective_language_code, self._auto_mode, ) try: @@ -290,7 +296,7 @@ class ElevenLabsTTS(BaseTTS): except OSError as exc: raise RuntimeError("ElevenLabs TTS WebSocket connection failed") from exc - def _build_websocket_url(self) -> str: + def _build_websocket_url(self, *, language_code: str | None = None) -> str: query = { "model_id": self._model_id, "output_format": self._output_format, @@ -299,8 +305,9 @@ class ElevenLabsTTS(BaseTTS): "sync_alignment": "false", "apply_text_normalization": "auto", } - if self._language_code: - query["language_code"] = self._language_code + effective_language = (str(language_code).strip() if language_code else "") or self._language_code + if effective_language: + query["language_code"] = effective_language encoded_voice_id = quote(self._voice_id, safe="") return f"{self._ws_base}/v1/text-to-speech/{encoded_voice_id}/stream-input?{urlencode(query)}"