[Wake/VAD/STT 与 Live runtime]:完成真实重复语音运行,包含云端ASR/TTS开关、run-live和临时上下文测试

This commit is contained in:
mkbk
2026-06-17 20:00:55 +08:00
parent ac97daa1e7
commit 4b7cd18a0f
20 changed files with 1043 additions and 68 deletions
+204 -10
View File
@@ -1,9 +1,20 @@
from __future__ import annotations
import io
import json
import re
import socket
import urllib.error
import urllib.request
import uuid
import wave
from collections.abc import Callable
from pathlib import Path
from typing import Any
from .config import AppConfig
from .models import AudioSegment, ErrorCode, ProviderError, Transcript
from .speech_models import stt_model_paths
_MEANINGFUL_TEXT = re.compile(r"[\w\u4e00-\u9fff]", re.UNICODE)
@@ -48,11 +59,92 @@ class MetadataSttProvider:
)
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",
)
wav_bytes = _segment_to_wav_bytes(segment)
boundary = "owner-voice-pet-" + uuid.uuid4().hex
body = _multipart_form_data(
boundary,
fields={"model": self.config.asr_model, "response_format": "json"},
files={"file": ("utterance.wav", "audio/wav", wav_bytes)},
)
request = urllib.request.Request(
self.config.api_url("/v1/audio/transcriptions"),
data=body,
headers={
"Authorization": f"Bearer {self.config.llm_api_key}",
"Content-Type": f"multipart/form-data; boundary={boundary}",
},
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
text = str(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") -> None:
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():
@@ -63,12 +155,43 @@ class SherpaOnnxSttProvider:
"sherpa-onnx-stt",
"stt",
)
paths = stt_model_paths(self.model_path)
missing = [name for name, path in paths.items() if name != "model_dir" 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:
import sherpa_onnx # type: ignore[import-not-found] # noqa: F401
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"sherpa_onnx is not available: {exc}",
f"failed to load sherpa-onnx STT model: {exc}",
False,
"sherpa-onnx-stt",
"stt",
@@ -76,7 +199,7 @@ class SherpaOnnxSttProvider:
self.loaded = True
def transcribe(self, segment: AudioSegment) -> Transcript:
if not self.loaded:
if not self.loaded or self._recognizer is None:
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
"sherpa-onnx STT provider is not loaded",
@@ -84,10 +207,81 @@ class SherpaOnnxSttProvider:
"sherpa-onnx-stt",
"stt",
)
raise ProviderError(
ErrorCode.STT_TRANSCRIBE_FAILED,
"sherpa-onnx runtime transcription adapter requires a concrete model profile",
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
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()
def _multipart_form_data(
boundary: str,
*,
fields: dict[str, str],
files: dict[str, tuple[str, str, bytes]],
) -> bytes:
body = bytearray()
for name, value in fields.items():
body.extend(f"--{boundary}\r\n".encode())
body.extend(f'Content-Disposition: form-data; name="{name}"\r\n\r\n'.encode())
body.extend(value.encode("utf-8"))
body.extend(b"\r\n")
for name, (filename, content_type, data) in files.items():
body.extend(f"--{boundary}\r\n".encode())
body.extend(
f'Content-Disposition: form-data; name="{name}"; filename="{filename}"\r\n'.encode()
)
body.extend(f"Content-Type: {content_type}\r\n\r\n".encode())
body.extend(data)
body.extend(b"\r\n")
body.extend(f"--{boundary}--\r\n".encode())
return bytes(body)