[Wake/VAD/STT 与 Live runtime]:完成真实重复语音运行,包含云端ASR/TTS开关、run-live和临时上下文测试
This commit is contained in:
+204
-10
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user