Resolve initial STT language with separate prompts

This commit is contained in:
Magzhan Zhumabayev
2026-05-02 02:34:39 +05:00
parent ce1d4df696
commit 113db24c51
+121 -22
View File
@@ -284,6 +284,8 @@ _RUSSIAN_LANGUAGE_MARKERS = frozenset({
"орыс",
"орысша",
})
_ARABIC_SCRIPT_RE = re.compile(r"[\u0600-\u06ff]")
_CYRILLIC_SCRIPT_RE = re.compile(r"[А-Яа-яЁёӘәҒғҚқҢңӨөҰұҮүҺһІі]")
def _detect_preferred_session_language(text: str) -> str | None:
@@ -333,6 +335,51 @@ def _detect_session_language(text: str) -> str | None:
return None
def _stt_language_candidate_score(text: str, expected_language: str) -> int:
normalized_language = expected_language.lower()
normalized_text = str(text or "").strip().lower()
if not normalized_text:
return -100
score = 0
detected_preference = _detect_preferred_session_language(normalized_text)
if detected_preference == normalized_language:
score += 40
elif detected_preference:
score -= 40
if any(ch in _KAZAKH_SPECIFIC_LETTERS for ch in normalized_text):
score += 12 if normalized_language == "kk" else -8
if any(ch in _RUSSIAN_SPECIFIC_LETTERS for ch in normalized_text):
score += 8 if normalized_language == "ru" else -4
if _CYRILLIC_SCRIPT_RE.search(normalized_text):
score += 4
if _ARABIC_SCRIPT_RE.search(normalized_text):
score -= 30
latin_chars = sum(1 for ch in normalized_text if "a" <= ch <= "z")
cyrillic_chars = len(_CYRILLIC_SCRIPT_RE.findall(normalized_text))
if latin_chars > cyrillic_chars and latin_chars >= 6:
score -= 12
return score
def _choose_language_specific_transcript(ru_transcript: str, kk_transcript: str) -> tuple[str, str | None, int, int]:
ru_text = str(ru_transcript or "").strip()
kk_text = str(kk_transcript or "").strip()
ru_score = _stt_language_candidate_score(ru_text, "ru")
kk_score = _stt_language_candidate_score(kk_text, "kk")
if ru_score >= kk_score and ru_text:
detected_language = _detect_name_collection_language(ru_text) or _detect_session_language(ru_text)
return ru_text, detected_language, ru_score, kk_score
if kk_text:
detected_language = _detect_name_collection_language(kk_text) or _detect_session_language(kk_text)
return kk_text, detected_language, ru_score, kk_score
return ru_text or kk_text, None, ru_score, kk_score
_TTS_LANGUAGE_CODE_MAP = {
"ru": "ru",
"kk": "kk",
@@ -1714,18 +1761,7 @@ class CallSession:
len(audio_bytes),
_audio_duration_ms(audio_bytes, sample_rate_hz=self.transport.sample_rate_hz),
)
pending_transcript = (
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 _stt_language_for(self._session_language)
),
force_batch=self._awaiting_customer_name,
)
).strip()
pending_transcript = (await self._transcribe_batch_audio(audio_bytes)).strip()
except Exception:
LOGGER.exception("realtime session %s failed to transcribe pending unanswered audio", self.session_id)
continue
@@ -1751,6 +1787,78 @@ class CallSession:
)
return raw or None
async def _transcribe_batch_audio(self, audio_bytes: bytes) -> str:
if self._awaiting_customer_name and self._session_language is None and self._name_capture_language_code() is None:
return await self._transcribe_initial_language_choice(audio_bytes)
return 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 _stt_language_for(self._session_language)
),
force_batch=self._awaiting_customer_name,
)
async def _transcribe_initial_language_choice(self, audio_bytes: bytes) -> str:
LOGGER.info(
"realtime session %s resolving initial language choice with ru/kk STT requests",
self.session_id,
)
ru_task = self._stt.transcribe(
audio_bytes,
keyterms=self._name_capture_keyterms(),
language_code="ru",
force_batch=True,
)
kk_task = self._stt.transcribe(
audio_bytes,
keyterms=self._name_capture_keyterms(),
language_code="kk",
force_batch=True,
)
ru_raw, kk_raw = await asyncio.gather(ru_task, kk_task, return_exceptions=True)
if isinstance(ru_raw, Exception):
LOGGER.warning(
"realtime session %s initial Russian STT request failed: %s",
self.session_id,
ru_raw,
)
ru_result = ""
else:
ru_result = str(ru_raw or "")
if isinstance(kk_raw, Exception):
LOGGER.warning(
"realtime session %s initial Kazakh STT request failed: %s",
self.session_id,
kk_raw,
)
kk_result = ""
else:
kk_result = str(kk_raw or "")
transcript, detected_language, ru_score, kk_score = _choose_language_specific_transcript(ru_result, kk_result)
LOGGER.info(
"realtime session %s initial language STT selected: detected=%s ru_score=%s kk_score=%s "
"ru=%r kk=%r selected=%r",
self.session_id,
detected_language,
ru_score,
kk_score,
_preview_text(ru_result),
_preview_text(kk_result),
_preview_text(transcript),
)
if detected_language and self._session_language is None:
self._session_language = detected_language
LOGGER.info(
"realtime session %s language locked during initial language choice: language=%s transcript=%r",
self.session_id,
detected_language,
_preview_text(transcript),
)
return transcript
def _build_name_collection_response(self, transcript: str) -> str | None:
if not self._awaiting_customer_name:
return None
@@ -2386,16 +2494,7 @@ class CallSession:
len(utterance_audio),
_audio_duration_ms(utterance_audio, sample_rate_hz=self.transport.sample_rate_hz),
)
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 _stt_language_for(self._session_language)
),
force_batch=self._awaiting_customer_name,
)
return await self._transcribe_batch_audio(utterance_audio)
def _maybe_dump_audio(self, audio_bytes: bytes) -> None:
if str(os.getenv("ENABLE_AUDIO_DUMP", "")).strip().lower() not in {"1", "true", "yes", "on"}: