407 lines
14 KiB
Python
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()
|