feat(voice): switch voice ASR to Yandex SpeechKit
This commit is contained in:
+12
-2
@@ -52,13 +52,13 @@ AI_VOICE_RUNTIME_SERVICE_URL=http://localhost:8018
|
||||
AI_VOICE_ENABLED=0
|
||||
AI_VOICE_POLICY_MODE=v2_fast_conversational
|
||||
AI_VOICE_QUEUE_CONFIG_JSON={}
|
||||
AI_VOICE_ASR_PROVIDER=openai
|
||||
AI_VOICE_ASR_PROVIDER=yandex
|
||||
AI_VOICE_TTS_PROVIDER=yandex
|
||||
AI_VOICE_V2_ENABLED=1
|
||||
AI_VOICE_V2_QUEUE_CODES=voice_lab_ai
|
||||
AI_VOICE_V2_ACK_MODE=immediate_short
|
||||
AI_VOICE_V2_DUPLEX_ENABLED=1
|
||||
AI_VOICE_V2_STREAMING_ASR_BACKEND=local_sidecar
|
||||
AI_VOICE_V2_STREAMING_ASR_BACKEND=yandex_speechkit
|
||||
AI_VOICE_V2_STREAMING_ASR_BASE_URL=http://127.0.0.1:8021
|
||||
AI_VOICE_V2_STREAMING_ASR_TIMEOUT_SECONDS=4
|
||||
AI_VOICE_V2_STREAMING_ASR_MODEL=base
|
||||
@@ -76,6 +76,16 @@ AI_VOICE_AUDIOSOCKET_ENABLED=0
|
||||
AI_VOICE_AUDIOSOCKET_HOST=0.0.0.0
|
||||
AI_VOICE_AUDIOSOCKET_PORT=9019
|
||||
AI_VOICE_ASR_MODEL=gpt-4o-mini-transcribe
|
||||
AI_VOICE_ASR_YANDEX_API_BASE=https://stt.api.ml.yandexcloud.kz
|
||||
AI_VOICE_ASR_YANDEX_API_VERSION=auto
|
||||
AI_VOICE_ASR_YANDEX_API_KEY=
|
||||
AI_VOICE_ASR_YANDEX_IAM_TOKEN=
|
||||
AI_VOICE_ASR_YANDEX_FOLDER_ID=
|
||||
AI_VOICE_ASR_YANDEX_LANGUAGE=ru-RU
|
||||
AI_VOICE_ASR_YANDEX_TOPIC=general
|
||||
AI_VOICE_ASR_YANDEX_SAMPLE_RATE_HZ=8000
|
||||
AI_VOICE_ASR_YANDEX_TIMEOUT_SECONDS=8
|
||||
AI_VOICE_ASR_YANDEX_POLL_INTERVAL_SECONDS=0.25
|
||||
AI_VOICE_TTS_MODEL=gpt-4o-mini-tts
|
||||
AI_VOICE_TTS_VOICE=alloy
|
||||
AI_VOICE_TTS_YANDEX_API_BASE=https://tts.api.ml.yandexcloud.kz
|
||||
|
||||
@@ -22,6 +22,9 @@ x-app-env: &app_env
|
||||
KB_SERVICE_URL: http://kb-service:8000
|
||||
REPORTING_SERVICE_URL: http://reporting-service:8000
|
||||
SUPERVISOR_SERVICE_URL: http://supervisor-service:8000
|
||||
AI_VOICE_ASR_PROVIDER: yandex
|
||||
AI_VOICE_TTS_PROVIDER: yandex
|
||||
AI_VOICE_V2_STREAMING_ASR_BACKEND: yandex_speechkit
|
||||
AI_VOICE_V2_STREAMING_ASR_BASE_URL: http://streaming-asr-sidecar:8021
|
||||
AI_VOICE_V2_STREAMING_ASR_TIMEOUT_SECONDS: "4"
|
||||
|
||||
|
||||
@@ -22,6 +22,9 @@ x-app-env: &app_env
|
||||
KB_SERVICE_URL: http://kb-service:8000
|
||||
REPORTING_SERVICE_URL: http://reporting-service:8000
|
||||
SUPERVISOR_SERVICE_URL: http://supervisor-service:8000
|
||||
AI_VOICE_ASR_PROVIDER: yandex
|
||||
AI_VOICE_TTS_PROVIDER: yandex
|
||||
AI_VOICE_V2_STREAMING_ASR_BACKEND: yandex_speechkit
|
||||
AI_VOICE_V2_STREAMING_ASR_BASE_URL: http://streaming-asr-sidecar:8021
|
||||
AI_VOICE_V2_STREAMING_ASR_TIMEOUT_SECONDS: "4"
|
||||
|
||||
|
||||
@@ -22,6 +22,9 @@ x-app-env: &app_env
|
||||
KB_SERVICE_URL: http://kb-service:8000
|
||||
REPORTING_SERVICE_URL: http://reporting-service:8000
|
||||
SUPERVISOR_SERVICE_URL: http://supervisor-service:8000
|
||||
AI_VOICE_ASR_PROVIDER: yandex
|
||||
AI_VOICE_TTS_PROVIDER: yandex
|
||||
AI_VOICE_V2_STREAMING_ASR_BACKEND: yandex_speechkit
|
||||
AI_VOICE_V2_STREAMING_ASR_BASE_URL: http://streaming-asr-sidecar:8021
|
||||
AI_VOICE_V2_STREAMING_ASR_TIMEOUT_SECONDS: "4"
|
||||
|
||||
|
||||
@@ -19,7 +19,7 @@ x-app-env: &app_env
|
||||
AI_VOICE_ENABLED: "0"
|
||||
AI_VOICE_QUEUE_CONFIG_JSON: "{}"
|
||||
AI_VOICE_POLICY_MODE: v2_fast_conversational
|
||||
AI_VOICE_ASR_PROVIDER: openai
|
||||
AI_VOICE_ASR_PROVIDER: yandex
|
||||
AI_VOICE_TTS_PROVIDER: yandex
|
||||
AI_VOICE_AUDIOSOCKET_ENABLED: "1"
|
||||
AI_VOICE_AUDIOSOCKET_HOST: 0.0.0.0
|
||||
@@ -35,7 +35,7 @@ x-app-env: &app_env
|
||||
AI_VOICE_V2_QUEUE_CODES: voice_lab_ai
|
||||
AI_VOICE_V2_ACK_MODE: immediate_short
|
||||
AI_VOICE_V2_DUPLEX_ENABLED: "1"
|
||||
AI_VOICE_V2_STREAMING_ASR_BACKEND: local_sidecar
|
||||
AI_VOICE_V2_STREAMING_ASR_BACKEND: yandex_speechkit
|
||||
AI_VOICE_V2_STREAMING_ASR_BASE_URL: http://streaming-asr-sidecar:8021
|
||||
AI_VOICE_V2_STREAMING_ASR_TIMEOUT_SECONDS: "4"
|
||||
AI_VOICE_V2_PREBAKED_ACK_ENABLED: "1"
|
||||
|
||||
@@ -47,8 +47,8 @@ recording:
|
||||
voiceAI:
|
||||
enabled: "0"
|
||||
queueConfigJson: "{}"
|
||||
asrProvider: openai
|
||||
ttsProvider: openai
|
||||
asrProvider: yandex
|
||||
ttsProvider: yandex
|
||||
audioSocketEnabled: "1"
|
||||
audioSocketHost: 0.0.0.0
|
||||
audioSocketPort: "9019"
|
||||
|
||||
@@ -309,14 +309,14 @@ def spawn_service(spec: dict[str, Any], runtime_dir: Path, data_dir: Path, base_
|
||||
env["AI_VOICE_RUNTIME_SERVICE_URL"] = env.get("AI_VOICE_RUNTIME_SERVICE_URL", "http://127.0.0.1:8018")
|
||||
env["AI_VOICE_ENABLED"] = env.get("AI_VOICE_ENABLED", "0")
|
||||
env["AI_VOICE_QUEUE_CONFIG_JSON"] = env.get("AI_VOICE_QUEUE_CONFIG_JSON", "{}")
|
||||
env["AI_VOICE_ASR_PROVIDER"] = env.get("AI_VOICE_ASR_PROVIDER", "openai")
|
||||
env["AI_VOICE_ASR_PROVIDER"] = env.get("AI_VOICE_ASR_PROVIDER", "yandex")
|
||||
env["AI_VOICE_TTS_PROVIDER"] = env.get("AI_VOICE_TTS_PROVIDER", "yandex")
|
||||
env["AI_VOICE_POLICY_MODE"] = env.get("AI_VOICE_POLICY_MODE", "v2_fast_conversational")
|
||||
env["AI_VOICE_V2_ENABLED"] = env.get("AI_VOICE_V2_ENABLED", "1")
|
||||
env["AI_VOICE_V2_QUEUE_CODES"] = env.get("AI_VOICE_V2_QUEUE_CODES", "voice_lab_ai")
|
||||
env["AI_VOICE_V2_ACK_MODE"] = env.get("AI_VOICE_V2_ACK_MODE", "immediate_short")
|
||||
env["AI_VOICE_V2_DUPLEX_ENABLED"] = env.get("AI_VOICE_V2_DUPLEX_ENABLED", "1")
|
||||
env["AI_VOICE_V2_STREAMING_ASR_BACKEND"] = env.get("AI_VOICE_V2_STREAMING_ASR_BACKEND", "local_sidecar")
|
||||
env["AI_VOICE_V2_STREAMING_ASR_BACKEND"] = env.get("AI_VOICE_V2_STREAMING_ASR_BACKEND", "yandex_speechkit")
|
||||
env["AI_VOICE_V2_STREAMING_ASR_BASE_URL"] = env.get("AI_VOICE_V2_STREAMING_ASR_BASE_URL", "http://127.0.0.1:8021")
|
||||
env["AI_VOICE_V2_STREAMING_ASR_TIMEOUT_SECONDS"] = env.get("AI_VOICE_V2_STREAMING_ASR_TIMEOUT_SECONDS", "4")
|
||||
env["AI_VOICE_V2_STREAMING_ASR_MODEL"] = env.get("AI_VOICE_V2_STREAMING_ASR_MODEL", "base")
|
||||
|
||||
@@ -84,7 +84,7 @@ def _interaction_service_url() -> str:
|
||||
|
||||
|
||||
def _asr_provider_name() -> str:
|
||||
return os.getenv("AI_VOICE_ASR_PROVIDER", "openai").strip() or "openai"
|
||||
return os.getenv("AI_VOICE_ASR_PROVIDER", "yandex").strip() or "yandex"
|
||||
|
||||
|
||||
def _tts_provider_name() -> str:
|
||||
@@ -164,7 +164,7 @@ def _voice_v2_duplex_enabled() -> bool:
|
||||
|
||||
|
||||
def _voice_v2_streaming_asr_backend() -> str:
|
||||
return str(os.getenv("AI_VOICE_V2_STREAMING_ASR_BACKEND", "local_sidecar") or "local_sidecar").strip().lower()
|
||||
return str(os.getenv("AI_VOICE_V2_STREAMING_ASR_BACKEND", "yandex_speechkit") or "yandex_speechkit").strip().lower()
|
||||
|
||||
|
||||
def _voice_v2_prebaked_ack_enabled() -> bool:
|
||||
|
||||
@@ -1,7 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
import wave
|
||||
from dataclasses import dataclass
|
||||
|
||||
import httpx
|
||||
@@ -30,6 +35,99 @@ def _openai_asr_model() -> str:
|
||||
return os.getenv("AI_VOICE_ASR_MODEL", "gpt-4o-mini-transcribe").strip() or "gpt-4o-mini-transcribe"
|
||||
|
||||
|
||||
def _yandex_asr_api_base() -> str:
|
||||
explicit = os.getenv("AI_VOICE_ASR_YANDEX_API_BASE", "").strip()
|
||||
if explicit:
|
||||
return explicit.rstrip("/")
|
||||
tts_base = os.getenv("AI_VOICE_TTS_YANDEX_API_BASE", "https://tts.api.ml.yandexcloud.kz").strip().lower()
|
||||
if "yandexcloud.kz" in tts_base:
|
||||
return "https://stt.api.ml.yandexcloud.kz"
|
||||
return "https://stt.api.cloud.yandex.net"
|
||||
|
||||
|
||||
def _yandex_asr_api_key() -> str:
|
||||
return (
|
||||
os.getenv("AI_VOICE_ASR_YANDEX_API_KEY", "").strip()
|
||||
or os.getenv("AI_VOICE_TTS_YANDEX_API_KEY", "").strip()
|
||||
)
|
||||
|
||||
|
||||
def _yandex_asr_iam_token() -> str:
|
||||
return (
|
||||
os.getenv("AI_VOICE_ASR_YANDEX_IAM_TOKEN", "").strip()
|
||||
or os.getenv("AI_VOICE_TTS_YANDEX_IAM_TOKEN", "").strip()
|
||||
)
|
||||
|
||||
|
||||
def _yandex_asr_folder_id() -> str:
|
||||
return (
|
||||
os.getenv("AI_VOICE_ASR_YANDEX_FOLDER_ID", "").strip()
|
||||
or os.getenv("AI_VOICE_TTS_YANDEX_FOLDER_ID", "").strip()
|
||||
)
|
||||
|
||||
|
||||
def _yandex_asr_timeout_seconds() -> float:
|
||||
raw = os.getenv("AI_VOICE_ASR_YANDEX_TIMEOUT_SECONDS", "").strip()
|
||||
if not raw:
|
||||
return _timeout_seconds()
|
||||
try:
|
||||
value = float(raw)
|
||||
except ValueError:
|
||||
value = _timeout_seconds()
|
||||
return max(value, 3.0)
|
||||
|
||||
|
||||
def _yandex_asr_sample_rate_hz() -> int:
|
||||
raw = os.getenv("AI_VOICE_ASR_YANDEX_SAMPLE_RATE_HZ", "8000").strip()
|
||||
try:
|
||||
value = int(raw)
|
||||
except ValueError:
|
||||
value = 8000
|
||||
if value < 8000 or value > 48000:
|
||||
return 8000
|
||||
return value
|
||||
|
||||
|
||||
def _yandex_asr_topic() -> str:
|
||||
return os.getenv("AI_VOICE_ASR_YANDEX_TOPIC", "general").strip() or "general"
|
||||
|
||||
|
||||
def _yandex_asr_api_version(api_base: str) -> str:
|
||||
raw = os.getenv("AI_VOICE_ASR_YANDEX_API_VERSION", "auto").strip().lower() or "auto"
|
||||
if raw in {"v1", "v1_sync", "sync"}:
|
||||
return "v1"
|
||||
if raw in {"v3", "v3_async", "async"}:
|
||||
return "v3"
|
||||
if "yandexcloud.kz" in str(api_base or "").lower():
|
||||
return "v3"
|
||||
return "v1"
|
||||
|
||||
|
||||
def _yandex_asr_poll_interval_seconds() -> float:
|
||||
raw = os.getenv("AI_VOICE_ASR_YANDEX_POLL_INTERVAL_SECONDS", "0.25").strip()
|
||||
try:
|
||||
value = float(raw)
|
||||
except ValueError:
|
||||
value = 0.25
|
||||
return max(min(value, 2.0), 0.05)
|
||||
|
||||
|
||||
def _yandex_asr_default_language() -> str:
|
||||
return os.getenv("AI_VOICE_ASR_YANDEX_LANGUAGE", "ru-RU").strip() or "ru-RU"
|
||||
|
||||
|
||||
def _normalize_yandex_asr_language(language: str | None) -> str:
|
||||
raw = str(language or "").strip()
|
||||
if not raw:
|
||||
return _yandex_asr_default_language()
|
||||
lowered = raw.replace("_", "-").lower()
|
||||
if lowered in {"ru", "ru-ru"}:
|
||||
return "ru-RU"
|
||||
if lowered in {"kk", "kz", "kk-kz", "kz-kz"}:
|
||||
return "kk-KZ"
|
||||
return raw
|
||||
|
||||
|
||||
def _streaming_asr_api_base() -> str:
|
||||
return (
|
||||
os.getenv("AI_VOICE_V2_STREAMING_ASR_BASE_URL", "http://127.0.0.1:8021").strip()
|
||||
@@ -152,6 +250,214 @@ class OpenAIASRProvider(ASRProvider):
|
||||
return self.transcribe(audio_bytes, language_hint=language_hint)
|
||||
|
||||
|
||||
class YandexSpeechKitASRProvider(ASRProvider):
|
||||
name = "yandex"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
iam_token: str | None = None,
|
||||
folder_id: str | None = None,
|
||||
timeout_seconds: float | None = None,
|
||||
sample_rate_hz: int | None = None,
|
||||
topic: str | None = None,
|
||||
) -> None:
|
||||
self._api_base = str(api_base or _yandex_asr_api_base()).strip().rstrip("/")
|
||||
self._api_key = str(api_key if api_key is not None else _yandex_asr_api_key()).strip()
|
||||
self._iam_token = str(iam_token if iam_token is not None else _yandex_asr_iam_token()).strip()
|
||||
self._folder_id = str(folder_id if folder_id is not None else _yandex_asr_folder_id()).strip()
|
||||
self._timeout_seconds = max(
|
||||
float(timeout_seconds if timeout_seconds is not None else _yandex_asr_timeout_seconds()),
|
||||
3.0,
|
||||
)
|
||||
self._sample_rate_hz = int(sample_rate_hz if sample_rate_hz is not None else _yandex_asr_sample_rate_hz())
|
||||
self._topic = str(topic if topic is not None else _yandex_asr_topic()).strip()
|
||||
self._api_version = _yandex_asr_api_version(self._api_base)
|
||||
self._poll_interval_seconds = _yandex_asr_poll_interval_seconds()
|
||||
|
||||
@staticmethod
|
||||
def _lpcm_from_audio_bytes(audio_bytes: bytes, *, default_sample_rate_hz: int) -> tuple[bytes, int]:
|
||||
if not audio_bytes:
|
||||
return b"", default_sample_rate_hz
|
||||
try:
|
||||
with wave.open(io.BytesIO(audio_bytes), "rb") as wav_file:
|
||||
sample_width = wav_file.getsampwidth()
|
||||
channels = wav_file.getnchannels()
|
||||
sample_rate = wav_file.getframerate()
|
||||
if sample_width != 2 or channels != 1:
|
||||
return audio_bytes, default_sample_rate_hz
|
||||
return wav_file.readframes(wav_file.getnframes()), int(sample_rate or default_sample_rate_hz)
|
||||
except (wave.Error, EOFError):
|
||||
return audio_bytes, default_sample_rate_hz
|
||||
|
||||
def _headers(self, *, content_type: str = "application/octet-stream") -> dict[str, str]:
|
||||
if not self._api_key and not self._iam_token:
|
||||
raise RuntimeError(
|
||||
"AI_VOICE_ASR_YANDEX_API_KEY or AI_VOICE_ASR_YANDEX_IAM_TOKEN is required for Yandex ASR"
|
||||
)
|
||||
headers = {"Content-Type": content_type}
|
||||
if self._api_key:
|
||||
headers["Authorization"] = f"Api-Key {self._api_key}"
|
||||
else:
|
||||
headers["Authorization"] = f"Bearer {self._iam_token}"
|
||||
if self._folder_id and self._api_version == "v3":
|
||||
headers["x-folder-id"] = self._folder_id
|
||||
return headers
|
||||
|
||||
def transcribe(self, audio_bytes: bytes, *, language_hint: str | None = None) -> ASRTranscription:
|
||||
pcm_bytes, sample_rate_hz = self._lpcm_from_audio_bytes(
|
||||
audio_bytes,
|
||||
default_sample_rate_hz=self._sample_rate_hz,
|
||||
)
|
||||
if not pcm_bytes:
|
||||
return ASRTranscription(text="", language=language_hint, confidence=None)
|
||||
|
||||
language = _normalize_yandex_asr_language(language_hint)
|
||||
if self._api_version == "v3":
|
||||
return self._transcribe_v3_async(pcm_bytes, sample_rate_hz=sample_rate_hz, language=language)
|
||||
return self._transcribe_v1_sync(pcm_bytes, sample_rate_hz=sample_rate_hz, language=language)
|
||||
|
||||
def _transcribe_v1_sync(self, pcm_bytes: bytes, *, sample_rate_hz: int, language: str) -> ASRTranscription:
|
||||
params: dict[str, str] = {
|
||||
"lang": language,
|
||||
"format": "lpcm",
|
||||
"sampleRateHertz": str(sample_rate_hz),
|
||||
}
|
||||
if self._topic:
|
||||
params["topic"] = self._topic
|
||||
if self._folder_id:
|
||||
params["folderId"] = self._folder_id
|
||||
|
||||
with httpx.Client(timeout=self._timeout_seconds) as client:
|
||||
response = client.post(
|
||||
f"{self._api_base}/speech/v1/stt:recognize",
|
||||
headers=self._headers(),
|
||||
params=params,
|
||||
content=pcm_bytes,
|
||||
)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
return ASRTranscription(
|
||||
text=str(payload.get("result") or "").strip(),
|
||||
language=language,
|
||||
confidence=None,
|
||||
)
|
||||
|
||||
def _transcribe_v3_async(self, pcm_bytes: bytes, *, sample_rate_hz: int, language: str) -> ASRTranscription:
|
||||
payload = {
|
||||
"content": base64.b64encode(pcm_bytes).decode("ascii"),
|
||||
"recognitionModel": {
|
||||
"model": self._topic or "general",
|
||||
"audioFormat": {
|
||||
"rawAudio": {
|
||||
"audioEncoding": "LINEAR16_PCM",
|
||||
"sampleRateHertz": str(sample_rate_hz),
|
||||
"audioChannelCount": "1",
|
||||
}
|
||||
},
|
||||
"textNormalization": {
|
||||
"textNormalization": "TEXT_NORMALIZATION_ENABLED",
|
||||
"profanityFilter": False,
|
||||
"literatureText": False,
|
||||
"phoneFormattingMode": "PHONE_FORMATTING_MODE_DISABLED",
|
||||
},
|
||||
"languageRestriction": {
|
||||
"restrictionType": "WHITELIST",
|
||||
"languageCode": [language],
|
||||
},
|
||||
"audioProcessingType": "FULL_DATA",
|
||||
},
|
||||
}
|
||||
headers = self._headers(content_type="application/json")
|
||||
deadline = time.monotonic() + self._timeout_seconds
|
||||
with httpx.Client(timeout=self._timeout_seconds) as client:
|
||||
response = client.post(
|
||||
f"{self._api_base}/stt/v3/recognizeFileAsync",
|
||||
headers=headers,
|
||||
json=payload,
|
||||
)
|
||||
response.raise_for_status()
|
||||
operation_payload = response.json()
|
||||
operation_id = str(operation_payload.get("id") or "").strip()
|
||||
if not operation_id:
|
||||
raise RuntimeError("Yandex ASR v3 response is missing operation id")
|
||||
|
||||
while not bool(operation_payload.get("done")):
|
||||
if time.monotonic() >= deadline:
|
||||
raise TimeoutError("Yandex ASR v3 recognition timed out")
|
||||
time.sleep(self._poll_interval_seconds)
|
||||
operation_response = client.get(
|
||||
f"{self._api_base}/operations/{operation_id}",
|
||||
headers=headers,
|
||||
)
|
||||
operation_response.raise_for_status()
|
||||
operation_payload = operation_response.json()
|
||||
|
||||
error_payload = operation_payload.get("error")
|
||||
if isinstance(error_payload, dict) and error_payload:
|
||||
raise RuntimeError(str(error_payload.get("message") or error_payload)[:500])
|
||||
|
||||
result_response = client.get(
|
||||
f"{self._api_base}/stt/v3/getRecognition",
|
||||
headers=headers,
|
||||
params={"operationId": operation_id},
|
||||
)
|
||||
result_response.raise_for_status()
|
||||
result_payload = result_response.json()
|
||||
|
||||
return ASRTranscription(
|
||||
text=self._extract_v3_text(result_payload),
|
||||
language=language,
|
||||
confidence=None,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _extract_v3_text(cls, payload: object) -> str:
|
||||
if isinstance(payload, list):
|
||||
texts = [cls._extract_v3_text(item) for item in payload]
|
||||
return " ".join(text for text in texts if text).strip()
|
||||
if not isinstance(payload, dict):
|
||||
return ""
|
||||
|
||||
final_refinement = payload.get("finalRefinement")
|
||||
if isinstance(final_refinement, dict):
|
||||
normalized = final_refinement.get("normalizedText")
|
||||
text = cls._extract_alternatives_text(normalized)
|
||||
if text:
|
||||
return text
|
||||
|
||||
for key in ("final", "partial"):
|
||||
text = cls._extract_alternatives_text(payload.get(key))
|
||||
if text:
|
||||
return text
|
||||
|
||||
response = payload.get("response")
|
||||
if isinstance(response, (dict, list)):
|
||||
return cls._extract_v3_text(response)
|
||||
return ""
|
||||
|
||||
@staticmethod
|
||||
def _extract_alternatives_text(payload: object) -> str:
|
||||
if not isinstance(payload, dict):
|
||||
return ""
|
||||
alternatives = payload.get("alternatives")
|
||||
if not isinstance(alternatives, list):
|
||||
return ""
|
||||
texts: list[str] = []
|
||||
for alternative in alternatives:
|
||||
if not isinstance(alternative, dict):
|
||||
continue
|
||||
text = str(alternative.get("text") or "").strip()
|
||||
if text:
|
||||
texts.append(text)
|
||||
return " ".join(texts).strip()
|
||||
|
||||
def transcribe_partial(self, audio_bytes: bytes, *, language_hint: str | None = None) -> ASRTranscription:
|
||||
return self.transcribe(audio_bytes, language_hint=language_hint)
|
||||
|
||||
|
||||
class LocalSidecarStreamingASRProvider(StreamingASRProvider):
|
||||
name = "local-sidecar"
|
||||
supports_streaming = True
|
||||
@@ -245,10 +551,60 @@ class LocalSidecarStreamingASRProvider(StreamingASRProvider):
|
||||
return
|
||||
|
||||
|
||||
class YandexSpeechKitBufferedStreamingASRProvider(StreamingASRProvider):
|
||||
name = "yandex-speechkit-buffered"
|
||||
supports_streaming = True
|
||||
|
||||
def __init__(self, *, asr_provider: YandexSpeechKitASRProvider | None = None) -> None:
|
||||
self._asr_provider = asr_provider or YandexSpeechKitASRProvider()
|
||||
self._streams: dict[str, dict[str, object]] = {}
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def open_stream(self, session_id: str, *, language_hint: str | None = None) -> str:
|
||||
stream_id = f"yasr_{uuid.uuid4().hex}"
|
||||
with self._lock:
|
||||
self._streams[stream_id] = {
|
||||
"session_id": str(session_id or "").strip(),
|
||||
"language_hint": str(language_hint or "").strip() or None,
|
||||
"pcm": bytearray(),
|
||||
}
|
||||
return stream_id
|
||||
|
||||
def push_pcm(self, stream_id: str, pcm_8k_chunk: bytes) -> None:
|
||||
if not pcm_8k_chunk:
|
||||
return
|
||||
with self._lock:
|
||||
stream = self._streams.get(stream_id)
|
||||
if stream is None:
|
||||
raise StreamingASRUnavailable("Yandex ASR stream is not active")
|
||||
pcm = stream.get("pcm")
|
||||
if isinstance(pcm, bytearray):
|
||||
pcm.extend(pcm_8k_chunk)
|
||||
|
||||
def poll_partial(self, stream_id: str) -> StreamingASRPartial | None:
|
||||
del stream_id
|
||||
return None
|
||||
|
||||
def finalize(self, stream_id: str) -> ASRTranscription:
|
||||
with self._lock:
|
||||
stream = self._streams.get(stream_id)
|
||||
if stream is None:
|
||||
raise StreamingASRUnavailable("Yandex ASR stream is not active")
|
||||
pcm = bytes(stream.get("pcm") or b"")
|
||||
language_hint = str(stream.get("language_hint") or "").strip() or None
|
||||
return self._asr_provider.transcribe(pcm, language_hint=language_hint)
|
||||
|
||||
def close_stream(self, stream_id: str) -> None:
|
||||
with self._lock:
|
||||
self._streams.pop(stream_id, None)
|
||||
|
||||
|
||||
def build_asr_provider(name: str) -> ASRProvider:
|
||||
normalized = str(name or "stub").strip().lower()
|
||||
if normalized == "openai":
|
||||
return OpenAIASRProvider()
|
||||
if normalized in {"yandex", "yandex_speechkit", "speechkit"}:
|
||||
return YandexSpeechKitASRProvider()
|
||||
return ASRProvider()
|
||||
|
||||
|
||||
@@ -256,4 +612,6 @@ def build_streaming_asr_provider(name: str) -> StreamingASRProvider:
|
||||
normalized = str(name or "disabled").strip().lower()
|
||||
if normalized in {"local_sidecar", "local-sidecar", "sidecar"}:
|
||||
return LocalSidecarStreamingASRProvider()
|
||||
if normalized in {"yandex", "yandex_speechkit", "yandex-speechkit", "speechkit"}:
|
||||
return YandexSpeechKitBufferedStreamingASRProvider()
|
||||
return StreamingASRProvider()
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from services.ai_voice_runtime_service.audiosocket import pcm16le_to_wav_bytes
|
||||
from services.ai_voice_runtime_service.providers import asr as asr_module
|
||||
|
||||
|
||||
class _DummyResponse:
|
||||
def __init__(self, json_payload: dict | None = None) -> None:
|
||||
self.content = b"{}"
|
||||
self._json_payload = json_payload or {}
|
||||
|
||||
def raise_for_status(self) -> None:
|
||||
return None
|
||||
|
||||
def json(self) -> dict:
|
||||
return self._json_payload
|
||||
|
||||
|
||||
def test_yandex_asr_provider_posts_lpcm_with_tts_credential_fallback(monkeypatch):
|
||||
calls: list[dict] = []
|
||||
pcm = b"\x01\x00" * 160
|
||||
wav_bytes = pcm16le_to_wav_bytes(pcm, sample_rate_hz=8000)
|
||||
|
||||
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],
|
||||
content: bytes,
|
||||
) -> _DummyResponse:
|
||||
calls.append(
|
||||
{
|
||||
"url": url,
|
||||
"headers": headers,
|
||||
"params": params,
|
||||
"content": content,
|
||||
"timeout": self.timeout,
|
||||
}
|
||||
)
|
||||
return _DummyResponse({"result": "almaty schedule"})
|
||||
|
||||
monkeypatch.delenv("AI_VOICE_ASR_YANDEX_API_KEY", raising=False)
|
||||
monkeypatch.delenv("AI_VOICE_ASR_YANDEX_IAM_TOKEN", raising=False)
|
||||
monkeypatch.setenv("AI_VOICE_TTS_YANDEX_API_KEY", "tts-yandex-key")
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_API_BASE", "https://stt.example.test")
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_FOLDER_ID", "folder-test")
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_TIMEOUT_SECONDS", "5")
|
||||
monkeypatch.setattr(asr_module.httpx, "Client", _DummyClient)
|
||||
|
||||
provider = asr_module.YandexSpeechKitASRProvider()
|
||||
result = provider.transcribe(wav_bytes, language_hint="ru")
|
||||
|
||||
assert result.text == "almaty schedule"
|
||||
assert result.language == "ru-RU"
|
||||
assert len(calls) == 1
|
||||
assert calls[0]["url"] == "https://stt.example.test/speech/v1/stt:recognize"
|
||||
assert calls[0]["headers"]["Authorization"] == "Api-Key tts-yandex-key"
|
||||
assert calls[0]["headers"]["Content-Type"] == "application/octet-stream"
|
||||
assert calls[0]["params"]["lang"] == "ru-RU"
|
||||
assert calls[0]["params"]["format"] == "lpcm"
|
||||
assert calls[0]["params"]["sampleRateHertz"] == "8000"
|
||||
assert calls[0]["params"]["topic"] == "general"
|
||||
assert calls[0]["params"]["folderId"] == "folder-test"
|
||||
assert calls[0]["content"] == pcm
|
||||
assert calls[0]["timeout"] == 5.0
|
||||
|
||||
|
||||
def test_yandex_asr_provider_uses_iam_token(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],
|
||||
content: bytes,
|
||||
) -> _DummyResponse:
|
||||
calls.append({"url": url, "headers": headers, "params": params, "content": content})
|
||||
return _DummyResponse({"result": "operator"})
|
||||
|
||||
monkeypatch.delenv("AI_VOICE_ASR_YANDEX_API_KEY", raising=False)
|
||||
monkeypatch.delenv("AI_VOICE_TTS_YANDEX_API_KEY", raising=False)
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_IAM_TOKEN", "iam-token")
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_API_BASE", "https://stt.example.test")
|
||||
monkeypatch.setattr(asr_module.httpx, "Client", _DummyClient)
|
||||
|
||||
provider = asr_module.YandexSpeechKitASRProvider()
|
||||
result = provider.transcribe(b"\x02\x00" * 80, language_hint="kk")
|
||||
|
||||
assert result.text == "operator"
|
||||
assert result.language == "kk-KZ"
|
||||
assert calls[0]["headers"]["Authorization"] == "Bearer iam-token"
|
||||
assert calls[0]["params"]["lang"] == "kk-KZ"
|
||||
|
||||
|
||||
def test_yandex_asr_provider_supports_kz_v3_async_rest(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], json: dict) -> _DummyResponse:
|
||||
calls.append({"method": "POST", "url": url, "headers": headers, "json": json})
|
||||
return _DummyResponse({"id": "operation-1", "done": True})
|
||||
|
||||
def get(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
headers: dict[str, str],
|
||||
params: dict[str, str] | None = None,
|
||||
) -> _DummyResponse:
|
||||
calls.append({"method": "GET", "url": url, "headers": headers, "params": params or {}})
|
||||
return _DummyResponse(
|
||||
{
|
||||
"finalRefinement": {
|
||||
"normalizedText": {
|
||||
"alternatives": [
|
||||
{
|
||||
"text": "almaty schedule",
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_API_KEY", "asr-key")
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_API_BASE", "https://stt.api.ml.yandexcloud.kz")
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_FOLDER_ID", "folder-test")
|
||||
monkeypatch.setenv("AI_VOICE_ASR_YANDEX_POLL_INTERVAL_SECONDS", "0.05")
|
||||
monkeypatch.setattr(asr_module.httpx, "Client", _DummyClient)
|
||||
|
||||
provider = asr_module.YandexSpeechKitASRProvider()
|
||||
result = provider.transcribe(b"\x03\x00" * 80, language_hint="ru")
|
||||
|
||||
assert result.text == "almaty schedule"
|
||||
assert result.language == "ru-RU"
|
||||
assert calls[0]["url"] == "https://stt.api.ml.yandexcloud.kz/stt/v3/recognizeFileAsync"
|
||||
assert calls[0]["headers"]["Authorization"] == "Api-Key asr-key"
|
||||
assert calls[0]["headers"]["Content-Type"] == "application/json"
|
||||
assert calls[0]["headers"]["x-folder-id"] == "folder-test"
|
||||
assert calls[0]["json"]["recognitionModel"]["audioFormat"]["rawAudio"]["sampleRateHertz"] == "8000"
|
||||
assert calls[0]["json"]["recognitionModel"]["languageRestriction"]["languageCode"] == ["ru-RU"]
|
||||
assert calls[1]["url"] == "https://stt.api.ml.yandexcloud.kz/stt/v3/getRecognition"
|
||||
assert calls[1]["params"] == {"operationId": "operation-1"}
|
||||
|
||||
|
||||
def test_yandex_buffered_streaming_provider_buffers_pcm_until_finalize():
|
||||
calls: list[dict] = []
|
||||
|
||||
class _FakeYandexASR:
|
||||
def transcribe(self, audio_bytes: bytes, *, language_hint: str | None = None):
|
||||
calls.append({"audio_bytes": audio_bytes, "language_hint": language_hint})
|
||||
return asr_module.ASRTranscription(text="need schedule", language=language_hint, confidence=None)
|
||||
|
||||
provider = asr_module.YandexSpeechKitBufferedStreamingASRProvider(asr_provider=_FakeYandexASR()) # type: ignore[arg-type]
|
||||
stream_id = provider.open_stream("session-1", language_hint="ru")
|
||||
|
||||
provider.push_pcm(stream_id, b"\x10\x00")
|
||||
provider.push_pcm(stream_id, b"\x20\x00")
|
||||
|
||||
assert provider.poll_partial(stream_id) is None
|
||||
|
||||
result = provider.finalize(stream_id)
|
||||
provider.close_stream(stream_id)
|
||||
|
||||
assert result.text == "need schedule"
|
||||
assert calls == [{"audio_bytes": b"\x10\x00\x20\x00", "language_hint": "ru"}]
|
||||
|
||||
|
||||
def test_yandex_asr_builders():
|
||||
assert isinstance(asr_module.build_asr_provider("yandex"), asr_module.YandexSpeechKitASRProvider)
|
||||
assert isinstance(asr_module.build_asr_provider("speechkit"), asr_module.YandexSpeechKitASRProvider)
|
||||
assert isinstance(
|
||||
asr_module.build_streaming_asr_provider("yandex_speechkit"),
|
||||
asr_module.YandexSpeechKitBufferedStreamingASRProvider,
|
||||
)
|
||||
Reference in New Issue
Block a user