Files
Owner/src/owner_voice_pet/stt.py
T

407 lines
14 KiB
Python

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()