diff --git a/.env.example b/.env.example index 50b94cd..a8a43b2 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/deployment/docker-compose.parallel.server.yml b/deployment/docker-compose.parallel.server.yml index 85a8a2c..bdec4f1 100644 --- a/deployment/docker-compose.parallel.server.yml +++ b/deployment/docker-compose.parallel.server.yml @@ -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" diff --git a/deployment/docker-compose.server.registry.yml b/deployment/docker-compose.server.registry.yml index eee6e53..7d06b06 100644 --- a/deployment/docker-compose.server.registry.yml +++ b/deployment/docker-compose.server.registry.yml @@ -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" diff --git a/deployment/docker-compose.server.yml b/deployment/docker-compose.server.yml index be92bd8..b21d236 100644 --- a/deployment/docker-compose.server.yml +++ b/deployment/docker-compose.server.yml @@ -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" diff --git a/deployment/docker-compose.yml b/deployment/docker-compose.yml index ea37f12..f5e5aba 100644 --- a/deployment/docker-compose.yml +++ b/deployment/docker-compose.yml @@ -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" diff --git a/deployment/helm/values.yaml b/deployment/helm/values.yaml index a4e575a..85fdf58 100644 --- a/deployment/helm/values.yaml +++ b/deployment/helm/values.yaml @@ -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" diff --git a/scripts/local_stack.py b/scripts/local_stack.py index a875fe6..49c3b30 100644 --- a/scripts/local_stack.py +++ b/scripts/local_stack.py @@ -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") diff --git a/services/ai_voice_runtime_service/app.py b/services/ai_voice_runtime_service/app.py index ed2eecf..aa28722 100644 --- a/services/ai_voice_runtime_service/app.py +++ b/services/ai_voice_runtime_service/app.py @@ -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: diff --git a/services/ai_voice_runtime_service/providers/asr.py b/services/ai_voice_runtime_service/providers/asr.py index bb55467..66b30f8 100644 --- a/services/ai_voice_runtime_service/providers/asr.py +++ b/services/ai_voice_runtime_service/providers/asr.py @@ -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() diff --git a/tests/test_ai_voice_asr_provider.py b/tests/test_ai_voice_asr_provider.py new file mode 100644 index 0000000..b03d350 --- /dev/null +++ b/tests/test_ai_voice_asr_provider.py @@ -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, + )