[唤醒应答与快速端点]:完成唤醒后我在播报和本地VAD端点优化,包含缓冲清理、配置项和测试覆盖

This commit is contained in:
mkbk
2026-06-17 20:51:14 +08:00
parent 860c2909fc
commit 44f708a3dc
12 changed files with 224 additions and 7 deletions
+8
View File
@@ -10,6 +10,14 @@ OWNER_WAKE_PROVIDER=local_kws
OWNER_WAKE_KEYWORDS_FILE= OWNER_WAKE_KEYWORDS_FILE=
OWNER_WAKE_KWS_THRESHOLD=0.25 OWNER_WAKE_KWS_THRESHOLD=0.25
OWNER_WAKE_KWS_SCORE=1.0 OWNER_WAKE_KWS_SCORE=1.0
OWNER_WAKE_ACK_TEXT=我在
OWNER_POST_PLAYBACK_DRAIN_MS=250
OWNER_VAD_PROVIDER=local
OWNER_VAD_THRESHOLD=0.5
OWNER_VAD_MIN_DURATION_MS=250
OWNER_VAD_END_SILENCE_MS=350
OWNER_VAD_NO_SPEECH_TIMEOUT_MS=5000
OWNER_VAD_MAX_RECORDING_MS=12000
OWNER_SPEECH_PROVIDER=cloud OWNER_SPEECH_PROVIDER=cloud
OWNER_ASR_MODEL=mimo-v2.5-asr OWNER_ASR_MODEL=mimo-v2.5-asr
OWNER_TTS_MODEL=mimo-v2.5-tts OWNER_TTS_MODEL=mimo-v2.5-tts
+12 -2
View File
@@ -13,7 +13,7 @@
- `OWNER_SPEECH_PROVIDER=cloud|local`:选择云端语音模型或本地语音模型。 - `OWNER_SPEECH_PROVIDER=cloud|local`:选择云端语音模型或本地语音模型。
- 默认云端语音模型:`mimo-v2.5-asr``mimo-v2.5-tts` - 默认云端语音模型:`mimo-v2.5-asr``mimo-v2.5-tts`
- 本地设备:`sounddevice` 读取麦克风,扬声器或 `afplay` 播放。 - 本地设备:`sounddevice` 读取麦克风,扬声器或 `afplay` 播放。
- 本地模型:`models/` 存放 `sherpa-onnx` VAD/STT 模型,目录不提交 Git。 - 本地模型:`models/` 存放 `sherpa-onnx` wake/VAD/STT 模型,目录不提交 Git。
- 临时上下文:同一次 `run-live` 进程内携带最近 user/assistant 历史,退出即清空。 - 临时上下文:同一次 `run-live` 进程内携带最近 user/assistant 历史,退出即清空。
## 首次准备 ## 首次准备
@@ -36,13 +36,21 @@ OWNER_WAKE_PROVIDER=local_kws
OWNER_WAKE_KEYWORDS_FILE= OWNER_WAKE_KEYWORDS_FILE=
OWNER_WAKE_KWS_THRESHOLD=0.25 OWNER_WAKE_KWS_THRESHOLD=0.25
OWNER_WAKE_KWS_SCORE=1.0 OWNER_WAKE_KWS_SCORE=1.0
OWNER_WAKE_ACK_TEXT=我在
OWNER_POST_PLAYBACK_DRAIN_MS=250
OWNER_VAD_PROVIDER=local
OWNER_VAD_THRESHOLD=0.5
OWNER_VAD_MIN_DURATION_MS=250
OWNER_VAD_END_SILENCE_MS=350
OWNER_VAD_NO_SPEECH_TIMEOUT_MS=5000
OWNER_VAD_MAX_RECORDING_MS=12000
OWNER_SPEECH_PROVIDER=cloud OWNER_SPEECH_PROVIDER=cloud
OWNER_ASR_MODEL=mimo-v2.5-asr OWNER_ASR_MODEL=mimo-v2.5-asr
OWNER_TTS_MODEL=mimo-v2.5-tts OWNER_TTS_MODEL=mimo-v2.5-tts
OWNER_TTS_VOICE=mimo_default OWNER_TTS_VOICE=mimo_default
``` ```
`OWNER_WAKE_PROVIDER=local_kws` 表示唤醒词“小杰小杰”由本地模型检测。`OWNER_SPEECH_PROVIDER=cloud` 只影响唤醒后的正式问题 ASR/TTS:它会把 VAD 切出来的用户问题片段发送到云端 ASR,不会上传连续麦克风流,也不会用云端判断唤醒。改成 `local` 时使用项目 `models/` 下的本地语音模型路径。 `OWNER_WAKE_PROVIDER=local_kws` 表示唤醒词“小杰小杰”由本地模型检测。唤醒命中后会先本地播报 `OWNER_WAKE_ACK_TEXT=我在`,再开始听取问题。`OWNER_SPEECH_PROVIDER=cloud` 只影响唤醒后的正式问题 ASR/TTS:它会把 VAD 切出来的用户问题片段发送到云端 ASR,不会上传连续麦克风流,也不会用云端判断唤醒。改成 `local` 时使用项目 `models/` 下的本地语音模型路径。
## 本地模型 ## 本地模型
@@ -57,6 +65,8 @@ python3.11 scripts/download_speech_models.py --dir models
本地唤醒关键词文件位于 `models/wake/keywords.txt`。如果真人唤醒不灵敏,可以先把 `.env``OWNER_WAKE_KWS_THRESHOLD` 调低,例如 `0.15`,再重新运行。 本地唤醒关键词文件位于 `models/wake/keywords.txt`。如果真人唤醒不灵敏,可以先把 `.env``OWNER_WAKE_KWS_THRESHOLD` 调低,例如 `0.15`,再重新运行。
如果唤醒后你已经停说但还长时间显示“录音中”,优先调小 `OWNER_VAD_END_SILENCE_MS`,例如 `250`;如果房间噪声较大,再略调高 `OWNER_VAD_THRESHOLD`,例如 `0.6`
## 设备检查 ## 设备检查
```bash ```bash
@@ -11,6 +11,28 @@ The live runtime SHALL display the recognized user utterance text in the termina
- **WHEN** STT returns empty text, punctuation-only text, or an invalid transcript - **WHEN** STT returns empty text, punctuation-only text, or an invalid transcript
- **THEN** the runtime SHALL NOT emit a misleading transcript as valid user input and SHALL recover to standby without invoking the LLM - **THEN** the runtime SHALL NOT emit a misleading transcript as valid user input and SHALL recover to standby without invoking the LLM
### Requirement: Wake acknowledgement before recording
The live runtime SHALL provide an audible local acknowledgement after local wake detection and before it starts recording the user's formal question.
#### Scenario: Wake is detected
- **WHEN** the local wakeword provider detects “小杰小杰”
- **THEN** the runtime SHALL play a short acknowledgement such as “我在” before entering the user utterance recording state
#### Scenario: Acknowledgement playback finishes
- **WHEN** the acknowledgement playback completes
- **THEN** the runtime SHALL clear buffered microphone input captured during acknowledgement playback before starting VAD recording for the user's question
### Requirement: Fast user utterance endpointing
The live runtime SHALL use the project-local VAD model by default for user utterance endpoint detection and SHALL expose configurable silence timing so the recording stops promptly after the user stops speaking.
#### Scenario: User stops speaking after wake
- **WHEN** VAD observes the configured continuous silence duration after a started utterance
- **THEN** the runtime SHALL close the utterance segment and proceed to STT without waiting for the maximum recording duration
#### Scenario: Local VAD model is available
- **WHEN** `OWNER_VAD_PROVIDER=local`
- **THEN** the runtime SHALL use the project-local `sherpa-onnx` VAD model rather than a raw energy threshold for live user utterance endpointing
## MODIFIED Requirements ## MODIFIED Requirements
### Requirement: Wake word detection ### Requirement: Wake word detection
@@ -56,7 +78,7 @@ The system SHALL provide a `run-live` command that performs real repeated voice
#### Scenario: Live runtime completes two turns #### Scenario: Live runtime completes two turns
- **WHEN** the user wakes the system with “小杰小杰”, asks a question, hears the reply, then wakes it again and asks another question - **WHEN** the user wakes the system with “小杰小杰”, asks a question, hears the reply, then wakes it again and asks another question
- **THEN** the system SHALL complete local wake, recording, STT, transcript display, LLM, TTS, playback for both turns and SHALL return to standby after each turn - **THEN** the system SHALL complete local wake, audible acknowledgement, recording, STT, transcript display, LLM, TTS, playback for both turns and SHALL return to standby after each turn
#### Scenario: Once mode completes one turn #### Scenario: Once mode completes one turn
- **WHEN** the user runs `PYTHONPATH=src python3.11 -m owner_voice_pet run-live --once` - **WHEN** the user runs `PYTHONPATH=src python3.11 -m owner_voice_pet run-live --once`
@@ -31,3 +31,10 @@
- [ ] 4.3 执行真实 run-live 验收;前置条件:模型、设备、.env 齐全;验收标准:唤醒命中使用本地模型,终端显示转写结果;测试要点:单轮或两轮状态输出;优先级:P0;预计:60 分钟。 - [ ] 4.3 执行真实 run-live 验收;前置条件:模型、设备、.env 齐全;验收标准:唤醒命中使用本地模型,终端显示转写结果;测试要点:单轮或两轮状态输出;优先级:P0;预计:60 分钟。
- [x] 4.4 最终门禁;前置条件:全部实现完成;验收标准:compileall、unittest、security-check、model-check、device-check、OpenSpec strict 全通过;优先级:P0;预计:30 分钟。 - [x] 4.4 最终门禁;前置条件:全部实现完成;验收标准:compileall、unittest、security-check、model-check、device-check、OpenSpec strict 全通过;优先级:P0;预计:30 分钟。
- [ ] 4.5 归档变更并提交;前置条件:4.4 通过;验收标准:主 spec 更新,archive 完成,最终 commit`git status --short` 为空;测试要点:中文提交信息;优先级:P0;预计:20 分钟。 - [ ] 4.5 归档变更并提交;前置条件:4.4 通过;验收标准:主 spec 更新,archive 完成,最终 commit`git status --short` 为空;测试要点:中文提交信息;优先级:P0;预计:20 分钟。
## 5. 唤醒应答与快速端点修正
- [x] 5.1 增加唤醒后本地语音应答;前置条件:本地 KWS 已可唤醒;验收标准:wake 命中后播放“我在”再进入录音;测试要点:fake runtime 播放顺序;优先级:P0;预计:45 分钟。
- [x] 5.2 清理应答播放期间的麦克风缓冲;前置条件:5.1 完成;验收标准:应答音频不进入正式问题 VAD/STT;测试要点:transport flush 测试;优先级:P0;预计:45 分钟。
- [x] 5.3 live 默认改用本地 `sherpa-onnx` VAD 并缩短静音端点;前置条件:模型已下载;验收标准:`.env` 可配置 VAD provider、静音结束时间和最大录音时长;测试要点:配置和 runtime 构造测试;优先级:P0;预计:45 分钟。
- [x] 5.4 验证并提交“唤醒应答与快速端点”模块;前置条件:5.1 至 5.3 完成;验收标准:compileall、unittest、security-check、model-check、OpenSpec strict 通过后 commit;优先级:P0;预计:20 分钟。
+16
View File
@@ -57,6 +57,14 @@ def main(argv: list[str] | None = None) -> int:
"wake_keywords_file": str(config.wake_keywords_file) if config.wake_keywords_file else "", "wake_keywords_file": str(config.wake_keywords_file) if config.wake_keywords_file else "",
"wake_kws_threshold": config.wake_kws_threshold, "wake_kws_threshold": config.wake_kws_threshold,
"wake_kws_score": config.wake_kws_score, "wake_kws_score": config.wake_kws_score,
"wake_ack_text": config.wake_ack_text,
"post_playback_drain_ms": config.post_playback_drain_ms,
"vad_provider": config.vad_provider,
"vad_threshold": config.vad_threshold,
"vad_min_duration_ms": config.vad_min_duration_ms,
"vad_end_silence_ms": config.vad_end_silence_ms,
"vad_no_speech_timeout_ms": config.vad_no_speech_timeout_ms,
"vad_max_recording_ms": config.vad_max_recording_ms,
"speech_provider": config.speech_provider, "speech_provider": config.speech_provider,
"asr_model": config.asr_model, "asr_model": config.asr_model,
"tts_model": config.tts_model, "tts_model": config.tts_model,
@@ -144,6 +152,14 @@ def main(argv: list[str] | None = None) -> int:
wake_keywords_file=config.wake_keywords_file, wake_keywords_file=config.wake_keywords_file,
wake_kws_threshold=config.wake_kws_threshold, wake_kws_threshold=config.wake_kws_threshold,
wake_kws_score=config.wake_kws_score, wake_kws_score=config.wake_kws_score,
wake_ack_text=config.wake_ack_text,
post_playback_drain_ms=config.post_playback_drain_ms,
vad_provider=config.vad_provider,
vad_threshold=config.vad_threshold,
vad_min_duration_ms=config.vad_min_duration_ms,
vad_end_silence_ms=config.vad_end_silence_ms,
vad_no_speech_timeout_ms=config.vad_no_speech_timeout_ms,
vad_max_recording_ms=config.vad_max_recording_ms,
speech_provider=config.speech_provider, speech_provider=config.speech_provider,
asr_model=config.asr_model, asr_model=config.asr_model,
tts_model=config.tts_model, tts_model=config.tts_model,
+62
View File
@@ -24,6 +24,14 @@ class AppConfig:
wake_keywords_file: Path | None = None wake_keywords_file: Path | None = None
wake_kws_threshold: float = 0.25 wake_kws_threshold: float = 0.25
wake_kws_score: float = 1.0 wake_kws_score: float = 1.0
wake_ack_text: str = "我在"
post_playback_drain_ms: int = 250
vad_provider: str = "local"
vad_threshold: float = 0.5
vad_min_duration_ms: int = 250
vad_end_silence_ms: int = 350
vad_no_speech_timeout_ms: int = 5000
vad_max_recording_ms: int = 12000
speech_provider: str = "cloud" speech_provider: str = "cloud"
asr_model: str = "mimo-v2.5-asr" asr_model: str = "mimo-v2.5-asr"
tts_model: str = "mimo-v2.5-tts" tts_model: str = "mimo-v2.5-tts"
@@ -57,6 +65,14 @@ class AppConfig:
wake_keywords_file=Path(value) if (value := get("WAKE_KEYWORDS_FILE")) else None, wake_keywords_file=Path(value) if (value := get("WAKE_KEYWORDS_FILE")) else None,
wake_kws_threshold=float(get("WAKE_KWS_THRESHOLD", "0.25") or "0.25"), wake_kws_threshold=float(get("WAKE_KWS_THRESHOLD", "0.25") or "0.25"),
wake_kws_score=float(get("WAKE_KWS_SCORE", "1.0") or "1.0"), wake_kws_score=float(get("WAKE_KWS_SCORE", "1.0") or "1.0"),
wake_ack_text=get("WAKE_ACK_TEXT", "我在") or "我在",
post_playback_drain_ms=int(get("POST_PLAYBACK_DRAIN_MS", "250") or "250"),
vad_provider=(get("VAD_PROVIDER", "local") or "local").lower(),
vad_threshold=float(get("VAD_THRESHOLD", "0.5") or "0.5"),
vad_min_duration_ms=int(get("VAD_MIN_DURATION_MS", "250") or "250"),
vad_end_silence_ms=int(get("VAD_END_SILENCE_MS", "350") or "350"),
vad_no_speech_timeout_ms=int(get("VAD_NO_SPEECH_TIMEOUT_MS", "5000") or "5000"),
vad_max_recording_ms=int(get("VAD_MAX_RECORDING_MS", "12000") or "12000"),
speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(), speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(),
asr_model=get("ASR_MODEL", "mimo-v2.5-asr") or "mimo-v2.5-asr", 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_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts",
@@ -152,6 +168,52 @@ class AppConfig:
"startup", "startup",
) )
) )
if self.post_playback_drain_ms < 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_POST_PLAYBACK_DRAIN_MS must be non-negative",
False,
"config",
"startup",
)
)
if self.vad_provider not in {"local", "energy"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_VAD_PROVIDER must be local or energy",
False,
"config",
"startup",
)
)
if self.vad_threshold <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_VAD_THRESHOLD must be positive",
False,
"config",
"startup",
)
)
for name, value in {
"OWNER_VAD_MIN_DURATION_MS": self.vad_min_duration_ms,
"OWNER_VAD_END_SILENCE_MS": self.vad_end_silence_ms,
"OWNER_VAD_NO_SPEECH_TIMEOUT_MS": self.vad_no_speech_timeout_ms,
"OWNER_VAD_MAX_RECORDING_MS": self.vad_max_recording_ms,
}.items():
if value <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
f"{name} must be positive",
False,
"config",
"startup",
)
)
if not self.llm_base_url.startswith(("http://", "https://")): if not self.llm_base_url.startswith(("http://", "https://")):
errors.append( errors.append(
ProviderError( ProviderError(
+3
View File
@@ -28,6 +28,9 @@ class AudioTransport(Protocol):
def play_pcm(self, segment: AudioSegment, interrupt: bool = False) -> PlaybackResult: def play_pcm(self, segment: AudioSegment, interrupt: bool = False) -> PlaybackResult:
... ...
def flush_input(self) -> int:
...
def stop(self) -> None: def stop(self) -> None:
... ...
+48 -2
View File
@@ -12,7 +12,7 @@ from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, Wa
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .transport import SoundDeviceAudioTransport from .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
from .vad import EnergyVadProvider, VadRecorder from .vad import EnergyVadProvider, SherpaOnnxVadProvider, VadRecorder
from .wakeword import SherpaOnnxKeywordWakeWordProvider from .wakeword import SherpaOnnxKeywordWakeWordProvider
@@ -71,6 +71,7 @@ class LiveVoiceRuntime:
llm: LlmProvider, llm: LlmProvider,
tts: TtsProvider, tts: TtsProvider,
context: ConversationContext, context: ConversationContext,
ack_tts: TtsProvider | None = None,
reporter: RuntimeReporter | None = None, reporter: RuntimeReporter | None = None,
sentence_buffer: SentenceBuffer | None = None, sentence_buffer: SentenceBuffer | None = None,
) -> None: ) -> None:
@@ -81,6 +82,7 @@ class LiveVoiceRuntime:
self.stt = stt self.stt = stt
self.llm = llm self.llm = llm
self.tts = tts self.tts = tts
self.ack_tts = ack_tts or tts
self.context = context self.context = context
self.reporter = reporter or TerminalRuntimeReporter() self.reporter = reporter or TerminalRuntimeReporter()
self.sentence_buffer = sentence_buffer or SentenceBuffer() self.sentence_buffer = sentence_buffer or SentenceBuffer()
@@ -91,6 +93,8 @@ class LiveVoiceRuntime:
self.vad_recorder.provider.load() self.vad_recorder.provider.load()
self.stt.load() self.stt.load()
self.tts.load() self.tts.load()
if self.ack_tts is not self.tts:
self.ack_tts.load()
def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary: def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary:
completed = 0 completed = 0
@@ -142,6 +146,9 @@ class LiveVoiceRuntime:
if wake_error is not None: if wake_error is not None:
return wake_error return wake_error
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id) self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id)
ack_error = self._acknowledge_wake(turn_id)
if ack_error is not None:
return ack_error
user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题") user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题")
if isinstance(user_segment, ProviderError): if isinstance(user_segment, ProviderError):
return user_segment return user_segment
@@ -186,6 +193,33 @@ class LiveVoiceRuntime:
if isinstance(result, AudioSegment): if isinstance(result, AudioSegment):
return result return result
def _acknowledge_wake(self, turn_id: int) -> ProviderError | None:
text = self.config.wake_ack_text.strip()
if not text:
self._drain_input_after_playback()
return None
try:
self._state(PipelineState.SPEAKING, f"应答中:{text}", turn_id=turn_id)
segment = self.ack_tts.synthesize(text)
playback = self.transport.play_pcm(segment)
if playback.error:
return playback.error
self._drain_input_after_playback()
return None
except ProviderError as exc:
return exc
def _drain_input_after_playback(self) -> None:
if self.config.post_playback_drain_ms <= 0:
return
self.transport.flush_input()
remaining_ms = self.config.post_playback_drain_ms
while remaining_ms > 0:
timeout_ms = min(50, remaining_ms)
self.transport.read_frames(timeout_ms=timeout_ms)
remaining_ms -= timeout_ms
self.transport.flush_input()
def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult: def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult:
self.context.append_user(user_text) self.context.append_user(user_text)
self._state(PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id) self._state(PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id)
@@ -220,6 +254,7 @@ class LiveVoiceRuntime:
playback = self.transport.play_pcm(segment) playback = self.transport.play_pcm(segment)
if playback.error: if playback.error:
raise playback.error raise playback.error
self._drain_input_after_playback()
def _recover(self, error: ProviderError, turn_id: int) -> TurnResult: def _recover(self, error: ProviderError, turn_id: int) -> TurnResult:
self.reporter.error(error.stage, error.code.value, error.message, turn_id=turn_id) self.reporter.error(error.stage, error.code.value, error.message, turn_id=turn_id)
@@ -242,6 +277,10 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
else: else:
stt = SherpaOnnxSttProvider(str(config.speech_models_dir)) stt = SherpaOnnxSttProvider(str(config.speech_models_dir))
tts = MacSayTtsProvider() tts = MacSayTtsProvider()
if config.vad_provider == "local":
vad_provider = SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold)
else:
vad_provider = EnergyVadProvider()
return LiveVoiceRuntime( return LiveVoiceRuntime(
config=config, config=config,
transport=SoundDeviceAudioTransport(output_device=config.audio_output_device), transport=SoundDeviceAudioTransport(output_device=config.audio_output_device),
@@ -252,10 +291,17 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
threshold=config.wake_kws_threshold, threshold=config.wake_kws_threshold,
score=config.wake_kws_score, score=config.wake_kws_score,
), ),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=300, end_silence_ms=500, no_speech_timeout_ms=8000), vad_recorder=VadRecorder(
vad_provider,
min_duration_ms=config.vad_min_duration_ms,
end_silence_ms=config.vad_end_silence_ms,
no_speech_timeout_ms=config.vad_no_speech_timeout_ms,
max_recording_ms=config.vad_max_recording_ms,
),
stt=stt, stt=stt,
llm=OpenAICompatibleLlmProvider(config), llm=OpenAICompatibleLlmProvider(config),
tts=tts, tts=tts,
ack_tts=MacSayTtsProvider(),
context=ConversationContext( context=ConversationContext(
max_messages=config.context_max_messages, max_messages=config.context_max_messages,
max_chars=config.context_max_chars, max_chars=config.context_max_chars,
+15
View File
@@ -103,6 +103,11 @@ class MemoryAudioTransport:
self.played_segments.append(segment) self.played_segments.append(segment)
return PlaybackResult(True, segment.duration_ms) return PlaybackResult(True, segment.duration_ms)
def flush_input(self) -> int:
count = len(self._frames)
self._frames.clear()
return count
def stop(self) -> None: def stop(self) -> None:
self.started = False self.started = False
@@ -290,6 +295,16 @@ class SoundDeviceAudioTransport:
break break
return None return None
def flush_input(self) -> int:
count = 0
while not self._queue.empty():
try:
self._queue.get_nowait()
count += 1
except queue.Empty:
break
return count
def health(self) -> TransportHealth: def health(self) -> TransportHealth:
if self._sd is None: if self._sd is None:
return TransportHealth(False, False, "sounddevice unavailable") return TransportHealth(False, False, "sounddevice unavailable")
+3 -2
View File
@@ -80,7 +80,7 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
tts = SineTtsProvider() tts = SineTtsProvider()
reporter = RecordingReporter() reporter = RecordingReporter()
runtime = LiveVoiceRuntime( runtime = LiveVoiceRuntime(
config=AppConfig(llm_api_key="secret", speech_provider="cloud"), config=AppConfig(llm_api_key="secret", speech_provider="cloud", post_playback_drain_ms=0),
transport=transport, transport=transport,
wakeword=KeywordWakeWordProvider(), wakeword=KeywordWakeWordProvider(),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
@@ -101,8 +101,9 @@ class LiveRuntimeTests(unittest.TestCase):
self.assertEqual(summary.completed_turns, 2) self.assertEqual(summary.completed_turns, 2)
self.assertEqual(len(stt.calls), 2) self.assertEqual(len(stt.calls), 2)
self.assertEqual(len(llm.calls), 2) self.assertEqual(len(llm.calls), 2)
self.assertEqual(len(transport.played_segments), 2) self.assertEqual(len(transport.played_segments), 4)
self.assertEqual(reporter.transcripts, ["第一问", "第二问"]) self.assertEqual(reporter.transcripts, ["第一问", "第二问"])
self.assertIn("应答中:我在", reporter.statuses)
self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1]) self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1])
def test_temporary_context_is_sent_to_second_llm_call(self) -> None: def test_temporary_context_is_sent_to_second_llm_call(self) -> None:
+13
View File
@@ -72,6 +72,14 @@ class ModelsConfigTests(unittest.TestCase):
self.assertEqual(config.wake_provider, "local_kws") self.assertEqual(config.wake_provider, "local_kws")
self.assertEqual(config.wake_kws_threshold, 0.25) self.assertEqual(config.wake_kws_threshold, 0.25)
self.assertEqual(config.wake_kws_score, 1.0) self.assertEqual(config.wake_kws_score, 1.0)
self.assertEqual(config.wake_ack_text, "我在")
self.assertEqual(config.post_playback_drain_ms, 250)
self.assertEqual(config.vad_provider, "local")
self.assertEqual(config.vad_threshold, 0.5)
self.assertEqual(config.vad_min_duration_ms, 250)
self.assertEqual(config.vad_end_silence_ms, 350)
self.assertEqual(config.vad_no_speech_timeout_ms, 5000)
self.assertEqual(config.vad_max_recording_ms, 12000)
self.assertEqual(config.speech_provider, "cloud") self.assertEqual(config.speech_provider, "cloud")
self.assertEqual(config.asr_model, "mimo-v2.5-asr") self.assertEqual(config.asr_model, "mimo-v2.5-asr")
self.assertEqual(config.tts_model, "mimo-v2.5-tts") self.assertEqual(config.tts_model, "mimo-v2.5-tts")
@@ -90,6 +98,11 @@ class ModelsConfigTests(unittest.TestCase):
errors = config.validate_basic() errors = config.validate_basic()
self.assertTrue(any("OWNER_WAKE_PROVIDER" in error.message for error in errors)) self.assertTrue(any("OWNER_WAKE_PROVIDER" in error.message for error in errors))
def test_vad_provider_must_be_local_or_energy(self) -> None:
config = AppConfig(vad_provider="invalid")
errors = config.validate_basic()
self.assertTrue(any("OWNER_VAD_PROVIDER" in error.message for error in errors))
def test_missing_dotenv_uses_non_secret_defaults(self) -> None: def test_missing_dotenv_uses_non_secret_defaults(self) -> None:
config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env") config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env")
self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1") self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
+14
View File
@@ -42,6 +42,12 @@ class TransportTests(unittest.TestCase):
self.assertFalse(result.played) self.assertFalse(result.played)
self.assertEqual(result.error.code, ErrorCode.AUDIO_OUTPUT_DEVICE_MISSING) self.assertEqual(result.error.code, ErrorCode.AUDIO_OUTPUT_DEVICE_MISSING)
def test_memory_transport_flushes_input_frames(self) -> None:
transport = MemoryAudioTransport([frame(1, 0), frame(2, 20)])
transport.start_input()
self.assertEqual(transport.flush_input(), 2)
self.assertEqual(transport.read_frames(10), [])
def test_file_replay_roundtrip(self) -> None: def test_file_replay_roundtrip(self) -> None:
frames = [frame(1, 0, {"wake": True}), frame(2, 20, {"speech": True})] frames = [frame(1, 0, {"wake": True}), frame(2, 20, {"speech": True})]
with tempfile.TemporaryDirectory() as tmp: with tempfile.TemporaryDirectory() as tmp:
@@ -76,6 +82,14 @@ class TransportTests(unittest.TestCase):
self.assertEqual(frames[0].pcm, b"\x01\x00\x02\x00") self.assertEqual(frames[0].pcm, b"\x01\x00\x02\x00")
self.assertEqual(frames[0].sample_rate, 16000) self.assertEqual(frames[0].sample_rate, 16000)
def test_sounddevice_transport_flushes_queued_input(self) -> None:
fake = FakeSoundDevice()
transport = SoundDeviceAudioTransport(sounddevice_module=fake)
transport.start_input(sample_rate=16000, channels=1)
self.assertEqual(transport.flush_input(), 1)
self.assertEqual(transport.read_frames(0), [])
transport.stop()
def test_sounddevice_transport_plays_pcm_to_raw_output_stream(self) -> None: def test_sounddevice_transport_plays_pcm_to_raw_output_stream(self) -> None:
fake = FakeSoundDevice() fake = FakeSoundDevice()
transport = SoundDeviceAudioTransport(sounddevice_module=fake) transport = SoundDeviceAudioTransport(sounddevice_module=fake)