diff --git a/.env.example b/.env.example index 854d9f8..7e5757c 100644 --- a/.env.example +++ b/.env.example @@ -1,4 +1,4 @@ -OWNER_LLM_BASE_URL=https://newapi.mkbk.shop +OWNER_LLM_BASE_URL=https://token-plan-cn.xiaomimimo.com/v1 OWNER_LLM_API_KEY= OWNER_LLM_MODEL=gpt-5.4-mini OWNER_LLM_API_STYLE=chat_completions @@ -6,6 +6,10 @@ OWNER_AUDIO_INPUT_DEVICE= OWNER_AUDIO_OUTPUT_DEVICE= OWNER_ASSET_DIR=assets/pet OWNER_LOG_DIR=logs +OWNER_SPEECH_PROVIDER=cloud +OWNER_ASR_MODEL=mimo-v2.5-asr +OWNER_TTS_MODEL=mimo-v2.5-tts +OWNER_TTS_VOICE=alloy OWNER_SPEECH_MODELS_DIR=models OWNER_CONTEXT_MAX_MESSAGES=12 OWNER_CONTEXT_MAX_CHARS=12000 diff --git a/openspec/changes/complete-live-repeat-voice-runtime/design.md b/openspec/changes/complete-live-repeat-voice-runtime/design.md index 71df36d..2edf0bd 100644 --- a/openspec/changes/complete-live-repeat-voice-runtime/design.md +++ b/openspec/changes/complete-live-repeat-voice-runtime/design.md @@ -4,7 +4,7 @@ `add-voice-pet-pipeline` 已归档并生成 `voice-pet-pipeline` 主规范。当前仓库已有 Python 包、CLI acceptance、`.env` 配置读取、Provider 协议、mock 测试、桌宠资产和 OpenSpec 主规范,但真实运行能力仍不完整:没有 `run-live` 常驻命令,没有真实麦克风输入和扬声器输出,没有本地模型下载/检查,没有证明同一运行进程内可以连续多轮携带历史。 -本设计聚焦第一版“无 GUI 的真实实时语音版”。GUI 桌宠窗口、长期记忆、播放中打断、跨设备 Transport 和云端 STT/TTS 均不在本变更范围内。 +本设计聚焦第一版“无 GUI 的真实实时语音版”。GUI 桌宠窗口、长期记忆、播放中打断和跨设备 Transport 不在本变更范围内;ASR/TTS 由 `.env` 的 `OWNER_SPEECH_PROVIDER=cloud|local` 控制,默认先走云端以保证真实可用。 ## Goals @@ -12,8 +12,8 @@ 2. 提供 `--once`,便于单轮真实验收和自动化测试。 3. 使用 `.env` 作为默认配置来源,避免要求用户导出 shell 环境变量。 4. 使用 `sounddevice` 真实读取本机麦克风。 -5. 使用 `sherpa-onnx` 本地 VAD/STT 模型,模型存放在项目 `models/`。 -6. 使用 macOS `say/afplay` 先保证本地 TTS 真实播出。 +5. 使用 `OWNER_SPEECH_PROVIDER=cloud|local` 明确选择语音模型路径。 +6. 默认 cloud 模式使用 NewAPI 的 `mimo-v2.5-asr` 和 `mimo-v2.5-tts`,local 模式保留 `sherpa-onnx` 和本地 TTS 路径。 7. 在同一个 `run-live` 进程内保留最近 user/assistant 历史,进程退出即丢弃。 8. 新增自动化测试证明两轮 runtime、临时上下文、进程隔离和错误恢复。 @@ -34,10 +34,10 @@ CLI -> LiveVoiceRuntimeFactory -> SoundDeviceAudioTransport -> WakeDetector("小杰小杰") - -> SherpaOnnxVadProvider - -> SherpaOnnxSttProvider + -> Energy/Sherpa VAD provider + -> CloudAsrSttProvider or SherpaOnnxSttProvider -> OpenAI/NewAPI LlmProvider - -> MacOsSayTtsProvider + -> CloudTtsProvider or MacOsSayTtsProvider -> ConversationContext(empty per process) -> TerminalRuntimeReporter @@ -58,7 +58,7 @@ LiveVoiceRuntime.run() 1. Parse CLI args. 2. Load `.env` and merge only explicitly supported defaults. 3. Validate LLM config. -4. Validate models when live mode requires local VAD/STT. +4. Validate speech provider mode; cloud mode validates credentials and model names, local mode validates project-local VAD/STT models. 5. Validate audio devices. 6. Create one `ConversationContext` for the process. 7. Enter loop: @@ -184,12 +184,13 @@ Test requirements: ## Wake/VAD/STT Design -First implementation can use one of two local strategies: +First implementation can use one of two strategies: -1. Preferred: sherpa-onnx VAD for endpointing plus sherpa-onnx STT for short wake windows and user utterances. -2. Future replacement: dedicated local KWS provider for “小杰小杰”. +1. Default cloud mode: local VAD cuts speech segments, `mimo-v2.5-asr` transcribes wake and user segments through NewAPI-compatible `/v1/audio/transcriptions`. +2. Local mode: sherpa-onnx VAD/STT uses project-local model assets under `models/`. +3. Future replacement: dedicated local KWS provider for “小杰小杰”. -The implementation must not send raw audio to the cloud for wake, VAD, or STT. If a dedicated wake model is unavailable, wake detection may transcribe short local segments and search for the normalized wake word. +Cloud mode sends only VAD-cut speech segments to ASR, not the continuous microphone stream. Local mode does not send wake, VAD, or STT audio to the cloud. If a dedicated wake model is unavailable, wake detection may transcribe short speech segments and search for the normalized wake word. ## Model Download Design @@ -241,4 +242,4 @@ No persistent data migration is needed. Existing `.env` remains valid. New optio 1. Which exact sherpa-onnx Chinese STT model gives the best latency/accuracy tradeoff on this Mac. 2. Whether future wake detection should use a dedicated KWS model instead of short-window STT. -3. Whether TTS should later move from macOS `say` to sherpa-onnx TTS for consistent voice and offline packaging. +3. Whether local TTS should later move from macOS `say` to sherpa-onnx TTS for consistent voice and offline packaging. diff --git a/openspec/changes/complete-live-repeat-voice-runtime/proposal.md b/openspec/changes/complete-live-repeat-voice-runtime/proposal.md index 757ea6a..012788d 100644 --- a/openspec/changes/complete-live-repeat-voice-runtime/proposal.md +++ b/openspec/changes/complete-live-repeat-voice-runtime/proposal.md @@ -4,7 +4,7 @@ ### 完整业务价值 -本变更把现有语音桌宠从“可用 mock/fixture 验收链路”升级为“可真实运行的无 GUI 实时语音程序”。完成后,用户可以在 macOS 本机启动 `owner_voice_pet run-live`,程序常驻监听本机麦克风,说出唤醒词“小杰小杰”后进入录音,完成 VAD 端点检测、STT 转写、云端 LLM 回复、本地 TTS 语音合成和扬声器播放,然后自动回到待机继续监听下一轮。 +本变更把现有语音桌宠从“可用 mock/fixture 验收链路”升级为“可真实运行的无 GUI 实时语音程序”。完成后,用户可以在 macOS 本机启动 `owner_voice_pet run-live`,程序常驻监听本机麦克风,说出唤醒词“小杰小杰”后进入录音,完成 VAD 端点检测、按 `.env` 选择的 ASR 转写、云端 LLM 回复、按 `.env` 选择的 TTS 语音合成和扬声器播放,然后自动回到待机继续监听下一轮。 该能力的核心价值是让桌宠语音链路从工程骨架变成可反复使用的真实语音入口。用户不需要手动运行一次性 acceptance,也不需要手工提供文本输入;只要进程存活,就可以多轮唤醒、多轮提问、多轮播放,并且同一次运行进程内的第二轮、第三轮会携带前面 user/assistant 历史,从而支持“刚才那句话”“继续解释上一轮”这类连续对话。 @@ -33,8 +33,8 @@ 2. CLI 增加 `run-live`、`model-check`、`device-check`,README 增加 `.env`、模型下载和真实运行说明。 3. Runtime 增加常驻循环,成功或失败都返回待机;`--once` 只用于测试和单轮人工验收。 4. Audio Transport 从占位实现升级为 `sounddevice` 本机麦克风输入和本机扬声器/播放器输出。 -5. VAD/STT 从占位边界升级为 `sherpa-onnx` 本地模型 Provider 或明确的本地 Provider 装载路径;模型下载到项目 `models/`,不提交 Git。 -6. TTS 第一版选择 macOS 本地 `say/afplay`,优先保证真实播出;后续可用 OpenSpec 新变更替换为 `sherpa-onnx` TTS。 +5. ASR/TTS 从占位边界升级为可配置 Provider:`OWNER_SPEECH_PROVIDER=cloud` 默认使用 NewAPI `mimo-v2.5-asr` 和 `mimo-v2.5-tts`,`OWNER_SPEECH_PROVIDER=local` 使用本地模型/本地播放路径。 +6. 本地模型仍下载到项目 `models/`,不提交 Git,并通过 `model-check` 验证。 7. 对话上下文策略明确为“进程内临时历史”:同一个 `run-live` 进程保留多轮 user/assistant,进程退出即丢弃,不落盘。 8. 测试套件新增重复 runtime、临时上下文、进程本地上下文隔离、模型检查、设备检查的自动化覆盖。 @@ -67,7 +67,7 @@ 6. `run-live` 必须使用本机扬声器播放回复音频,TTS 第一版允许通过 macOS `say` 生成音频,再用 `afplay` 或 transport 播放。 7. 唤醒词必须固定为“小杰小杰”;第一版可通过短语音段 STT 命中唤醒词,或通过本地 KWS Provider 命中唤醒词,但检测链路必须在本机执行。 8. 唤醒命中后必须进入录音和 VAD 端点检测;录音结束后进入 STT。 -9. STT 必须优先本地实现,默认候选为 `sherpa-onnx`;模型文件必须位于项目 `models/`。 +9. STT/ASR 必须由 `.env` 的 `OWNER_SPEECH_PROVIDER` 决定;`cloud` 使用 `OWNER_ASR_MODEL=mimo-v2.5-asr`,`local` 使用项目 `models/` 下的本地模型。 10. LLM 必须使用 `.env` 配置的云端 NewAPI/OpenAI 兼容接口;模型名来自配置,不能写死在核心逻辑。 11. 每轮 LLM 请求必须包含 system prompt、本次进程内最近 user/assistant 历史和当前 user message。 12. LLM 回复成功后必须追加 assistant 消息到同一个进程内 `ConversationContext`。 @@ -144,7 +144,10 @@ 4. `OWNER_LLM_API_STYLE`:第一版保留 `chat_completions` 或现有兼容值。 5. `OWNER_CONTEXT_MAX_MESSAGES`:本次进程内历史最大消息数。 6. `OWNER_CONTEXT_MAX_CHARS`:本次进程内历史最大字符数。 -7. `OWNER_SPEECH_MODELS_DIR`:可选,默认 `models/`。 +7. `OWNER_SPEECH_PROVIDER`:`cloud` 或 `local`,默认 `cloud`。 +8. `OWNER_ASR_MODEL`:云 ASR 模型,默认 `mimo-v2.5-asr`。 +9. `OWNER_TTS_MODEL`:云 TTS 模型,默认 `mimo-v2.5-tts`。 +10. `OWNER_SPEECH_MODELS_DIR`:本地模型目录,默认 `models/`。 #### CLI 输入 @@ -185,7 +188,7 @@ 1. 新增 `LiveVoiceRuntime` 或等价 runtime 编排器,管理常驻循环、一次性模式、状态输出、错误恢复和共享 `ConversationContext`。 2. 扩展 `SoundDeviceAudioTransport`,真实打开 `sounddevice` 输入流,提供录音片段读取,并支持本机播放或与 macOS 播放命令协作。 -3. 新增或完善 `SherpaOnnxVadProvider`、`SherpaOnnxSttProvider`,从 `models/` 加载本地模型。 +3. 新增或完善 `CloudAsrSttProvider`、`CloudTtsProvider`、`SherpaOnnxVadProvider`、`SherpaOnnxSttProvider`,由 `OWNER_SPEECH_PROVIDER` 选择云端或本地语音路径。 4. 新增模型下载脚本和检查命令,把模型资产生命周期从“人工假设”变成“可诊断前置条件”。 5. 修改 README 和 `.env.example`,用 `.env` 作为默认配置路径,明确运行命令和本地依赖安装命令。 6. 新增两轮 runtime 测试,证明真实运行层不是单轮 acceptance 的包装。 diff --git a/openspec/changes/complete-live-repeat-voice-runtime/specs/voice-pet-pipeline/spec.md b/openspec/changes/complete-live-repeat-voice-runtime/specs/voice-pet-pipeline/spec.md index 34bd45a..0a43f4b 100644 --- a/openspec/changes/complete-live-repeat-voice-runtime/specs/voice-pet-pipeline/spec.md +++ b/openspec/changes/complete-live-repeat-voice-runtime/specs/voice-pet-pipeline/spec.md @@ -1,7 +1,7 @@ ## ADDED Requirements ### Requirement: Live repeat voice runtime -The system SHALL provide a `run-live` command that performs real repeated voice conversation with local microphone input, local speech processing, cloud LLM reply generation, local TTS, local speaker playback, and automatic return to standby. +The system SHALL provide a `run-live` command that performs real repeated voice conversation with local microphone input, configured speech recognition and speech synthesis providers, cloud LLM reply generation, local speaker playback, and automatic return to standby. #### Scenario: Live runtime starts in standby - **WHEN** the user runs `PYTHONPATH=src python3.11 -m owner_voice_pet run-live` @@ -46,6 +46,21 @@ The system SHALL provide project-local speech model preparation and diagnostics - **WHEN** `sherpa-onnx` is unavailable, a model file is missing, or a model cannot be loaded - **THEN** `model-check` SHALL fail with a structured model error and SHALL NOT start live microphone listening +### Requirement: Configurable speech provider mode +The system SHALL read `OWNER_SPEECH_PROVIDER` from `.env` to choose between cloud speech providers and local speech providers for first-version live runtime. + +#### Scenario: Cloud speech provider is selected +- **WHEN** `OWNER_SPEECH_PROVIDER=cloud` +- **THEN** live ASR SHALL use the configured cloud model from `OWNER_ASR_MODEL`, live TTS SHALL use the configured cloud model from `OWNER_TTS_MODEL`, and the implementation SHALL default those models to `mimo-v2.5-asr` and `mimo-v2.5-tts` + +#### Scenario: Local speech provider is selected +- **WHEN** `OWNER_SPEECH_PROVIDER=local` +- **THEN** live ASR/VAD SHALL use project-local speech model assets and local TTS SHALL use a local playback-capable provider + +#### Scenario: Speech provider is invalid +- **WHEN** `OWNER_SPEECH_PROVIDER` is neither `cloud` nor `local` +- **THEN** startup validation SHALL fail with a structured configuration error + ### Requirement: Live audio device readiness The system SHALL provide live audio device diagnostics and SHALL use `sounddevice` for first-version real microphone and speaker access. @@ -123,11 +138,11 @@ The system SHALL use the local microphone as the first-version input Transport a - **WHEN** the microphone or speaker is missing, denied, unsupported, or inaccessible through `sounddevice` - **THEN** the system SHALL expose a recoverable Transport error with a stable error code and SHALL NOT silently fall back to fixture audio in live mode -### Requirement: Local STT transcription -The system SHALL transcribe captured user utterances through a local STT provider, with `sherpa-onnx` as the first live implementation candidate and project-local models under `models/`. +### Requirement: Configured STT transcription +The system SHALL transcribe captured user utterances through the configured STT provider; cloud mode SHALL use NewAPI-compatible ASR and local mode SHALL use project-local `sherpa-onnx` model assets under `models/`. #### Scenario: STT succeeds -- **WHEN** local STT returns non-empty text for a captured live audio segment +- **WHEN** configured STT returns non-empty text for a captured live audio segment - **THEN** the pipeline SHALL add the trimmed text as a user message to the current process conversation context #### Scenario: STT returns empty text @@ -135,11 +150,11 @@ The system SHALL transcribe captured user utterances through a local STT provide - **THEN** the live runtime SHALL skip LLM invocation and return to standby with a recoverable status #### Scenario: STT provider fails -- **WHEN** the STT provider raises an error or cannot load its local model +- **WHEN** the STT provider raises an error, cloud ASR fails, or local model loading fails - **THEN** the system SHALL emit an STT or model error code and SHALL recover to a state where future wake attempts are possible if startup can continue safely -### Requirement: Local TTS synthesis and playback -The system SHALL synthesize assistant replies through a local TTS provider and SHALL play synthesized speech through the local system output path, with macOS `say/afplay` accepted as the first live implementation. +### Requirement: Configured TTS synthesis and playback +The system SHALL synthesize assistant replies through the configured TTS provider and SHALL play synthesized speech through the local system output path; cloud mode SHALL use NewAPI-compatible TTS and local mode SHALL use a local playback-capable TTS provider. #### Scenario: Reply text is ready for speech - **WHEN** the LLM produces a non-empty reply for a live turn diff --git a/openspec/changes/complete-live-repeat-voice-runtime/tasks.md b/openspec/changes/complete-live-repeat-voice-runtime/tasks.md index 137d2a2..cd550e3 100644 --- a/openspec/changes/complete-live-repeat-voice-runtime/tasks.md +++ b/openspec/changes/complete-live-repeat-voice-runtime/tasks.md @@ -27,21 +27,21 @@ ## 4. Wake/VAD/STT 与模型检查 -- [ ] 4.1 实现 `model-check` CLI;前置条件:模型 manifest 和路径约定完成;验收标准:检查依赖、目录、关键文件和 Provider 可加载性;测试要点:缺模型、缺依赖、成功三类测试;优先级:P0;预计:45 分钟。 -- [ ] 4.2 实现 `SherpaOnnxVadProvider` 或等价本地 VAD 适配;前置条件:模型文件可用;验收标准:可分析帧并输出 speech/silence;测试要点:fake 模型或短音频 fixture;优先级:P0;预计:60 分钟。 -- [ ] 4.3 实现 `SherpaOnnxSttProvider` 真实转写;前置条件:STT 模型文件可用;验收标准:可从 AudioSegment 返回中文文本;测试要点:空音频、无模型、成功 fixture;优先级:P0;预计:60 分钟。 -- [ ] 4.4 实现 live 唤醒检测策略;前置条件:STT/VAD 可用;验收标准:能从实时音频中识别“小杰小杰”并进入录音,不把唤醒词传给 LLM;测试要点:命中/未命中测试;优先级:P0;预计:60 分钟。 -- [ ] 4.5 增加 VAD 端点和 STT 文本有效性测试;前置条件:4.1 至 4.4 完成;验收标准:无语音、空文本、最大录音时长均可恢复待机;测试要点:runtime 状态断言;优先级:P0;预计:45 分钟。 -- [ ] 4.6 验证并提交“Wake/VAD/STT 与模型检查”模块;前置条件:4.1 至 4.5 完成;验收标准:compileall、wake/vad/stt tests、model-check、OpenSpec strict 通过后 commit;测试要点:提交信息为中文格式;优先级:P0;预计:20 分钟。 +- [x] 4.1 实现 `model-check` CLI;前置条件:模型 manifest 和路径约定完成;验收标准:检查依赖、目录、关键文件和 Provider 可加载性;测试要点:缺模型、缺依赖、成功三类测试;优先级:P0;预计:45 分钟。 +- [x] 4.2 实现 `SherpaOnnxVadProvider` 或等价本地 VAD 适配;前置条件:模型文件可用;验收标准:可分析帧并输出 speech/silence;测试要点:fake 模型或短音频 fixture;优先级:P0;预计:60 分钟。 +- [x] 4.3 实现 `SherpaOnnxSttProvider` 真实转写;前置条件:STT 模型文件可用;验收标准:可从 AudioSegment 返回中文文本;测试要点:空音频、无模型、成功 fixture;优先级:P0;预计:60 分钟。 +- [x] 4.4 实现 live 唤醒检测策略;前置条件:STT/VAD 可用;验收标准:能从实时音频中识别“小杰小杰”并进入录音,不把唤醒词传给 LLM;测试要点:命中/未命中测试;优先级:P0;预计:60 分钟。 +- [x] 4.5 增加 VAD 端点和 STT 文本有效性测试;前置条件:4.1 至 4.4 完成;验收标准:无语音、空文本、最大录音时长均可恢复待机;测试要点:runtime 状态断言;优先级:P0;预计:45 分钟。 +- [x] 4.6 验证并提交“Wake/VAD/STT 与模型检查”模块;前置条件:4.1 至 4.5 完成;验收标准:compileall、wake/vad/stt tests、model-check、OpenSpec strict 通过后 commit;测试要点:提交信息为中文格式;优先级:P0;预计:20 分钟。 ## 5. Live runtime、LLM/TTS 与临时上下文 -- [ ] 5.1 新增 `LiveVoiceRuntime` 常驻循环;前置条件:Transport、Wake/VAD/STT 可用;验收标准:默认循环、`--once` 单轮、Ctrl-C 清理;测试要点:fake runtime 两轮;优先级:P0;预计:60 分钟。 -- [ ] 5.2 新增 `run-live` CLI;前置条件:LiveVoiceRuntime 可构造;验收标准:读取 `.env`、构造 Provider、输出中文状态;测试要点:CLI 参数和错误码测试;优先级:P0;预计:60 分钟。 -- [ ] 5.3 接入 LLM 请求临时历史;前置条件:ConversationContext 可复用;验收标准:第二轮请求包含第一轮 user/assistant;新 runtime 上下文为空;测试要点:temporary-context 和 process-local 测试;优先级:P0;预计:60 分钟。 -- [ ] 5.4 接入 macOS `say/afplay` TTS 播放;前置条件:TTS Provider 已有边界;验收标准:非空回复可生成并播放音频;测试要点:mock subprocess 成功/失败;优先级:P0;预计:45 分钟。 -- [ ] 5.5 增加 live 错误恢复;前置条件:runtime 主流程完成;验收标准:STT 空文本、LLM 失败、TTS 失败、播放失败后都恢复待机;测试要点:状态序列测试;优先级:P0;预计:60 分钟。 -- [ ] 5.6 验证并提交“Live runtime 与临时上下文”模块;前置条件:5.1 至 5.5 完成;验收标准:compileall、unittest、security-check、OpenSpec strict 通过后 commit;测试要点:提交信息为中文格式;优先级:P0;预计:20 分钟。 +- [x] 5.1 新增 `LiveVoiceRuntime` 常驻循环;前置条件:Transport、Wake/VAD/STT 可用;验收标准:默认循环、`--once` 单轮、Ctrl-C 清理;测试要点:fake runtime 两轮;优先级:P0;预计:60 分钟。 +- [x] 5.2 新增 `run-live` CLI;前置条件:LiveVoiceRuntime 可构造;验收标准:读取 `.env`、构造 Provider、输出中文状态;测试要点:CLI 参数和错误码测试;优先级:P0;预计:60 分钟。 +- [x] 5.3 接入 LLM 请求临时历史;前置条件:ConversationContext 可复用;验收标准:第二轮请求包含第一轮 user/assistant;新 runtime 上下文为空;测试要点:temporary-context 和 process-local 测试;优先级:P0;预计:60 分钟。 +- [x] 5.4 接入 macOS `say/afplay` TTS 播放;前置条件:TTS Provider 已有边界;验收标准:非空回复可生成并播放音频;测试要点:mock subprocess 成功/失败;优先级:P0;预计:45 分钟。 +- [x] 5.5 增加 live 错误恢复;前置条件:runtime 主流程完成;验收标准:STT 空文本、LLM 失败、TTS 失败、播放失败后都恢复待机;测试要点:状态序列测试;优先级:P0;预计:60 分钟。 +- [x] 5.6 验证并提交“Live runtime 与临时上下文”模块;前置条件:5.1 至 5.5 完成;验收标准:compileall、unittest、security-check、OpenSpec strict 通过后 commit;测试要点:提交信息为中文格式;优先级:P0;预计:20 分钟。 ## 6. 文档、真实验收、归档 diff --git a/src/owner_voice_pet/__init__.py b/src/owner_voice_pet/__init__.py index b44f566..58faabe 100644 --- a/src/owner_voice_pet/__init__.py +++ b/src/owner_voice_pet/__init__.py @@ -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", diff --git a/src/owner_voice_pet/cli.py b/src/owner_voice_pet/cli.py index c39ed03..7a8a071 100644 --- a/src/owner_voice_pet/cli.py +++ b/src/owner_voice_pet/cli.py @@ -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 diff --git a/src/owner_voice_pet/config.py b/src/owner_voice_pet/config.py index 8c4efe2..430b5f1 100644 --- a/src/owner_voice_pet/config.py +++ b/src/owner_voice_pet/config.py @@ -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(): diff --git a/src/owner_voice_pet/llm.py b/src/owner_voice_pet/llm.py index f322c52..16b43d0 100644 --- a/src/owner_voice_pet/llm.py +++ b/src/owner_voice_pet/llm.py @@ -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, ) diff --git a/src/owner_voice_pet/runtime.py b/src/owner_voice_pet/runtime.py new file mode 100644 index 0000000..dd10931 --- /dev/null +++ b/src/owner_voice_pet/runtime.py @@ -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()) diff --git a/src/owner_voice_pet/speech_models.py b/src/owner_voice_pet/speech_models.py index c9d3de4..cf5ca9f 100644 --- a/src/owner_voice_pet/speech_models.py +++ b/src/owner_voice_pet/speech_models.py @@ -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" diff --git a/src/owner_voice_pet/stt.py b/src/owner_voice_pet/stt.py index 33d34e3..43fe084 100644 --- a/src/owner_voice_pet/stt.py +++ b/src/owner_voice_pet/stt.py @@ -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) diff --git a/src/owner_voice_pet/transport.py b/src/owner_voice_pet/transport.py index c111b46..50a3dc8 100644 --- a/src/owner_voice_pet/transport.py +++ b/src/owner_voice_pet/transport.py @@ -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( diff --git a/src/owner_voice_pet/tts.py b/src/owner_voice_pet/tts.py index f8b6146..63f7c01 100644 --- a/src/owner_voice_pet/tts.py +++ b/src/owner_voice_pet/tts.py @@ -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"}) diff --git a/src/owner_voice_pet/vad.py b/src/owner_voice_pet/vad.py index ea337ed..f7cc5c2 100644 --- a/src/owner_voice_pet/vad.py +++ b/src/owner_voice_pet/vad.py @@ -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 diff --git a/tests/test_cli_acceptance.py b/tests/test_cli_acceptance.py index 7bfa808..6fb6055 100644 --- a/tests/test_cli_acceptance.py +++ b/tests/test_cli_acceptance.py @@ -57,11 +57,18 @@ class CliAcceptanceTests(unittest.TestCase): path = root / relative path.parent.mkdir(parents=True, exist_ok=True) path.write_bytes(b"placeholder") - with patch("importlib.util.find_spec", return_value=object()): + with ( + patch("importlib.util.find_spec", return_value=object()), + patch("owner_voice_pet.cli.SherpaOnnxVadProvider") as vad_cls, + patch("owner_voice_pet.cli.SherpaOnnxSttProvider") as stt_cls, + ): + vad_cls.return_value.load.return_value = None + stt_cls.return_value.load.return_value = None code, data = self.call("model-check", "--models-dir", str(root)) self.assertEqual(code, 0) self.assertTrue(data["ok"]) self.assertEqual(data["missing_files"], []) + self.assertTrue(data["provider_load_checked"]) def test_model_check_reports_missing_files(self) -> None: with tempfile.TemporaryDirectory() as tmp: @@ -77,6 +84,18 @@ class CliAcceptanceTests(unittest.TestCase): self.assertEqual(code, 0) self.assertTrue(data["ok"]) + def test_run_live_once_invokes_runtime(self) -> None: + class FakeRuntime: + def run(self, *, once: bool = False): + self.once = once + return type("Summary", (), {"completed_turns": 1, "interrupted": False})() + + fake_runtime = FakeRuntime() + with patch("owner_voice_pet.cli.build_live_runtime", return_value=fake_runtime): + code = main(["run-live", "--once"]) + self.assertEqual(code, 0) + self.assertTrue(fake_runtime.once) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_live_runtime.py b/tests/test_live_runtime.py new file mode 100644 index 0000000..050271d --- /dev/null +++ b/tests/test_live_runtime.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import unittest + +from owner_voice_pet.config import AppConfig +from owner_voice_pet.conversation import ConversationContext +from owner_voice_pet.llm import MockLlmProvider +from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript +from owner_voice_pet.runtime import LiveVoiceRuntime +from owner_voice_pet.transport import MemoryAudioTransport +from owner_voice_pet.tts import SineTtsProvider +from owner_voice_pet.vad import EnergyVadProvider, VadRecorder + + +def segment_frames(start_id: int, start_ms: int) -> list[AudioFrame]: + return [ + AudioFrame(b"\xff\x7f", 16000, 1, start_ms, start_id, {"duration_ms": 20, "speech": True}), + AudioFrame(b"\xff\x7f", 16000, 1, start_ms + 20, start_id + 1, {"duration_ms": 20, "speech": True}), + AudioFrame(b"\x00\x00", 16000, 1, start_ms + 40, start_id + 2, {"duration_ms": 20, "speech": False}), + AudioFrame(b"\x00\x00", 16000, 1, start_ms + 60, start_id + 3, {"duration_ms": 20, "speech": False}), + ] + + +class QueueSttProvider: + def __init__(self, texts: list[str]) -> None: + self.texts = list(texts) + self.calls: list[AudioSegment] = [] + self.loaded = False + + def load(self) -> None: + self.loaded = True + + def transcribe(self, segment: AudioSegment) -> Transcript: + self.calls.append(segment) + text = self.texts.pop(0) + return Transcript(text, "zh", 1.0, segment.duration_ms, "queue-stt") + + +class RecordingReporter: + def __init__(self) -> None: + self.statuses: list[str] = [] + self.errors: list[str] = [] + + def status(self, state: str, message: str, *, turn_id: int | None = None) -> None: + self.statuses.append(message) + + def error(self, stage: str, code: str, message: str, *, turn_id: int | None = None) -> None: + self.errors.append(f"{stage}:{code}:{message}") + + +def make_runtime(texts: list[str], context: ConversationContext | None = None) -> tuple[LiveVoiceRuntime, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]: + frames = [] + for idx in range(4): + frames.extend(segment_frames(idx * 4, idx * 80)) + transport = MemoryAudioTransport(frames) + stt = QueueSttProvider(texts) + llm = MockLlmProvider(["这是答复。"]) + tts = SineTtsProvider() + reporter = RecordingReporter() + runtime = LiveVoiceRuntime( + config=AppConfig(llm_api_key="secret", speech_provider="cloud"), + transport=transport, + vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40), + stt=stt, + llm=llm, + tts=tts, + context=context or ConversationContext(), + reporter=reporter, + ) + return runtime, stt, llm, transport, reporter + + +class LiveRuntimeTests(unittest.TestCase): + def test_repeated_runtime_runs_two_turns_and_returns_to_standby(self) -> None: + runtime, stt, llm, transport, reporter = make_runtime( + ["小杰小杰", "第一问", "小杰小杰", "第二问"] + ) + summary = runtime.run(max_turns=2) + + self.assertEqual(summary.completed_turns, 2) + self.assertEqual(len(stt.calls), 4) + self.assertEqual(len(llm.calls), 2) + self.assertEqual(len(transport.played_segments), 2) + self.assertIn("恢复待机:可继续唤醒", reporter.statuses[-1]) + + def test_temporary_context_is_sent_to_second_llm_call(self) -> None: + runtime, _, llm, _, _ = make_runtime(["小杰小杰", "第一问", "小杰小杰", "第二问"]) + runtime.run(max_turns=2) + + second_call_text = [message.content for message in llm.calls[1]] + self.assertIn("第一问", second_call_text) + self.assertIn("这是答复。", second_call_text) + self.assertEqual(second_call_text[-1], "第二问") + + def test_new_runtime_context_starts_empty(self) -> None: + first_context = ConversationContext() + first_runtime, _, _, _, _ = make_runtime(["小杰小杰", "第一问"], context=first_context) + first_runtime.run(max_turns=1) + self.assertGreater(len(first_context.messages()), 0) + + second_context = ConversationContext() + make_runtime(["小杰小杰", "第二问"], context=second_context) + self.assertEqual(second_context.messages(), ()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_models_config.py b/tests/test_models_config.py index 9ebf622..9b32382 100644 --- a/tests/test_models_config.py +++ b/tests/test_models_config.py @@ -59,25 +59,46 @@ class ModelsConfigTests(unittest.TestCase): handle.write( "\n".join( [ - "OWNER_LLM_BASE_URL=https://newapi.mkbk.shop/", + "OWNER_LLM_BASE_URL=https://token-plan-cn.xiaomimimo.com/v1/", "OWNER_LLM_API_KEY=secret-value", "OWNER_LLM_MODEL=test-model", ] ) ) config = AppConfig.from_dotenv(path) - self.assertEqual(config.llm_base_url, "https://newapi.mkbk.shop") + self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1") self.assertEqual(config.llm_api_key, "secret-value") self.assertEqual(config.llm_model, "test-model") + self.assertEqual(config.speech_provider, "cloud") + self.assertEqual(config.asr_model, "mimo-v2.5-asr") + self.assertEqual(config.tts_model, "mimo-v2.5-tts") + self.assertEqual(config.tts_voice, "alloy") self.assertEqual(str(config.speech_models_dir), "models") self.assertTrue(config.llm_stream) self.assertEqual(config.validate_basic(), []) + def test_speech_provider_must_be_cloud_or_local(self) -> None: + config = AppConfig(speech_provider="invalid") + errors = config.validate_basic() + self.assertTrue(any("OWNER_SPEECH_PROVIDER" in error.message for error in errors)) + def test_missing_dotenv_uses_non_secret_defaults(self) -> None: config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env") - self.assertEqual(config.llm_base_url, "https://newapi.mkbk.shop") + self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1") self.assertIsNone(config.llm_api_key) + def test_api_url_accepts_base_with_or_without_v1(self) -> None: + with_v1 = AppConfig(llm_base_url="https://token-plan-cn.xiaomimimo.com/v1") + without_v1 = AppConfig(llm_base_url="https://newapi.mkbk.shop") + self.assertEqual( + with_v1.api_url("/v1/chat/completions"), + "https://token-plan-cn.xiaomimimo.com/v1/chat/completions", + ) + self.assertEqual( + without_v1.api_url("/v1/chat/completions"), + "https://newapi.mkbk.shop/v1/chat/completions", + ) + def test_missing_llm_key_has_structured_error(self) -> None: config = AppConfig(llm_api_key=None) with self.assertRaises(ProviderError) as raised: diff --git a/tests/test_pipeline_llm_tts.py b/tests/test_pipeline_llm_tts.py index 2725e9d..d3495fa 100644 --- a/tests/test_pipeline_llm_tts.py +++ b/tests/test_pipeline_llm_tts.py @@ -11,7 +11,7 @@ from owner_voice_pet.models import AudioFrame, ErrorCode, Message, PipelineState from owner_voice_pet.pipeline import VoicePipeline from owner_voice_pet.stt import MetadataSttProvider from owner_voice_pet.transport import MemoryAudioTransport -from owner_voice_pet.tts import SentenceBuffer, SineTtsProvider +from owner_voice_pet.tts import CloudTtsProvider, SentenceBuffer, SineTtsProvider from owner_voice_pet.vad import EnergyVadProvider, VadRecorder from owner_voice_pet.wakeword import KeywordWakeWordProvider @@ -66,6 +66,36 @@ class PipelineLlmTtsTests(unittest.TestCase): self.assertGreater(len(segment.pcm), 0) self.assertGreater(segment.duration_ms, 0) + def test_cloud_tts_posts_speech_request(self) -> None: + class FakeResponse: + def __enter__(self): + return self + + def __exit__(self, *args) -> None: + return None + + def read(self) -> bytes: + return b"fake-mp3" + + requests = [] + + def fake_urlopen(request, timeout): + requests.append(request) + return FakeResponse() + + provider = CloudTtsProvider( + AppConfig(llm_api_key="secret", tts_model="mimo-v2.5-tts"), + urlopen=fake_urlopen, + ) + provider.load() + segment = provider.synthesize("你好") + + self.assertEqual(segment.metadata["format"], "mp3") + self.assertEqual(segment.pcm, b"fake-mp3") + body = json.loads(requests[0].data.decode()) + self.assertEqual(body["model"], "mimo-v2.5-tts") + self.assertIn("/v1/audio/speech", requests[0].full_url) + def test_pipeline_runs_from_wake_to_playback(self) -> None: frames = [ frame(0, 0, {"wake_word": "小杰小杰", "wake_confidence": 0.95}), @@ -136,7 +166,7 @@ class PipelineLlmTtsTests(unittest.TestCase): return FakeResponse() config = AppConfig( - llm_base_url="https://newapi.mkbk.shop", + llm_base_url="https://token-plan-cn.xiaomimimo.com/v1", llm_api_key="secret", llm_model="test-model", llm_api_style="chat_completions", @@ -160,7 +190,7 @@ class PipelineLlmTtsTests(unittest.TestCase): return json.dumps({"choices": [{"message": {"content": "非流式回复。"}}]}).encode() config = AppConfig( - llm_base_url="https://newapi.mkbk.shop", + llm_base_url="https://token-plan-cn.xiaomimimo.com/v1", llm_api_key="secret", llm_model="test-model", llm_api_style="chat_completions", diff --git a/tests/test_wake_vad_stt.py b/tests/test_wake_vad_stt.py index 5b132f3..f877c0b 100644 --- a/tests/test_wake_vad_stt.py +++ b/tests/test_wake_vad_stt.py @@ -1,10 +1,12 @@ from __future__ import annotations +import json import tempfile import unittest from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError -from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text +from owner_voice_pet.config import AppConfig +from owner_voice_pet.stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from owner_voice_pet.vad import EnergyVadProvider, VadRecorder from owner_voice_pet.wakeword import KeywordWakeWordProvider, MissingWakeWordModelProvider @@ -102,6 +104,34 @@ class WakeVadSttTests(unittest.TestCase): provider.load() self.assertEqual(raised.exception.code, ErrorCode.STT_MODEL_MISSING) + def test_cloud_asr_posts_audio_transcription_request(self) -> None: + class FakeResponse: + def __enter__(self): + return self + + def __exit__(self, *args) -> None: + return None + + def read(self) -> bytes: + return json.dumps({"text": "你好小杰", "language": "zh"}, ensure_ascii=False).encode() + + requests = [] + + def fake_urlopen(request, timeout): + requests.append(request) + return FakeResponse() + + provider = CloudAsrSttProvider( + AppConfig(llm_api_key="secret", asr_model="mimo-v2.5-asr"), + urlopen=fake_urlopen, + ) + provider.load() + transcript = provider.transcribe(AudioSegment(b"\x00\x00\x01\x00", 16000, 1, 0, 100)) + + self.assertEqual(transcript.text, "你好小杰") + self.assertIn("/v1/audio/transcriptions", requests[0].full_url) + self.assertIn(b'mimo-v2.5-asr', requests[0].data) + if __name__ == "__main__": unittest.main()