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