from __future__ import annotations import io import base64 import json import re import socket import urllib.error import urllib.request import wave from collections.abc import Callable from pathlib import Path from typing import Any from .config import AppConfig from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, Transcript from .speech_models import stt_model_paths _MEANINGFUL_TEXT = re.compile(r"[\w\u4e00-\u9fff]", re.UNICODE) _PARTIAL_MIN_MEANINGFUL_CHARS = 3 def is_valid_transcript_text(text: str) -> bool: return bool(_MEANINGFUL_TEXT.search(text.strip())) def _meaningful_text_length(text: str) -> int: return len(_MEANINGFUL_TEXT.findall(text.strip())) def should_emit_partial_transcript(text: str, last_text: str) -> bool: normalized = text.strip() if not is_valid_transcript_text(normalized) or _meaningful_text_length(normalized) < _PARTIAL_MIN_MEANINGFUL_CHARS: return False if normalized == last_text: return False if not last_text: return True return normalized.startswith(last_text) class MetadataSttProvider: def __init__(self, language: str = "zh") -> None: self.language = language self.loaded = False def load(self) -> None: self.loaded = True def transcribe(self, segment: AudioSegment) -> Transcript: if not self.loaded: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, "STT provider is not loaded", False, "metadata-stt", "stt", ) text = str(segment.metadata.get("transcript", "")).strip() if not is_valid_transcript_text(text): raise ProviderError( ErrorCode.STT_EMPTY_TRANSCRIPT, "STT produced no meaningful text", True, "metadata-stt", "stt", ) return Transcript( text=text, language=str(segment.metadata.get("language", self.language)), confidence=float(segment.metadata.get("stt_confidence", 1.0)), duration_ms=segment.duration_ms, provider="metadata-stt", raw_metadata=dict(segment.metadata), ) def start_stream(self) -> "MetadataRealtimeTranscriptSession": return MetadataRealtimeTranscriptSession(self.language) class MetadataRealtimeTranscriptSession: def __init__(self, language: str = "zh") -> None: self.language = language self._last_text = "" def accept_frame(self, frame: AudioFrame) -> Transcript | None: text = str( frame.metadata.get("partial_transcript") or frame.metadata.get("transcript") or "" ).strip() if not should_emit_partial_transcript(text, self._last_text): return None self._last_text = text return Transcript( text=text, language=str(frame.metadata.get("language", self.language)), confidence=float(frame.metadata.get("stt_confidence", 1.0)), duration_ms=int(frame.metadata.get("duration_ms", 20)), provider="metadata-stt-partial", raw_metadata=dict(frame.metadata), ) def finish(self) -> Transcript | None: return None class CloudAsrSttProvider: def __init__( self, config: AppConfig, timeout_s: float = 60.0, urlopen: Callable[..., Any] | None = None, ) -> None: self.config = config self.timeout_s = timeout_s self.urlopen = urlopen or urllib.request.urlopen self.loaded = False def load(self) -> None: self.config.require_llm_credentials() self.loaded = True def transcribe(self, segment: AudioSegment) -> Transcript: if not self.loaded: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, "cloud ASR provider is not loaded", False, "newapi-asr", "stt", ) audio_data = base64.b64encode(_segment_to_wav_bytes(segment)).decode("ascii") payload = { "model": self.config.asr_model, "messages": [ { "role": "user", "content": [ { "type": "input_audio", "input_audio": {"data": f"data:audio/wav;base64,{audio_data}"}, } ], } ], "asr_options": {"language": "auto"}, } request = urllib.request.Request( self.config.api_url("/v1/chat/completions"), data=json.dumps(payload).encode("utf-8"), headers={ "Authorization": f"Bearer {self.config.llm_api_key}", "Content-Type": "application/json", }, method="POST", ) try: with self.urlopen(request, timeout=self.timeout_s) as response: payload = json.loads(response.read().decode("utf-8")) except urllib.error.HTTPError as exc: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, f"cloud ASR HTTP error {exc.code}", exc.code >= 500, "newapi-asr", "stt", ) from exc except (urllib.error.URLError, TimeoutError, socket.timeout, json.JSONDecodeError) as exc: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, f"cloud ASR request failed: {exc}", True, "newapi-asr", "stt", ) from exc choices = payload.get("choices") or [] message = (choices[0].get("message") if choices else {}) or {} text = str(message.get("content") or payload.get("text") or "").strip() if not is_valid_transcript_text(text): raise ProviderError( ErrorCode.STT_EMPTY_TRANSCRIPT, "cloud ASR produced no meaningful text", True, "newapi-asr", "stt", ) return Transcript( text=text, language=str(payload.get("language") or "zh"), confidence=None, duration_ms=segment.duration_ms, provider="newapi-asr", raw_metadata={"model": self.config.asr_model}, ) class SherpaOnnxSttProvider: def __init__(self, model_path: str, language: str = "zh", sherpa_module: Any | None = None) -> None: self.model_path = Path(model_path) self.language = language self.loaded = False self._sherpa = sherpa_module self._recognizer: Any | None = None def load(self) -> None: if not self.model_path.exists(): raise ProviderError( ErrorCode.STT_MODEL_MISSING, f"sherpa-onnx STT model path does not exist: {self.model_path}", False, "sherpa-onnx-stt", "stt", ) paths = stt_model_paths(self.model_path) missing = [ name for name, path in paths.items() if name not in {"type", "model_dir"} and isinstance(path, Path) and not path.exists() ] if missing: raise ProviderError( ErrorCode.STT_MODEL_MISSING, "sherpa-onnx STT model files are missing: " + ", ".join(missing), False, "sherpa-onnx-stt", "stt", ) sherpa_onnx = self._sherpa if sherpa_onnx is None: try: import sherpa_onnx # type: ignore[import-not-found] except Exception as exc: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, f"sherpa_onnx is not available: {exc}", False, "sherpa-onnx-stt", "stt", ) from exc try: model_type = str(paths.get("type", "sherpa-onnx-streaming-transducer")) if model_type == "sherpa-onnx-streaming-zipformer2-ctc": self._recognizer = sherpa_onnx.OnlineRecognizer.from_zipformer2_ctc( tokens=str(paths["tokens"]), model=str(paths["model"]), num_threads=1, sample_rate=16000, feature_dim=80, enable_endpoint_detection=True, decoding_method="greedy_search", provider="cpu", ) else: self._recognizer = sherpa_onnx.OnlineRecognizer.from_transducer( tokens=str(paths["tokens"]), encoder=str(paths["encoder"]), decoder=str(paths["decoder"]), joiner=str(paths["joiner"]), num_threads=1, decoding_method="greedy_search", enable_endpoint_detection=True, provider="cpu", ) except Exception as exc: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, f"failed to load sherpa-onnx STT model: {exc}", False, "sherpa-onnx-stt", "stt", ) from exc self.loaded = True def start_stream(self) -> "SherpaOnnxRealtimeTranscriptSession": if not self.loaded or self._recognizer is None: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, "sherpa-onnx STT provider is not loaded", False, "sherpa-onnx-stt", "stt", ) return SherpaOnnxRealtimeTranscriptSession(self._recognizer, self.language) def transcribe(self, segment: AudioSegment) -> Transcript: if not self.loaded or self._recognizer is None: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, "sherpa-onnx STT provider is not loaded", False, "sherpa-onnx-stt", "stt", ) try: import numpy as np samples = _segment_to_float32(segment, np) stream = self._recognizer.create_stream() stream.accept_waveform(segment.sample_rate, samples) stream.accept_waveform(segment.sample_rate, np.zeros(int(0.5 * segment.sample_rate), dtype=np.float32)) stream.input_finished() while self._recognizer.is_ready(stream): self._recognizer.decode_stream(stream) result = self._recognizer.get_result_all(stream) text = str(getattr(result, "text", "")).strip() raw_json = result.as_json_string() if hasattr(result, "as_json_string") else "" except Exception as exc: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, f"sherpa-onnx transcription failed: {exc}", True, "sherpa-onnx-stt", "stt", ) from exc if not is_valid_transcript_text(text): raise ProviderError( ErrorCode.STT_EMPTY_TRANSCRIPT, "STT produced no meaningful text", True, "sherpa-onnx-stt", "stt", ) return Transcript( text=text, language=self.language, confidence=None, duration_ms=segment.duration_ms, provider="sherpa-onnx-stt", raw_metadata={"raw_json": raw_json}, ) def _segment_to_float32(segment: AudioSegment, np: Any) -> Any: samples = np.frombuffer(segment.pcm, dtype=np.int16).astype(np.float32) / 32768.0 if segment.channels > 1 and samples.size: samples = samples.reshape(-1, segment.channels).mean(axis=1) return samples class SherpaOnnxRealtimeTranscriptSession: def __init__(self, recognizer: Any, language: str = "zh") -> None: self.recognizer = recognizer self.language = language self.stream = recognizer.create_stream() self._last_text = "" def accept_frame(self, frame: AudioFrame) -> Transcript | None: try: import numpy as np samples = _frame_to_float32(frame, np) if samples.size == 0: return None self.stream.accept_waveform(frame.sample_rate, samples) while self.recognizer.is_ready(self.stream): self.recognizer.decode_stream(self.stream) text = _recognizer_result_text(self.recognizer, self.stream) except Exception as exc: raise ProviderError( ErrorCode.STT_TRANSCRIBE_FAILED, f"sherpa-onnx realtime transcription failed: {exc}", True, "sherpa-onnx-stt", "stt", ) from exc if not should_emit_partial_transcript(text, self._last_text): return None self._last_text = text return Transcript( text=text, language=self.language, confidence=None, duration_ms=int(frame.metadata.get("duration_ms", 20)), provider="sherpa-onnx-stt-partial", ) def finish(self) -> Transcript | None: return None def _frame_to_float32(frame: AudioFrame, np: Any) -> Any: samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0 if frame.channels > 1 and samples.size: samples = samples.reshape(-1, frame.channels).mean(axis=1) return samples def _recognizer_result_text(recognizer: Any, stream: Any) -> str: if hasattr(recognizer, "get_result"): result = recognizer.get_result(stream) else: result = recognizer.get_result_all(stream) if isinstance(result, str): return result.strip() return str(getattr(result, "text", "")).strip() def _segment_to_wav_bytes(segment: AudioSegment) -> bytes: buffer = io.BytesIO() with wave.open(buffer, "wb") as wav: wav.setnchannels(segment.channels) wav.setsampwidth(2) wav.setframerate(segment.sample_rate) wav.writeframes(segment.pcm) return buffer.getvalue()