From 5dda1e10a38fee52ec42ab6bd49cb906a92ade68 Mon Sep 17 00:00:00 2001 From: Magzhan Zhumabayev Date: Fri, 1 May 2026 18:08:46 +0500 Subject: [PATCH] . --- .env.example | 1 + core/session.py | 65 +++++++++++++++++++++++++ providers/base.py | 6 ++- providers/stt.py | 103 ++++++++++++++++++++++++++++++---------- providers/stt_openai.py | 23 +++++++-- 5 files changed, 169 insertions(+), 29 deletions(-) diff --git a/.env.example b/.env.example index dea51db..3b02d8f 100644 --- a/.env.example +++ b/.env.example @@ -33,6 +33,7 @@ OLLAMA_LLM_NUM_CTX=1024 STT_PROVIDER=openai STT_FALLBACK_PROVIDER=elevenlabs +STT_NAME_CAPTURE_LANGUAGE_CODE=kz STT_PROMPT=Это русская речь по телефону. Точно распознавай короткие ответы, имена, фамилии, слова благодарности и прощания. Особое внимание именам: Айнур, Айгерим, Алия, Данияр, Азамат, Арман, Мадина, Нурсултан, Руслан, Асель. Если человек представился одним словом, сохраняй его как имя. Номера телефонов записывай цифрами. TTS_PROVIDER=elevenlabs diff --git a/core/session.py b/core/session.py index df51367..718b813 100644 --- a/core/session.py +++ b/core/session.py @@ -557,6 +557,13 @@ def _sanitize_voice_text(text: str) -> str: DEFAULT_NAME_RETRY_TEXT = "Подскажите, пожалуйста, как я могу к вам обращаться?" DEFAULT_PERSONALIZED_GREETING_TEMPLATE = '{name}, я работаю на основе искусственного интеллекта и постараюсь максимально внимательно разобраться с вашим вопросом. Чем я могу помочь?' +_NAME_CAPTURE_CONFUSION_ALIASES = { + "парень": "Арнур", + "арнер": "Арнур", + "ар нур": "Арнур", + "шомарт": "Жомарт", + "шумарт": "Жомарт", +} def _voice_text_key(text: str | None) -> str: @@ -564,6 +571,9 @@ def _voice_text_key(text: str | None) -> str: return re.sub(r"\s+", " ", compact).strip() +_KAZAKH_NAME_BY_KEY = {_voice_text_key(name): name for name in TOP_KAZAKH_NAMES} + + def _voice_letters_key(text: str | None) -> str: return "".join(re.findall(r"[^\W\d_]+", str(text or "").lower(), flags=re.UNICODE)) @@ -582,6 +592,39 @@ def _looks_like_supported_name_text(text: str) -> bool: ) +def _correct_name_capture_candidate(name: str | None) -> str | None: + candidate = _canonical_name(str(name or "").strip()) + if not candidate: + return None + key = _voice_text_key(candidate) + exact_match = _KAZAKH_NAME_BY_KEY.get(key) + if exact_match: + return exact_match + alias_match = _NAME_CAPTURE_CONFUSION_ALIASES.get(key) + if alias_match: + return alias_match + + best_name = None + best_ratio = 0.0 + for gazetteer_key, gazetteer_name in _KAZAKH_NAME_BY_KEY.items(): + ratio = difflib.SequenceMatcher(None, key, gazetteer_key).ratio() + if ratio > best_ratio: + best_ratio = ratio + best_name = gazetteer_name + if best_name is None: + return candidate + threshold = 0.82 if len(key) >= 5 else 0.9 + if best_ratio >= threshold: + LOGGER.info( + "name capture candidate corrected via Kazakh names: candidate=%r corrected=%r ratio=%.3f", + candidate, + best_name, + best_ratio, + ) + return best_name + return candidate + + def _normalize_name_candidate(text: str | None) -> str | None: raw = str(text or "").strip(" \t\r\n,.;:!?\"'()[]{}") if not raw or any(ch.isdigit() for ch in raw): @@ -798,6 +841,16 @@ def _voice_start_name_outcome(text: str | None) -> tuple[str, str | None]: if not raw: return "name_not_obtained", None candidate, needs_followup = _extract_name_candidate(raw) + if candidate: + corrected_candidate = _correct_name_capture_candidate(candidate) or candidate + if corrected_candidate != candidate: + LOGGER.info( + "name capture candidate normalized: candidate=%r corrected=%r transcript=%r", + candidate, + corrected_candidate, + _preview_text(raw), + ) + candidate = corrected_candidate if candidate and not needs_followup: return "name_obtained", candidate if candidate: @@ -1493,6 +1546,8 @@ 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, + force_batch=self._awaiting_customer_name, ) ).strip() except Exception: @@ -1513,6 +1568,14 @@ class CallSession: def _name_capture_keyterms(self) -> tuple[str, ...]: return TOP_KAZAKH_NAMES + def _name_capture_language_code(self) -> str | None: + raw = ( + os.getenv("STT_NAME_CAPTURE_LANGUAGE_CODE", "").strip() + or os.getenv("STT_NAME_CAPTURE_LANGUAGE", "").strip() + or "kz" + ) + return raw or None + def _build_name_collection_response(self, transcript: str) -> str | None: if not self._awaiting_customer_name: return None @@ -2074,6 +2137,8 @@ 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, + force_batch=self._awaiting_customer_name, ) def _maybe_dump_audio(self, audio_bytes: bytes) -> None: diff --git a/providers/base.py b/providers/base.py index 02298ff..878fed9 100644 --- a/providers/base.py +++ b/providers/base.py @@ -37,6 +37,8 @@ class BaseSTT(ABC): audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, + language_code: str | None = None, + force_batch: bool = False, ) -> str: raise NotImplementedError @@ -88,8 +90,10 @@ class MockSTT(BaseSTT): audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, + language_code: str | None = None, + force_batch: bool = False, ) -> str: - del keyterms + del keyterms, language_code, force_batch await asyncio.sleep(self._latency_ms / 1000.0) self._call_count += 1 if self._scripted_transcripts: diff --git a/providers/stt.py b/providers/stt.py index 3caf690..739319c 100644 --- a/providers/stt.py +++ b/providers/stt.py @@ -52,6 +52,42 @@ def _normalize_keyterms( return normalized +def _normalize_elevenlabs_language_code(language_code: str | None) -> str | None: + normalized = str(language_code or "").strip() + if not normalized: + return None + lowered = normalized.lower().replace("_", "-") + mapping = { + "kz": "kaz", + "kk": "kaz", + "kaz": "kaz", + "kk-kz": "kaz", + "kz-kz": "kaz", + "ru": "rus", + "rus": "rus", + "ru-ru": "rus", + } + return mapping.get(lowered, normalized) + + +def _normalize_yandex_language_code(language_code: str | None) -> str | None: + normalized = str(language_code or "").strip() + if not normalized: + return None + lowered = normalized.lower().replace("_", "-") + mapping = { + "ru": "ru-RU", + "rus": "ru-RU", + "ru-ru": "ru-RU", + "kk": "kk-KZ", + "kaz": "kk-KZ", + "kz": "kk-KZ", + "kk-kz": "kk-KZ", + "kz-kz": "kk-KZ", + } + return mapping.get(lowered, normalized) + + def _api_base() -> str: return (os.getenv("ELEVENLABS_API_BASE", "https://api.elevenlabs.io").strip() or "https://api.elevenlabs.io").rstrip("/") @@ -107,18 +143,7 @@ def _yandex_language() -> str: or os.getenv("AI_VOICE_ASR_YANDEX_LANGUAGE", "").strip() or "ru-RU" ) - normalized = raw.lower().replace("_", "-") - mapping = { - "ru": "ru-RU", - "rus": "ru-RU", - "ru-ru": "ru-RU", - "kk": "kk-KZ", - "kaz": "kk-KZ", - "kz": "kk-KZ", - "kk-kz": "kk-KZ", - "kz-kz": "kk-KZ", - } - return mapping.get(normalized, raw) + return _normalize_yandex_language_code(raw) or "ru-RU" def _yandex_topic() -> str: @@ -562,21 +587,26 @@ class ElevenLabsSTT(BaseSTT): audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, + language_code: str | None = None, + force_batch: bool = False, ) -> str: if not audio_bytes: return "" 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 pcm_bytes, sample_rate_hz = self._extract_pcm(audio_bytes) LOGGER.info( "ElevenLabs STT transcribe start: input_bytes=%s extracted_pcm_bytes=%s sample_rate=%s duration_ms=%s " - "use_realtime=%s keyterms=%s", + "use_realtime=%s force_batch=%s language=%s keyterms=%s", len(audio_bytes), len(pcm_bytes), sample_rate_hz, _pcm_duration_ms(pcm_bytes, sample_rate_hz=sample_rate_hz), self._use_realtime, + force_batch, + effective_language_code, len(keyterms) if keyterms else 0, ) if sample_rate_hz != self._target_sample_rate_hz: @@ -596,12 +626,13 @@ class ElevenLabsSTT(BaseSTT): len(pcm_bytes), ) - if self._use_realtime: + if self._use_realtime and not force_batch: try: transcript = await self._transcribe_realtime( pcm_bytes=pcm_bytes, sample_rate_hz=sample_rate_hz, keyterms=keyterms, + language_code=effective_language_code, ) LOGGER.info( "ElevenLabs STT realtime transcript result: chars=%s transcript=%r", @@ -618,6 +649,7 @@ class ElevenLabsSTT(BaseSTT): pcm_bytes=pcm_bytes, sample_rate_hz=sample_rate_hz, keyterms=keyterms, + language_code=effective_language_code, ) LOGGER.info( "ElevenLabs STT batch transcript result: chars=%s transcript=%r", @@ -678,6 +710,7 @@ class ElevenLabsSTT(BaseSTT): pcm_bytes: bytes, sample_rate_hz: int, keyterms: Sequence[str] | None = None, + language_code: str | None = None, ) -> str: try: from websockets.exceptions import WebSocketException @@ -688,6 +721,7 @@ class ElevenLabsSTT(BaseSTT): websocket_url = self._build_realtime_websocket_url( sample_rate_hz=sample_rate_hz, keyterms=keyterms, + language_code=language_code, ) chunk_bytes = max(int(sample_rate_hz * self._realtime_chunk_duration_ms / 1000.0) * 2, 320) started_monotonic = time.perf_counter() @@ -790,6 +824,7 @@ class ElevenLabsSTT(BaseSTT): pcm_bytes: bytes, sample_rate_hz: int, keyterms: Sequence[str] | None = None, + language_code: str | None = None, ) -> str: # TODO: Architectural Bottleneck: Рассмотреть замену STT на Deepgram WebSocket API для достижения true-streaming latency. try: @@ -817,15 +852,16 @@ class ElevenLabsSTT(BaseSTT): len(pcm_bytes), len(wav_bytes), _pcm_duration_ms(pcm_bytes, sample_rate_hz=sample_rate_hz), - self._language_code, + language_code or self._language_code, len(normalized_keyterms), ) form = aiohttp.FormData() form.add_field("model_id", self._model_id) form.add_field("timestamps_granularity", "none") form.add_field("diarize", "false") - if self._language_code: - form.add_field("language_code", self._language_code) + effective_language_code = language_code or self._language_code + if effective_language_code: + form.add_field("language_code", effective_language_code) if normalized_keyterms: form.add_field("keyterms", json.dumps(normalized_keyterms, ensure_ascii=False)) form.add_field( @@ -927,6 +963,7 @@ class ElevenLabsSTT(BaseSTT): *, sample_rate_hz: int, keyterms: Sequence[str] | None = None, + language_code: str | None = None, ) -> str: query: list[tuple[str, str]] = [ ("model_id", self._realtime_model_id), @@ -934,8 +971,9 @@ class ElevenLabsSTT(BaseSTT): ("commit_strategy", "manual"), ("include_timestamps", "false"), ] - if self._language_code: - query.append(("language_code", self._language_code)) + effective_language_code = language_code or self._language_code + if effective_language_code: + query.append(("language_code", effective_language_code)) for term in _normalize_keyterms( keyterms, max_count=_REALTIME_MAX_KEYTERMS, @@ -1004,20 +1042,25 @@ class YandexSpeechKitSTT(BaseSTT): audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, + language_code: str | None = None, + force_batch: bool = False, ) -> str: - del keyterms + del keyterms, force_batch if not audio_bytes: return "" if not self._api_key and not self._iam_token: raise RuntimeError("YANDEX_STT_API_KEY or YANDEX_STT_IAM_TOKEN is required for Yandex STT") + effective_language = _normalize_yandex_language_code(language_code) or self._language pcm_bytes, sample_rate_hz = self._extract_pcm(audio_bytes) LOGGER.info( - "Yandex SpeechKit STT transcribe start: input_bytes=%s extracted_pcm_bytes=%s sample_rate=%s duration_ms=%s", + "Yandex SpeechKit STT transcribe start: input_bytes=%s extracted_pcm_bytes=%s sample_rate=%s " + "duration_ms=%s language=%s", len(audio_bytes), len(pcm_bytes), sample_rate_hz, _pcm_duration_ms(pcm_bytes, sample_rate_hz=sample_rate_hz), + effective_language, ) if sample_rate_hz != self._target_sample_rate_hz: before_rate_hz = sample_rate_hz @@ -1038,7 +1081,7 @@ class YandexSpeechKitSTT(BaseSTT): session = await self._get_client_session() params = { - "lang": self._language, + "lang": effective_language, "topic": self._topic, "format": "lpcm", "sampleRateHertz": str(sample_rate_hz), @@ -1145,16 +1188,28 @@ class FallbackSTT(BaseSTT): audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, + language_code: str | None = None, + force_batch: bool = False, ) -> str: try: - return await self._primary.transcribe(audio_bytes, keyterms=keyterms) + return await self._primary.transcribe( + audio_bytes, + keyterms=keyterms, + language_code=language_code, + force_batch=force_batch, + ) except Exception: LOGGER.exception( "Primary STT provider failed; falling back: primary=%s fallback=%s", type(self._primary).__name__, type(self._fallback).__name__, ) - return await self._fallback.transcribe(audio_bytes, keyterms=keyterms) + return await self._fallback.transcribe( + audio_bytes, + keyterms=keyterms, + language_code=language_code, + force_batch=force_batch, + ) async def close(self) -> None: for provider in (self._primary, self._fallback): diff --git a/providers/stt_openai.py b/providers/stt_openai.py index 1ba8064..e5581f0 100644 --- a/providers/stt_openai.py +++ b/providers/stt_openai.py @@ -15,6 +15,18 @@ from realtime_voice_service.providers.base import BaseSTT LOGGER = logging.getLogger("uvicorn.error") +def _normalize_openai_language_code(language_code: str | None) -> str | None: + normalized = str(language_code or "").strip() + if not normalized: + return None + lowered = normalized.lower() + if lowered in {"kz", "kk", "kaz", "kk-kz"}: + return "kk" + if lowered in {"ru", "rus", "ru-ru"}: + return "ru" + return normalized + + def _timeout_seconds() -> float: raw = os.getenv("OPENAI_STT_TIMEOUT_SECONDS") or os.getenv("OPENAI_TIMEOUT_SECONDS") if raw is None: @@ -100,13 +112,16 @@ class OpenAISTT(BaseSTT): audio_bytes: bytes, *, keyterms: Sequence[str] | None = None, + language_code: str | None = None, + force_batch: bool = False, ) -> str: - del keyterms + del keyterms, force_batch if not audio_bytes: return "" if not self._api_key: raise RuntimeError("OPENAI_API_KEY is required for OpenAI STT") + effective_language = _normalize_openai_language_code(language_code) or self._language wav_bytes = _pcm16le_to_wav_bytes(audio_bytes, sample_rate_hz=self._input_sample_rate_hz) audio_file = io.BytesIO(wav_bytes) audio_file.name = "utterance.wav" @@ -118,7 +133,7 @@ class OpenAISTT(BaseSTT): len(audio_bytes), len(wav_bytes), self._input_sample_rate_hz, - self._language, + effective_language, _preview_text(self._prompt), ) @@ -129,8 +144,8 @@ class OpenAISTT(BaseSTT): } if self._prompt: request["prompt"] = self._prompt - if self._language: - request["language"] = self._language + if effective_language: + request["language"] = effective_language try: response = await self._get_client().audio.transcriptions.create(**request)