[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
+11 -2
View File
@@ -18,11 +18,12 @@ from .models import (
from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport
from .wakeword import KeywordWakeWordProvider
from .vad import EnergyVadProvider, VadRecorder
from .stt import MetadataSttProvider, is_valid_transcript_text
from .stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .conversation import ConversationContext
from .llm import MockLlmProvider, OpenAICompatibleLlmProvider
from .pipeline import PipelineResult, VoicePipeline
from .tts import MacSayTtsProvider, SentenceBuffer, SineTtsProvider
from .runtime import LiveVoiceRuntime, RuntimeSummary, TerminalRuntimeReporter, TurnResult, build_live_runtime
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer, SineTtsProvider
from .assets import validate_pet_assets
from .ui import ConsolePetWindow, PetStateController, PetVisualState
@@ -37,13 +38,21 @@ __all__ = [
"KeywordWakeWordProvider",
"EnergyVadProvider",
"VadRecorder",
"CloudAsrSttProvider",
"MetadataSttProvider",
"SherpaOnnxSttProvider",
"is_valid_transcript_text",
"ConversationContext",
"MockLlmProvider",
"OpenAICompatibleLlmProvider",
"PipelineResult",
"VoicePipeline",
"LiveVoiceRuntime",
"RuntimeSummary",
"TerminalRuntimeReporter",
"TurnResult",
"build_live_runtime",
"CloudTtsProvider",
"MacSayTtsProvider",
"SentenceBuffer",
"SineTtsProvider",
+32 -3
View File
@@ -12,11 +12,12 @@ from .conversation import ConversationContext
from .llm import MockLlmProvider, OpenAICompatibleLlmProvider
from .models import AudioFrame, ProviderError
from .pipeline import VoicePipeline
from .runtime import build_live_runtime
from .speech_models import check_speech_models, model_status_errors
from .stt import MetadataSttProvider
from .stt import MetadataSttProvider, SherpaOnnxSttProvider
from .transport import MemoryAudioTransport, sounddevice_device_report
from .tts import SineTtsProvider
from .vad import EnergyVadProvider, VadRecorder
from .vad import EnergyVadProvider, SherpaOnnxVadProvider, VadRecorder
from .wakeword import KeywordWakeWordProvider
@@ -31,6 +32,8 @@ def main(argv: list[str] | None = None) -> int:
model_check = subparsers.add_parser("model-check", help="Validate local speech model files")
model_check.add_argument("--models-dir", default=None, help="Speech models directory. Defaults to .env or models")
subparsers.add_parser("device-check", help="Validate local microphone and speaker availability")
live = subparsers.add_parser("run-live", help="Run real repeated live voice conversation")
live.add_argument("--once", action="store_true", help="Run one completed live turn and exit")
smoke = subparsers.add_parser("llm-smoke", help="Call configured OpenAI/NewAPI endpoint")
smoke.add_argument("--message", default="用一句中文回复:小杰在线。")
smoke.add_argument("--no-stream", action="store_true")
@@ -50,6 +53,10 @@ def main(argv: list[str] | None = None) -> int:
"llm_stream": config.llm_stream,
"llm_api_key_present": bool(config.llm_api_key),
"asset_dir": str(config.asset_dir),
"speech_provider": config.speech_provider,
"asr_model": config.asr_model,
"tts_model": config.tts_model,
"tts_voice": config.tts_voice,
"speech_models_dir": str(config.speech_models_dir),
},
ensure_ascii=False,
@@ -73,8 +80,17 @@ def main(argv: list[str] | None = None) -> int:
models_dir = Path(args.models_dir) if args.models_dir else config.speech_models_dir
status = check_speech_models(models_dir, require_sherpa=True)
errors = model_status_errors(status)
provider_load_checked = False
if not errors:
try:
SherpaOnnxVadProvider(models_dir).load()
SherpaOnnxSttProvider(str(models_dir)).load()
provider_load_checked = True
except ProviderError as exc:
errors.append(exc)
data = status.to_json()
data["errors"] = [str(error) for error in errors]
data["provider_load_checked"] = provider_load_checked
print(json.dumps(data, ensure_ascii=False, sort_keys=True))
return 1 if errors else 0
@@ -83,6 +99,15 @@ def main(argv: list[str] | None = None) -> int:
print(json.dumps(report, ensure_ascii=False, sort_keys=True))
return 0 if report["ok"] else 1
if args.command == "run-live":
config = AppConfig.from_dotenv(args.env_file)
try:
summary = build_live_runtime(config).run(once=args.once)
except ProviderError as exc:
print(json.dumps({"ok": False, "code": exc.code.value, "message": exc.message}, ensure_ascii=False, sort_keys=True))
return 1
return 0 if summary.completed_turns > 0 or summary.interrupted else 1
if args.command == "acceptance":
result = run_acceptance()
print(json.dumps(result, ensure_ascii=False, sort_keys=True))
@@ -104,6 +129,10 @@ def main(argv: list[str] | None = None) -> int:
audio_output_device=config.audio_output_device,
asset_dir=config.asset_dir,
log_dir=config.log_dir,
speech_provider=config.speech_provider,
asr_model=config.asr_model,
tts_model=config.tts_model,
tts_voice=config.tts_voice,
speech_models_dir=config.speech_models_dir,
context_max_messages=config.context_max_messages,
context_max_chars=config.context_max_chars,
@@ -155,7 +184,7 @@ def run_acceptance() -> dict[str, object]:
def find_secret_leaks(root: Path) -> list[str]:
completed = subprocess.run(["git", "ls-files"], cwd=root, check=True, stdout=subprocess.PIPE, text=True)
pattern = re.compile(r"sk-[A-Za-z0-9_\\-]{16,}")
pattern = re.compile(r"(?:sk|tp)-[A-Za-z0-9_\\-]{16,}")
leaks: list[str] = []
for rel in completed.stdout.splitlines():
path = root / rel
+27 -2
View File
@@ -11,7 +11,7 @@ class AppConfig:
wake_word: str = "小杰小杰"
sample_rate: int = 16000
channels: int = 1
llm_base_url: str = "https://newapi.mkbk.shop"
llm_base_url: str = "https://token-plan-cn.xiaomimimo.com/v1"
llm_api_key: str | None = None
llm_model: str = "gpt-5.4-mini"
llm_api_style: str = "chat_completions"
@@ -20,6 +20,10 @@ class AppConfig:
audio_output_device: str | None = None
asset_dir: Path = Path("assets/pet")
log_dir: Path = Path("logs")
speech_provider: str = "cloud"
asr_model: str = "mimo-v2.5-asr"
tts_model: str = "mimo-v2.5-tts"
tts_voice: str = "alloy"
speech_models_dir: Path = Path("models")
context_max_messages: int = 12
context_max_chars: int = 12000
@@ -36,7 +40,7 @@ class AppConfig:
wake_word=get("WAKE_WORD", "小杰小杰") or "小杰小杰",
sample_rate=int(get("SAMPLE_RATE", "16000") or "16000"),
channels=int(get("CHANNELS", "1") or "1"),
llm_base_url=(get("LLM_BASE_URL", "https://newapi.mkbk.shop") or "").rstrip("/"),
llm_base_url=(get("LLM_BASE_URL", "https://token-plan-cn.xiaomimimo.com/v1") or "").rstrip("/"),
llm_api_key=get("LLM_API_KEY"),
llm_model=get("LLM_MODEL", "gpt-5.4-mini") or "gpt-5.4-mini",
llm_api_style=get("LLM_API_STYLE", "chat_completions") or "chat_completions",
@@ -45,6 +49,10 @@ class AppConfig:
audio_output_device=get("AUDIO_OUTPUT_DEVICE"),
asset_dir=Path(get("ASSET_DIR", "assets/pet") or "assets/pet"),
log_dir=Path(get("LOG_DIR", "logs") or "logs"),
speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(),
asr_model=get("ASR_MODEL", "mimo-v2.5-asr") or "mimo-v2.5-asr",
tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts",
tts_voice=get("TTS_VOICE", "alloy") or "alloy",
speech_models_dir=Path(get("SPEECH_MODELS_DIR", "models") or "models"),
context_max_messages=int(get("CONTEXT_MAX_MESSAGES", "12") or "12"),
context_max_chars=int(get("CONTEXT_MAX_CHARS", "12000") or "12000"),
@@ -96,6 +104,16 @@ class AppConfig:
"startup",
)
)
if self.speech_provider not in {"cloud", "local"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_SPEECH_PROVIDER must be cloud or local",
False,
"config",
"startup",
)
)
if not self.llm_base_url.startswith(("http://", "https://")):
errors.append(
ProviderError(
@@ -108,6 +126,13 @@ class AppConfig:
)
return errors
def api_url(self, path: str) -> str:
normalized = path if path.startswith("/") else f"/{path}"
base = self.llm_base_url.rstrip("/")
if base.endswith("/v1") and normalized.startswith("/v1/"):
return base + normalized[3:]
return base + normalized
def parse_dotenv(path: Path) -> dict[str, str]:
if not path.exists():
+2 -2
View File
@@ -51,7 +51,7 @@ class OpenAICompatibleLlmProvider:
"stream": self.config.llm_stream,
}
yield from self._post_stream(
f"{self.config.llm_base_url}/v1/chat/completions",
self.config.api_url("/v1/chat/completions"),
payload,
parser=_parse_chat_completion_sse,
)
@@ -63,7 +63,7 @@ class OpenAICompatibleLlmProvider:
"stream": self.config.llm_stream,
}
yield from self._post_stream(
f"{self.config.llm_base_url}/v1/responses",
self.config.api_url("/v1/responses"),
payload,
parser=_parse_responses_sse,
)
+267
View File
@@ -0,0 +1,267 @@
from __future__ import annotations
import sys
from dataclasses import dataclass, field
from typing import Protocol
from .config import AppConfig
from .conversation import ConversationContext
from .llm import OpenAICompatibleLlmProvider
from .models import AudioSegment, ErrorCode, PipelineState, ProviderError
from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
from .vad import EnergyVadProvider, VadRecorder
class RuntimeReporter(Protocol):
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
...
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
...
class TerminalRuntimeReporter:
def status(self, state: str, message: str, *, turn_id: int | None = None) -> None:
prefix = f"[第{turn_id}轮] " if turn_id is not None else ""
print(f"{prefix}{message}", flush=True)
def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None:
prefix = f"[第{turn_id}轮] " if turn_id is not None else ""
print(f"{prefix}{stage}失败:{code} {message}", file=sys.stderr, flush=True)
@dataclass(slots=True)
class TurnResult:
success: bool
transcript: str = ""
assistant_text: str = ""
error: ProviderError | None = None
states: list[PipelineState] = field(default_factory=list)
@dataclass(slots=True)
class RuntimeSummary:
completed_turns: int
failed_turns: int
interrupted: bool = False
last_error: ProviderError | None = None
class LiveVoiceRuntime:
def __init__(
self,
*,
config: AppConfig,
transport: AudioTransport,
vad_recorder: VadRecorder,
stt: SttProvider,
llm: LlmProvider,
tts: TtsProvider,
context: ConversationContext,
reporter: RuntimeReporter | None = None,
sentence_buffer: SentenceBuffer | None = None,
) -> None:
self.config = config
self.transport = transport
self.vad_recorder = vad_recorder
self.stt = stt
self.llm = llm
self.tts = tts
self.context = context
self.reporter = reporter or TerminalRuntimeReporter()
self.sentence_buffer = sentence_buffer or SentenceBuffer()
self._states: list[PipelineState] = []
def load(self) -> None:
self.vad_recorder.provider.load()
self.stt.load()
self.tts.load()
def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary:
completed = 0
failed = 0
last_error: ProviderError | None = None
self.load()
self.transport.start_input(
device_id=self.config.audio_input_device,
sample_rate=self.config.sample_rate,
channels=self.config.channels,
)
try:
while True:
turn_id = completed + failed + 1
result = self.run_turn(turn_id)
if result.success:
completed += 1
else:
failed += 1
last_error = result.error
if once:
break
if once and completed >= 1:
break
if max_turns is not None and completed >= max_turns:
break
except KeyboardInterrupt:
return RuntimeSummary(completed, failed, interrupted=True, last_error=last_error)
finally:
self.shutdown()
return RuntimeSummary(completed, failed, last_error=last_error)
def run_turn(self, turn_id: int) -> TurnResult:
self._states = []
try:
self._state(PipelineState.WAKE_LISTENING, "待机:等待唤醒词“小杰小杰”", turn_id=turn_id)
user_text = self._wait_for_wake_and_user_text(turn_id)
if isinstance(user_text, ProviderError):
return self._recover(user_text, turn_id)
return self._reply_to_user(user_text, turn_id)
except ProviderError as exc:
return self._recover(exc, turn_id)
def shutdown(self) -> None:
self.transport.stop()
def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError:
while True:
wake_segment = self._capture_segment(turn_id, state_message="待机:检测到语音,正在判断唤醒词")
if isinstance(wake_segment, ProviderError):
if wake_segment.code == ErrorCode.VAD_TIMEOUT_NO_SPEECH:
self._state(PipelineState.WAKE_LISTENING, "待机:继续等待唤醒词“小杰小杰”", turn_id=turn_id)
continue
return wake_segment
try:
wake_transcript = self.stt.transcribe(wake_segment).normalized_text
except ProviderError as exc:
if exc.code == ErrorCode.STT_EMPTY_TRANSCRIPT:
self._state(PipelineState.WAKE_LISTENING, "恢复待机:未听清唤醒词", turn_id=turn_id)
continue
raise
remainder = _text_after_wake_word(wake_transcript, self.config.wake_word)
if remainder is None:
self._state(PipelineState.WAKE_LISTENING, "恢复待机:未命中唤醒词", turn_id=turn_id)
continue
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id)
if is_valid_transcript_text(remainder):
return remainder
user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题")
if isinstance(user_segment, ProviderError):
return user_segment
self._state(PipelineState.TRANSCRIBING, "转写中:正在识别问题", turn_id=turn_id)
transcript = self.stt.transcribe(user_segment)
user_text = transcript.normalized_text
if not is_valid_transcript_text(user_text):
return ProviderError(
ErrorCode.STT_EMPTY_TRANSCRIPT,
"STT produced no meaningful user text",
True,
"live-runtime",
"stt",
)
return user_text
def _capture_segment(self, turn_id: int, *, state_message: str) -> AudioSegment | ProviderError:
self.vad_recorder.reset()
self.vad_recorder.provider.reset()
self._state(PipelineState.RECORDING, state_message, turn_id=turn_id)
while True:
frames = self.transport.read_frames(timeout_ms=100)
if not frames:
continue
for frame in frames:
result = self.vad_recorder.feed(frame)
if isinstance(result, ProviderError):
return result
if isinstance(result, AudioSegment):
return result
def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult:
self.context.append_user(user_text)
self._state(PipelineState.THINKING, f"思考中:{user_text}", turn_id=turn_id)
assistant_text = ""
try:
for delta in self.llm.stream_reply(self.context.build_llm_messages()):
assistant_text += delta.text_delta
for sentence in self.sentence_buffer.feed(delta.text_delta, bool(delta.finish_reason)):
self._speak(sentence, turn_id)
for sentence in self.sentence_buffer.flush():
self._speak(sentence, turn_id)
except ProviderError as exc:
return self._recover(exc, turn_id)
if not assistant_text.strip():
return self._recover(
ProviderError(
ErrorCode.LLM_EMPTY_REPLY,
"LLM returned no assistant text",
True,
"live-runtime",
"llm",
),
turn_id,
)
self.context.append_assistant(assistant_text)
self._state(PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id)
return TurnResult(True, user_text, assistant_text, states=list(self._states))
def _speak(self, sentence: str, turn_id: int) -> None:
self._state(PipelineState.SPEAKING, "播放中:正在播报回复", turn_id=turn_id)
segment = self.tts.synthesize(sentence)
playback = self.transport.play_pcm(segment)
if playback.error:
raise playback.error
def _recover(self, error: ProviderError, turn_id: int) -> TurnResult:
self.reporter.error(error.stage, error.code.value, error.message, turn_id=turn_id)
self._state(PipelineState.ERROR_RECOVERING, "恢复待机:本轮已结束", turn_id=turn_id)
self._state(PipelineState.WAKE_LISTENING, "恢复待机:可继续唤醒", turn_id=turn_id)
return TurnResult(False, error=error, states=list(self._states))
def _state(self, state: PipelineState, message: str, *, turn_id: int) -> None:
self._states.append(state)
self.reporter.status(state.value, message, turn_id=turn_id)
def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = None) -> LiveVoiceRuntime:
errors = config.validate_basic()
if errors:
raise errors[0]
if config.speech_provider == "cloud":
stt: SttProvider = CloudAsrSttProvider(config)
tts: TtsProvider = CloudTtsProvider(config)
else:
stt = SherpaOnnxSttProvider(str(config.speech_models_dir))
tts = MacSayTtsProvider()
return LiveVoiceRuntime(
config=config,
transport=SoundDeviceAudioTransport(output_device=config.audio_output_device),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=300, end_silence_ms=500, no_speech_timeout_ms=8000),
stt=stt,
llm=OpenAICompatibleLlmProvider(config),
tts=tts,
context=ConversationContext(
max_messages=config.context_max_messages,
max_chars=config.context_max_chars,
),
reporter=reporter,
)
def _text_after_wake_word(text: str, wake_word: str) -> str | None:
compact_text = _compact(text)
compact_wake = _compact(wake_word)
index = compact_text.find(compact_wake)
if index < 0:
return None
end = index + len(compact_wake)
compact_remainder = compact_text[end:].strip(",。.!!?? ")
if not compact_remainder:
return ""
original = text.replace(" ", "")
return original[-len(compact_remainder) :]
def _compact(text: str) -> str:
return "".join(ch for ch in text.strip() if not ch.isspace())
+28
View File
@@ -99,6 +99,34 @@ def required_model_files(models_dir: str | Path) -> tuple[str, ...]:
return tuple(str(item) for item in files)
def vad_model_path(models_dir: str | Path) -> Path:
root = Path(models_dir)
manifest = load_manifest(root)
path = manifest.get("providers", {}).get("vad", {}).get("path", "vad/silero_vad.onnx")
return root / str(path)
def stt_model_paths(model_path: str | Path) -> dict[str, Path]:
root = Path(model_path)
if (root / "manifest.json").exists() or (root / "stt").exists():
manifest = load_manifest(root)
stt = manifest.get("providers", {}).get("stt", {})
return {
"model_dir": root / str(stt.get("model_dir", f"stt/{DEFAULT_STT_DIR}")),
"tokens": root / str(stt.get("tokens", f"stt/{DEFAULT_STT_DIR}/tokens.txt")),
"encoder": root / str(stt.get("encoder", f"stt/{DEFAULT_STT_DIR}/encoder-epoch-99-avg-1.int8.onnx")),
"decoder": root / str(stt.get("decoder", f"stt/{DEFAULT_STT_DIR}/decoder-epoch-99-avg-1.onnx")),
"joiner": root / str(stt.get("joiner", f"stt/{DEFAULT_STT_DIR}/joiner-epoch-99-avg-1.int8.onnx")),
}
return {
"model_dir": root,
"tokens": root / "tokens.txt",
"encoder": root / "encoder-epoch-99-avg-1.int8.onnx",
"decoder": root / "decoder-epoch-99-avg-1.onnx",
"joiner": root / "joiner-epoch-99-avg-1.int8.onnx",
}
def check_speech_models(models_dir: str | Path, require_sherpa: bool = True) -> SpeechModelStatus:
root = Path(models_dir)
manifest_path = root / "manifest.json"
+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)
+1 -1
View File
@@ -247,7 +247,7 @@ class SoundDeviceAudioTransport:
"transport",
),
)
if segment.metadata.get("format") in {"aiff", "wav"}:
if segment.metadata.get("format") in {"aiff", "wav", "mp3", "m4a", "aac"}:
return _play_file_bytes_with_afplay(segment)
try:
with self._sd.RawOutputStream(
+86
View File
@@ -1,11 +1,18 @@
from __future__ import annotations
import math
import json
import socket
import struct
import subprocess
import tempfile
import urllib.error
import urllib.request
from collections.abc import Callable
from pathlib import Path
from typing import Any
from .config import AppConfig
from .models import AudioSegment, ErrorCode, ProviderError
@@ -121,3 +128,82 @@ class MacSayTtsProvider:
"tts",
)
return AudioSegment(data, 16000, 1, 0, max(120, len(clean) * 45), {"text": clean, "format": "aiff"})
class CloudTtsProvider:
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 synthesize(self, text: str) -> AudioSegment:
if not self.loaded:
raise ProviderError(
ErrorCode.TTS_SYNTHESIS_FAILED,
"cloud TTS provider is not loaded",
False,
"newapi-tts",
"tts",
)
clean = text.strip()
if not clean:
raise ProviderError(
ErrorCode.TTS_EMPTY_AUDIO,
"cannot synthesize empty text",
True,
"newapi-tts",
"tts",
)
payload = {
"model": self.config.tts_model,
"input": clean,
"voice": self.config.tts_voice,
"response_format": "mp3",
}
request = urllib.request.Request(
self.config.api_url("/v1/audio/speech"),
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:
data = response.read()
except urllib.error.HTTPError as exc:
raise ProviderError(
ErrorCode.TTS_SYNTHESIS_FAILED,
f"cloud TTS HTTP error {exc.code}",
exc.code >= 500,
"newapi-tts",
"tts",
) from exc
except (urllib.error.URLError, TimeoutError, socket.timeout) as exc:
raise ProviderError(
ErrorCode.TTS_SYNTHESIS_FAILED,
f"cloud TTS request failed: {exc}",
True,
"newapi-tts",
"tts",
) from exc
if not data:
raise ProviderError(
ErrorCode.TTS_EMPTY_AUDIO,
"cloud TTS returned empty audio",
True,
"newapi-tts",
"tts",
)
return AudioSegment(data, 16000, 1, 0, max(120, len(clean) * 45), {"text": clean, "format": "mp3"})
+109 -2
View File
@@ -1,12 +1,15 @@
from __future__ import annotations
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from .models import AudioFrame, AudioSegment, ErrorCode, ProviderError, VadResult
from .speech_models import vad_model_path
class EnergyVadProvider:
def __init__(self, threshold: int = 0) -> None:
def __init__(self, threshold: int = 500) -> None:
self.threshold = threshold
self.loaded = False
self._speech_ms = 0
@@ -47,12 +50,116 @@ class EnergyVadProvider:
return bool(frame.metadata["speech"])
if not frame.pcm:
return False
try:
import struct
sample_count = len(frame.pcm) // 2
if sample_count:
samples = struct.unpack("<" + "h" * sample_count, frame.pcm[: sample_count * 2])
return max(abs(sample) for sample in samples) > self.threshold
except Exception:
pass
return any(abs(byte - 128) > self.threshold for byte in frame.pcm)
class SherpaOnnxVadProvider:
def __init__(self, models_dir: str | Path, threshold: float = 0.5, sherpa_module: Any | None = None) -> None:
self.models_dir = Path(models_dir)
self.threshold = threshold
self.loaded = False
self._sherpa = sherpa_module
self._model: Any | None = None
self._speech_ms = 0
self._silence_ms = 0
def load(self) -> None:
model_path = vad_model_path(self.models_dir)
if not model_path.exists():
raise ProviderError(
ErrorCode.VAD_MODEL_LOAD_FAILED,
f"sherpa-onnx VAD model path does not exist: {model_path}",
False,
"sherpa-onnx-vad",
"vad",
)
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.VAD_MODEL_LOAD_FAILED,
f"sherpa_onnx is not available: {exc}",
False,
"sherpa-onnx-vad",
"vad",
) from exc
try:
config = sherpa_onnx.VadModelConfig(
silero_vad=sherpa_onnx.SileroVadModelConfig(model=str(model_path), threshold=self.threshold),
sample_rate=16000,
)
self._model = sherpa_onnx.VadModel.create(config)
except Exception as exc:
raise ProviderError(
ErrorCode.VAD_MODEL_LOAD_FAILED,
f"failed to load sherpa-onnx VAD model: {exc}",
False,
"sherpa-onnx-vad",
"vad",
) from exc
self.loaded = True
def analyze(self, frame: AudioFrame) -> VadResult:
if not self.loaded or self._model is None:
raise ProviderError(
ErrorCode.VAD_MODEL_LOAD_FAILED,
"sherpa-onnx VAD provider is not loaded",
False,
"sherpa-onnx-vad",
"vad",
)
try:
import numpy as np
samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0
window_size = int(self._model.window_size())
if samples.size < window_size:
samples = np.pad(samples, (0, window_size - samples.size))
elif samples.size > window_size:
samples = samples[-window_size:]
is_speech = bool(self._model.is_speech(samples))
except Exception as exc:
raise ProviderError(
ErrorCode.VAD_MODEL_LOAD_FAILED,
f"sherpa-onnx VAD analysis failed: {exc}",
True,
"sherpa-onnx-vad",
"vad",
) from exc
frame_ms = int(frame.metadata.get("duration_ms", 20))
if is_speech:
self._speech_ms += frame_ms
self._silence_ms = 0
else:
self._silence_ms += frame_ms
return VadResult(
is_speech=is_speech,
confidence=0.9 if is_speech else 0.1,
speech_ms=self._speech_ms,
silence_ms=self._silence_ms,
)
def reset(self) -> None:
self._speech_ms = 0
self._silence_ms = 0
if self._model is not None:
self._model.reset()
@dataclass(slots=True)
class VadRecorder:
provider: EnergyVadProvider
provider: Any
min_duration_ms: int = 300
end_silence_ms: int = 200
no_speech_timeout_ms: int = 1000