refactor: 迁移 agent/mcp/skills/voice 模块到独立 crate

- 创建 proxycast-agent crate(event_converter, mcp_bridge, prompt)
- 创建 proxycast-skills crate(ExecutionCallback/LlmProvider trait, skill_loader)
- 扩展 voice-core crate(device, threaded_recorder, types)
- 主 crate 各模块替换为 re-export 层 + Tauri 实现
- 更新 processor/mod.rs 为纯 re-export
- 更新 app/runner.rs MCP 初始化
- 添加 capabilities/default.json
- 清理已迁移的测试文件
This commit is contained in:
coso
2026-02-08 22:42:01 +08:00
parent bb319d7bb0
commit 98c124b0bb
36 changed files with 859 additions and 2214 deletions
+1
View File
@@ -10428,6 +10428,7 @@ dependencies = [
"futures-util",
"hmac",
"hound",
"parking_lot",
"reqwest 0.12.28",
"serde",
"serde_json",
+44
View File
@@ -0,0 +1,44 @@
{
"$schema": "../gen/schemas/desktop-schema.json",
"identifier": "default",
"description": "ProxyCast 默认权限配置",
"windows": ["main", "smart-input"],
"permissions": [
"core:default",
"core:event:default",
"core:event:allow-listen",
"core:event:allow-emit",
"core:window:default",
"core:window:allow-show",
"core:window:allow-hide",
"core:window:allow-close",
"core:window:allow-set-focus",
"core:window:allow-set-size",
"core:window:allow-set-position",
"core:window:allow-center",
"core:window:allow-set-fullscreen",
"core:window:allow-is-fullscreen",
"core:window:allow-set-title",
"core:window:allow-inner-size",
"core:window:allow-outer-size",
"core:webview:default",
"core:app:default",
"core:resources:default",
"core:image:default",
"core:tray:default",
"core:menu:default",
"shell:default",
"shell:allow-open",
"dialog:default",
"dialog:allow-open",
"dialog:allow-save",
"dialog:allow-message",
"dialog:allow-ask",
"dialog:allow-confirm",
"global-shortcut:allow-register",
"global-shortcut:allow-unregister",
"global-shortcut:allow-is-registered",
"autostart:default",
"deep-link:default"
]
}
+19
View File
@@ -0,0 +1,19 @@
[package]
name = "proxycast-agent"
version.workspace = true
edition.workspace = true
authors.workspace = true
repository.workspace = true
[dependencies]
proxycast-core.workspace = true
proxycast-mcp.workspace = true
aster.workspace = true
rmcp.workspace = true
serde.workspace = true
serde_json.workspace = true
tokio.workspace = true
tokio-util.workspace = true
async-trait.workspace = true
tracing.workspace = true
chrono.workspace = true
+12
View File
@@ -0,0 +1,12 @@
//! ProxyCast Agent Crate
//!
//! 包含 Agent 模块中不依赖主 crate 内部模块的纯逻辑部分。
//! 深耦合部分(aster_state、aster_agent、credential_bridge、subagent_scheduler)
//! 留在主 crate。
pub mod event_converter;
pub mod mcp_bridge;
pub mod prompt;
pub use event_converter::{convert_agent_event, convert_to_tauri_message, TauriAgentEvent};
pub use prompt::SystemPromptBuilder;
@@ -5,9 +5,8 @@
use aster::agents::mcp_client::{Error, McpClientTrait};
use rmcp::model::{
CallToolResult, GetPromptResult, InitializeResult, JsonObject,
ListPromptsResult, ListResourcesResult, ListToolsResult,
ReadResourceResult, ServerNotification,
CallToolResult, GetPromptResult, InitializeResult, JsonObject, ListPromptsResult,
ListResourcesResult, ListToolsResult, ReadResourceResult, ServerNotification,
};
use rmcp::service::RunningService;
use rmcp::RoleClient;
@@ -16,7 +15,7 @@ use std::sync::Arc;
use tokio::sync::{mpsc, Mutex};
use tokio_util::sync::CancellationToken;
use crate::mcp::client::ProxyCastMcpClient;
use proxycast_mcp::client::ProxyCastMcpClient;
/// MCP 桥接客户端
///
@@ -43,7 +43,6 @@ impl SystemPromptOptions {
}
}
/// System Prompt 构建器
pub struct SystemPromptBuilder {
options: SystemPromptOptions,
@@ -132,7 +131,6 @@ impl SystemPromptBuilder {
prompt
}
/// 构建环境信息部分
fn build_environment_info(&self) -> String {
let mut info = String::from("# 环境信息\n\n");
@@ -175,9 +173,7 @@ mod tests {
#[test]
fn test_build_with_working_dir() {
let prompt = SystemPromptBuilder::new()
.working_dir("/tmp/test")
.build();
let prompt = SystemPromptBuilder::new().working_dir("/tmp/test").build();
assert!(prompt.contains("/tmp/test"));
}
}
@@ -7,8 +7,8 @@
//! - templates - 提示词模板定义
//! - builder - 提示词构建器
pub mod templates;
pub mod builder;
pub mod templates;
pub use builder::SystemPromptBuilder;
pub use templates::*;
@@ -16,7 +16,6 @@ pub const CORE_IDENTITY: &str = r#"你是 ProxyCast Agent,一个强大的 AI
- 拒绝破坏性技术、DoS 攻击、大规模攻击、供应链攻击的请求
- 永远不要生成或猜测 URL,除非你确信这些 URL 是用于帮助用户编程"#;
/// 工具使用指南
pub const TOOL_GUIDELINES: &str = r#"# 工具使用策略
@@ -46,7 +45,6 @@ pub const TOOL_GUIDELINES: &str = r#"# 工具使用策略
3. **先读后改**:修改文件前必须先读取文件内容
4. **最小权限**:只执行必要的操作,避免不必要的文件修改"#;
/// 代码编写指南
pub const CODING_GUIDELINES: &str = r#"# 代码编写指南
@@ -70,7 +68,6 @@ pub const CODING_GUIDELINES: &str = r#"# 代码编写指南
- 优先编辑现有文件而不是创建新文件
- 删除未使用的代码,不要留下注释掉的代码"#;
/// 任务管理指南
pub const TASK_MANAGEMENT: &str = r#"# 任务管理
@@ -89,7 +86,6 @@ pub const TASK_MANAGEMENT: &str = r#"# 任务管理
不要批量完成多个任务后再标记,应该完成一个标记一个。"#;
/// Git 操作指南
pub const GIT_GUIDELINES: &str = r#"# Git 操作
@@ -101,7 +97,6 @@ pub const GIT_GUIDELINES: &str = r#"# Git 操作
- 在 amend 之前:始终检查作者信息(git log -1 --format='%an %ae')
- 永远不要提交更改,除非用户明确要求"#;
/// 输出风格指南
pub const OUTPUT_STYLE: &str = r#"# 输出风格
+14
View File
@@ -0,0 +1,14 @@
[package]
name = "proxycast-skills"
version.workspace = true
edition.workspace = true
authors.workspace = true
repository.workspace = true
[dependencies]
serde.workspace = true
serde_json.workspace = true
async-trait.workspace = true
tracing.workspace = true
regex.workspace = true
dirs.workspace = true
@@ -0,0 +1,70 @@
//! Skill 执行回调 trait 和 Payload 类型
//!
//! 定义 Skill 执行过程中的回调接口和事件数据类型。
//! Tauri 实现(TauriExecutionCallback)留在主 crate。
use serde::Serialize;
/// 步骤开始事件 Payload
#[derive(Debug, Clone, Serialize)]
pub struct StepStartPayload {
pub execution_id: String,
pub step_id: String,
pub step_name: String,
pub current_step: usize,
pub total_steps: usize,
}
/// 步骤完成事件 Payload
#[derive(Debug, Clone, Serialize)]
pub struct StepCompletePayload {
pub execution_id: String,
pub step_id: String,
pub output: String,
}
/// 步骤错误事件 Payload
#[derive(Debug, Clone, Serialize)]
pub struct StepErrorPayload {
pub execution_id: String,
pub step_id: String,
pub error: String,
pub will_retry: bool,
}
/// 执行完成事件 Payload
#[derive(Debug, Clone, Serialize)]
pub struct ExecutionCompletePayload {
pub execution_id: String,
pub success: bool,
pub output: Option<String>,
pub error: Option<String>,
}
/// Tauri 事件名称常量
pub mod events {
pub const STEP_START: &str = "skill:step_start";
pub const STEP_COMPLETE: &str = "skill:step_complete";
pub const STEP_ERROR: &str = "skill:step_error";
pub const COMPLETE: &str = "skill:complete";
}
/// ExecutionCallback Trait
///
/// 定义 Skill 执行过程中的回调接口。
/// 应用层需要实现此 trait 以接收执行进度更新。
pub trait ExecutionCallback: Send + Sync {
fn on_step_start(
&self,
step_id: &str,
step_name: &str,
current_step: usize,
total_steps: usize,
);
fn on_step_complete(&self, step_id: &str, output: &str);
fn on_step_error(&self, step_id: &str, error: &str, will_retry: bool);
fn on_complete(&self, success: bool, final_output: Option<&str>, error: Option<&str>);
}
+19
View File
@@ -0,0 +1,19 @@
//! ProxyCast Skills Crate
//!
//! 包含 Skills 系统的 trait 定义和纯逻辑部分。
//! Tauri 相关的实现(TauriExecutionCallback、ProxyCastLlmProvider)留在主 crate。
mod execution_callback;
mod llm_provider;
mod skill_loader;
pub use execution_callback::{
events, ExecutionCallback, ExecutionCompletePayload, StepCompletePayload, StepErrorPayload,
StepStartPayload,
};
pub use llm_provider::{LlmProvider, SkillError};
pub use skill_loader::{
find_skill_by_name, get_proxycast_skills_dir, load_skill_from_file, load_skills_from_directory,
parse_allowed_tools, parse_boolean, parse_skill_frontmatter, LoadedSkillDefinition,
SkillFrontmatter,
};
@@ -0,0 +1,41 @@
//! LLM Provider trait 和错误类型
//!
//! 定义 Skill 执行引擎调用 LLM 的接口。
//! 具体实现(ProxyCastLlmProvider)留在主 crate。
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
/// Skill 执行错误类型
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum SkillError {
ProviderError(String),
ExecutionError(String),
ConfigError(String),
}
impl std::fmt::Display for SkillError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SkillError::ProviderError(msg) => write!(f, "Provider error: {}", msg),
SkillError::ExecutionError(msg) => write!(f, "Execution error: {}", msg),
SkillError::ConfigError(msg) => write!(f, "Config error: {}", msg),
}
}
}
impl std::error::Error for SkillError {}
/// LLM Provider Trait
///
/// 定义 Skill 执行引擎调用 LLM 的接口。
/// 应用层需要实现此 trait 以提供 LLM 调用能力。
#[async_trait]
pub trait LlmProvider: Send + Sync {
async fn chat(
&self,
system_prompt: &str,
user_message: &str,
model: Option<&str>,
) -> Result<String, SkillError>;
}
@@ -1,7 +1,6 @@
//! Skill 定义加载器
//!
//! 负责从 `~/.proxycast/skills/<skill>/SKILL.md` 加载并解析 Skill 定义。
//! 命令层只负责编排执行,不再持有文件解析细节。
use std::path::{Path, PathBuf};
@@ -9,63 +8,42 @@ use serde::{Deserialize, Serialize};
/// Skill 前置元数据
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub(crate) struct SkillFrontmatter {
/// Skill 名称
pub struct SkillFrontmatter {
pub name: Option<String>,
/// Skill 描述
pub description: Option<String>,
/// 允许的工具
#[serde(rename = "allowed-tools")]
pub allowed_tools: Option<String>,
/// 参数提示
#[serde(rename = "argument-hint")]
pub argument_hint: Option<String>,
/// 使用场景
#[serde(rename = "when-to-use")]
pub when_to_use: Option<String>,
/// 版本
pub version: Option<String>,
/// 偏好模型
pub model: Option<String>,
/// 偏好 Provider
pub provider: Option<String>,
/// 是否禁用模型调用
#[serde(rename = "disable-model-invocation")]
pub disable_model_invocation: Option<String>,
/// 执行模式
#[serde(rename = "execution-mode")]
pub execution_mode: Option<String>,
}
/// 内部 Skill 定义(用于加载和执行)
#[derive(Debug, Clone)]
pub(crate) struct LoadedSkillDefinition {
/// Skill 名称
pub struct LoadedSkillDefinition {
pub skill_name: String,
/// 显示名称
pub display_name: String,
/// 描述
pub description: String,
/// Markdown 内容(System Prompt)
pub markdown_content: String,
/// 允许的工具
pub allowed_tools: Option<Vec<String>>,
/// 参数提示
pub argument_hint: Option<String>,
/// 使用场景
pub when_to_use: Option<String>,
/// 偏好模型
pub model: Option<String>,
/// 偏好 Provider
pub provider: Option<String>,
/// 是否禁用模型调用
pub disable_model_invocation: bool,
/// 执行模式
pub execution_mode: String,
}
/// 解析 Skill 文件的 frontmatter
pub(crate) fn parse_skill_frontmatter(content: &str) -> (SkillFrontmatter, String) {
pub fn parse_skill_frontmatter(content: &str) -> (SkillFrontmatter, String) {
let regex = regex::Regex::new(r"^---\s*\n([\s\S]*?)---\s*\n?").unwrap();
if let Some(captures) = regex.captures(content) {
@@ -111,7 +89,7 @@ pub(crate) fn parse_skill_frontmatter(content: &str) -> (SkillFrontmatter, Strin
}
/// 解析 allowed-tools 字段
pub(crate) fn parse_allowed_tools(value: Option<&str>) -> Option<Vec<String>> {
pub fn parse_allowed_tools(value: Option<&str>) -> Option<Vec<String>> {
value.and_then(|v| {
if v.is_empty() {
return None;
@@ -130,7 +108,7 @@ pub(crate) fn parse_allowed_tools(value: Option<&str>) -> Option<Vec<String>> {
}
/// 解析布尔值字段
pub(crate) fn parse_boolean(value: Option<&str>, default: bool) -> bool {
pub fn parse_boolean(value: Option<&str>, default: bool) -> bool {
value
.map(|v| {
let lower = v.to_lowercase();
@@ -140,7 +118,7 @@ pub(crate) fn parse_boolean(value: Option<&str>, default: bool) -> bool {
}
/// 从文件加载 Skill 定义
pub(crate) fn load_skill_from_file(
pub fn load_skill_from_file(
skill_name: &str,
file_path: &Path,
) -> Result<LoadedSkillDefinition, String> {
@@ -178,12 +156,12 @@ pub(crate) fn load_skill_from_file(
}
/// 获取 ProxyCast Skills 目录
pub(crate) fn get_proxycast_skills_dir() -> Option<PathBuf> {
pub fn get_proxycast_skills_dir() -> Option<PathBuf> {
dirs::home_dir().map(|home| home.join(".proxycast").join("skills"))
}
/// 从目录加载所有 Skills
pub(crate) fn load_skills_from_directory(dir_path: &Path) -> Vec<LoadedSkillDefinition> {
pub fn load_skills_from_directory(dir_path: &Path) -> Vec<LoadedSkillDefinition> {
let mut results = Vec::new();
if !dir_path.exists() {
@@ -216,7 +194,7 @@ pub(crate) fn load_skills_from_directory(dir_path: &Path) -> Vec<LoadedSkillDefi
}
/// 根据名称查找 Skill
pub(crate) fn find_skill_by_name(skill_name: &str) -> Result<LoadedSkillDefinition, String> {
pub fn find_skill_by_name(skill_name: &str) -> Result<LoadedSkillDefinition, String> {
let skills_dir =
get_proxycast_skills_dir().ok_or_else(|| "无法获取 Skills 目录".to_string())?;
+1
View File
@@ -32,6 +32,7 @@ reqwest = { version = "0.12", features = ["json", "multipart"] }
# 异步运行时
tokio = { version = "1", features = ["sync", "time"] }
parking_lot = "0.12"
# WebSocket 客户端(讯飞 ASR)
tokio-tungstenite = { version = "0.24", features = ["native-tls"] }
+2
View File
@@ -16,7 +16,9 @@ src/
├── lib.rs # 库入口
├── types.rs # 类型定义
├── error.rs # 错误类型
├── device.rs # 音频设备枚举
├── recorder.rs # 音频录制
├── threaded_recorder.rs # 线程化录音服务(可跨线程控制)
├── transcriber.rs # Whisper 本地识别
├── output.rs # 文字输出
└── asr_client/ # 云端 ASR
+43
View File
@@ -0,0 +1,43 @@
//! 音频输入设备枚举
//!
//! 提供跨平台的麦克风设备发现能力。
use cpal::traits::{DeviceTrait, HostTrait};
use serde::{Deserialize, Serialize};
use crate::error::{Result, VoiceError};
/// 麦克风设备信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AudioDeviceInfo {
/// 设备 ID(用于选择设备)
pub id: String,
/// 设备名称
pub name: String,
/// 是否为默认设备
pub is_default: bool,
}
/// 获取所有可用的麦克风设备
pub fn list_audio_devices() -> Result<Vec<AudioDeviceInfo>> {
let host = cpal::default_host();
let default_device = host.default_input_device();
let default_name = default_device.as_ref().and_then(|d| d.name().ok());
let devices = host
.input_devices()
.map_err(|e| VoiceError::RecorderError(format!("无法枚举音频设备: {e}")))?
.filter_map(|device| {
let name = device.name().ok()?;
let is_default = default_name.as_ref().map(|n| n == &name).unwrap_or(false);
Some(AudioDeviceInfo {
id: name.clone(),
name,
is_default,
})
})
.collect();
Ok(devices)
}
+4
View File
@@ -4,16 +4,20 @@
//! 不依赖 Tauri,可被任何 Rust 项目使用。
pub mod asr_client;
pub mod device;
pub mod error;
pub mod output;
pub mod recorder;
pub mod threaded_recorder;
#[cfg(feature = "local-whisper")]
pub mod transcriber;
pub mod types;
pub use device::{list_audio_devices, AudioDeviceInfo};
pub use error::{Result, VoiceError};
pub use output::OutputHandler;
pub use recorder::AudioRecorder;
pub use threaded_recorder::{RecordingCommand, RecordingResponse, RecordingService};
#[cfg(feature = "local-whisper")]
pub use transcriber::WhisperTranscriber;
pub use types::*;
@@ -0,0 +1,469 @@
//! 录音服务
//!
//! 管理录音状态,提供录音控制接口。
//!
//! ## 线程安全设计
//!
//! 由于 `cpal::Stream` 不实现 `Send` trait,无法直接在 Tauri 的 async 命令中使用。
//! 本模块采用**独立线程 + channel 通信**的方案:
//!
//! ```text
//! ┌─────────────────┐ Command ┌─────────────────┐
//! │ Tauri Command │ ───────────────> │ Recording │
//! │ (async) │ │ Thread │
//! │ │ <─────────────── │ (owns Stream) │
//! └─────────────────┘ Response └─────────────────┘
//! ```
//!
//! - 录音线程拥有 `cpal::Stream`,在独立线程中运行
//! - Tauri 命令通过 channel 发送控制指令
//! - 录音线程通过 channel 返回结果
use parking_lot::Mutex;
use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use std::sync::mpsc::{self, Receiver, Sender};
use std::sync::Arc;
use std::thread::{self, JoinHandle};
use std::time::Instant;
use crate::types::AudioData;
/// 录音控制命令
#[derive(Debug)]
pub enum RecordingCommand {
/// 开始录音(可选指定设备 ID)
Start(Option<String>),
/// 停止录音
Stop,
/// 取消录音
Cancel,
/// 关闭录音线程
Shutdown,
}
/// 录音响应
#[derive(Debug)]
pub enum RecordingResponse {
/// 操作成功
Ok,
/// 停止录音成功,返回音频数据
AudioData(AudioData),
/// 操作失败
Error(String),
}
/// 录音服务
///
/// 使用独立线程管理 cpal::Stream,通过 channel 与 Tauri 命令通信
pub struct RecordingService {
/// 命令发送端
command_tx: Option<Sender<RecordingCommand>>,
/// 响应接收端
response_rx: Option<Receiver<RecordingResponse>>,
/// 录音线程句柄
thread_handle: Option<JoinHandle<()>>,
/// 是否正在录音(共享状态,用于快速查询)
is_recording: Arc<AtomicBool>,
/// 当前音量级别(共享状态,用于快速查询)
volume_level: Arc<AtomicU32>,
/// 录音开始时间(共享状态)
start_time: Arc<Mutex<Option<Instant>>>,
}
impl RecordingService {
/// 创建新的录音服务
pub fn new() -> Self {
Self {
command_tx: None,
response_rx: None,
thread_handle: None,
is_recording: Arc::new(AtomicBool::new(false)),
volume_level: Arc::new(AtomicU32::new(0)),
start_time: Arc::new(Mutex::new(None)),
}
}
/// 确保录音线程已启动
fn ensure_thread_started(&mut self) {
if self.command_tx.is_some() {
return;
}
let (cmd_tx, cmd_rx) = mpsc::channel::<RecordingCommand>();
let (resp_tx, resp_rx) = mpsc::channel::<RecordingResponse>();
let is_recording = Arc::clone(&self.is_recording);
let volume_level = Arc::clone(&self.volume_level);
let start_time = Arc::clone(&self.start_time);
let handle = thread::spawn(move || {
recording_thread_main(cmd_rx, resp_tx, is_recording, volume_level, start_time);
});
self.command_tx = Some(cmd_tx);
self.response_rx = Some(resp_rx);
self.thread_handle = Some(handle);
tracing::info!("[录音服务] 录音线程已启动");
}
/// 开始录音(可选指定设备 ID)
pub fn start(&mut self, device_id: Option<String>) -> Result<(), String> {
self.ensure_thread_started();
let tx = self.command_tx.as_ref().ok_or("录音线程未启动")?;
let rx = self.response_rx.as_ref().ok_or("录音线程未启动")?;
tx.send(RecordingCommand::Start(device_id))
.map_err(|e| format!("发送命令失败: {e}"))?;
match rx.recv() {
Ok(RecordingResponse::Ok) => {
tracing::info!("[录音服务] 开始录音");
Ok(())
}
Ok(RecordingResponse::Error(e)) => Err(e),
Ok(_) => Err("意外的响应".to_string()),
Err(e) => Err(format!("接收响应失败: {e}")),
}
}
/// 停止录音并返回音频数据
pub fn stop(&mut self) -> Result<AudioData, String> {
let tx = self.command_tx.as_ref().ok_or("录音线程未启动")?;
let rx = self.response_rx.as_ref().ok_or("录音线程未启动")?;
tx.send(RecordingCommand::Stop)
.map_err(|e| format!("发送命令失败: {e}"))?;
match rx.recv() {
Ok(RecordingResponse::AudioData(audio)) => {
tracing::info!("[录音服务] 停止录音,时长: {:.2}s", audio.duration_secs);
Ok(audio)
}
Ok(RecordingResponse::Error(e)) => Err(e),
Ok(_) => Err("意外的响应".to_string()),
Err(e) => Err(format!("接收响应失败: {e}")),
}
}
/// 取消录音
pub fn cancel(&mut self) {
if let Some(tx) = &self.command_tx {
let _ = tx.send(RecordingCommand::Cancel);
// 使用 try_recv 避免阻塞,或者设置超时
if let Some(rx) = &self.response_rx {
// 尝试接收响应,但不阻塞太久
use std::time::Duration;
match rx.recv_timeout(Duration::from_millis(500)) {
Ok(_) => tracing::info!("[录音服务] 取消录音成功"),
Err(std::sync::mpsc::RecvTimeoutError::Timeout) => {
tracing::warn!("[录音服务] 取消录音超时,强制继续");
}
Err(e) => {
tracing::warn!("[录音服务] 取消录音响应错误: {}", e);
}
}
}
}
// 无论如何都重置状态
self.is_recording.store(false, Ordering::SeqCst);
self.volume_level.store(0, Ordering::SeqCst);
*self.start_time.lock() = None;
}
/// 获取当前音量级别(0-100)
pub fn get_volume(&self) -> u32 {
self.volume_level.load(Ordering::SeqCst)
}
/// 获取录音时长(秒)
pub fn get_duration(&self) -> f32 {
self.start_time
.lock()
.map(|t| t.elapsed().as_secs_f32())
.unwrap_or(0.0)
}
/// 是否正在录音
pub fn is_recording(&self) -> bool {
self.is_recording.load(Ordering::SeqCst)
}
/// 关闭录音服务
pub fn shutdown(&mut self) {
if let Some(tx) = self.command_tx.take() {
let _ = tx.send(RecordingCommand::Shutdown);
}
if let Some(handle) = self.thread_handle.take() {
let _ = handle.join();
}
self.response_rx = None;
tracing::info!("[录音服务] 已关闭");
}
}
impl Default for RecordingService {
fn default() -> Self {
Self::new()
}
}
impl Drop for RecordingService {
fn drop(&mut self) {
self.shutdown();
}
}
/// 录音线程主函数
///
/// 在独立线程中运行,拥有 cpal::Stream
fn recording_thread_main(
cmd_rx: Receiver<RecordingCommand>,
resp_tx: Sender<RecordingResponse>,
is_recording: Arc<AtomicBool>,
volume_level: Arc<AtomicU32>,
start_time: Arc<Mutex<Option<Instant>>>,
) {
use cpal::traits::{DeviceTrait, HostTrait, StreamTrait};
// 录音数据缓冲区
let samples: Arc<Mutex<Vec<i16>>> = Arc::new(Mutex::new(Vec::new()));
// 当前活跃的音频流
let mut active_stream: Option<cpal::Stream> = None;
// 实际使用的采样率和声道数
let mut actual_sample_rate: u32 = 16000;
#[allow(unused_assignments)]
let mut actual_channels: u16 = 1;
tracing::debug!("[录音线程] 开始运行");
loop {
match cmd_rx.recv() {
Ok(RecordingCommand::Start(device_id)) => {
// 如果已在录音,返回错误
if is_recording.load(Ordering::SeqCst) {
let _ = resp_tx.send(RecordingResponse::Error("已在录音中".to_string()));
continue;
}
// 清空缓冲区
samples.lock().clear();
// 获取输入设备
let host = cpal::default_host();
let device = if let Some(ref id) = device_id {
// 查找指定设备
host.input_devices()
.ok()
.and_then(|mut devices| {
devices.find(|d| d.name().ok().as_ref() == Some(id))
})
.or_else(|| {
tracing::warn!("[录音线程] 未找到指定设备 {},使用默认设备", id);
host.default_input_device()
})
} else {
host.default_input_device()
};
let device = match device {
Some(d) => d,
None => {
let _ =
resp_tx.send(RecordingResponse::Error("未找到麦克风设备".to_string()));
continue;
}
};
tracing::info!("[录音线程] 使用麦克风: {:?}", device.name());
// 获取设备支持的配置
let supported_config = match device.default_input_config() {
Ok(c) => c,
Err(e) => {
let _ = resp_tx
.send(RecordingResponse::Error(format!("获取音频配置失败: {e}")));
continue;
}
};
tracing::info!(
"[录音线程] 设备支持配置: 采样率={}, 声道={}",
supported_config.sample_rate().0,
supported_config.channels()
);
// 使用设备默认配置
actual_sample_rate = supported_config.sample_rate().0;
actual_channels = supported_config.channels();
let config = cpal::StreamConfig {
channels: actual_channels,
sample_rate: supported_config.sample_rate(),
buffer_size: cpal::BufferSize::Default,
};
// 创建共享状态的克隆
let samples_clone = Arc::clone(&samples);
let volume_clone = Arc::clone(&volume_level);
let is_rec_clone = Arc::clone(&is_recording);
let channels = actual_channels;
// 回调计数器(用于调试)
let callback_count = Arc::new(AtomicU32::new(0));
let callback_count_clone = Arc::clone(&callback_count);
// 创建输入流
let stream = match device.build_input_stream(
&config,
move |data: &[f32], _: &cpal::InputCallbackInfo| {
if !is_rec_clone.load(Ordering::SeqCst) {
return;
}
// 增加回调计数
let count = callback_count_clone.fetch_add(1, Ordering::SeqCst);
if count == 0 {
tracing::info!("[录音线程] 首次收到音频数据,数据长度: {}", data.len());
} else if count % 100 == 0 {
tracing::debug!("[录音线程] 已收到 {} 次音频回调", count);
}
// 计算音量级别(使用 RMS 均方根,更准确反映音量)
let sum_sq: f32 = data.iter().map(|s| s * s).sum();
let rms = (sum_sq / data.len() as f32).sqrt();
// 将 RMS 值映射到 0-100 范围
// 静音时 RMS 约 0.001-0.01,说话时约 0.02-0.1
// 使用更高的系数来提高灵敏度
let level = ((rms * 1500.0).min(100.0)) as u32;
// 每 50 次回调打印一次音量(用于调试)
if count % 50 == 0 {
tracing::debug!("[录音线程] RMS: {:.6}, 音量: {}%", rms, level);
}
volume_clone.store(level, Ordering::SeqCst);
// 如果是多声道,转换为单声道
let mono_data: Vec<f32> = if channels > 1 {
data.chunks(channels as usize)
.map(|chunk| chunk.iter().sum::<f32>() / channels as f32)
.collect()
} else {
data.to_vec()
};
// 转换为 i16 并存储
let i16_samples: Vec<i16> = mono_data
.iter()
.map(|&s| (s.clamp(-1.0, 1.0) * i16::MAX as f32) as i16)
.collect();
samples_clone.lock().extend(i16_samples);
},
|err| {
tracing::error!("[录音线程] 录音流错误: {}", err);
},
None,
) {
Ok(s) => s,
Err(e) => {
let _ =
resp_tx.send(RecordingResponse::Error(format!("创建音频流失败: {e}")));
continue;
}
};
// 开始播放(录音)
if let Err(e) = stream.play() {
let _ = resp_tx.send(RecordingResponse::Error(format!("启动录音失败: {e}")));
continue;
}
tracing::info!("[录音线程] stream.play() 成功,等待音频数据...");
// 保存流和状态
active_stream = Some(stream);
is_recording.store(true, Ordering::SeqCst);
*start_time.lock() = Some(Instant::now());
let _ = resp_tx.send(RecordingResponse::Ok);
tracing::info!(
"[录音线程] 开始录音,采样率: {}, 声道: {}",
actual_sample_rate,
actual_channels
);
}
Ok(RecordingCommand::Stop) => {
if !is_recording.load(Ordering::SeqCst) {
let _ = resp_tx.send(RecordingResponse::Error("未在录音中".to_string()));
continue;
}
// 停止录音
is_recording.store(false, Ordering::SeqCst);
// 停止并释放流
if let Some(stream) = active_stream.take() {
drop(stream);
}
// 获取录音数据(已转换为单声道)
let audio_samples = samples.lock().clone();
let audio = AudioData::new(audio_samples, actual_sample_rate, 1);
// 重置开始时间
*start_time.lock() = None;
volume_level.store(0, Ordering::SeqCst);
// 检查录音时长
if !audio.is_valid() {
let _ = resp_tx.send(RecordingResponse::Error(
"录音时间过短(需要至少 0.5 秒)".to_string(),
));
continue;
}
let _ = resp_tx.send(RecordingResponse::AudioData(audio));
tracing::info!("[录音线程] 停止录音");
}
Ok(RecordingCommand::Cancel) => {
// 停止录音
is_recording.store(false, Ordering::SeqCst);
// 停止并释放流
if let Some(stream) = active_stream.take() {
drop(stream);
}
// 清空缓冲区
samples.lock().clear();
// 重置状态
*start_time.lock() = None;
volume_level.store(0, Ordering::SeqCst);
let _ = resp_tx.send(RecordingResponse::Ok);
tracing::info!("[录音线程] 取消录音");
}
Ok(RecordingCommand::Shutdown) => {
// 清理资源
is_recording.store(false, Ordering::SeqCst);
if let Some(stream) = active_stream.take() {
drop(stream);
}
tracing::info!("[录音线程] 收到关闭命令,退出");
break;
}
Err(_) => {
// channel 已关闭,退出线程
tracing::info!("[录音线程] channel 已关闭,退出");
break;
}
}
}
}
+18
View File
@@ -34,6 +34,24 @@ impl AudioData {
self.duration_secs >= 0.5
}
/// 从 PCM16 LE 字节创建音频数据
pub fn from_pcm16le_bytes(bytes: &[u8], sample_rate: u32, channels: u16) -> Self {
let samples = bytes
.chunks_exact(2)
.map(|chunk| i16::from_le_bytes([chunk[0], chunk[1]]))
.collect();
Self::new(samples, sample_rate, channels)
}
/// 转换为 PCM16 LE 字节
pub fn to_pcm16le_bytes(&self) -> Vec<u8> {
self.samples
.iter()
.flat_map(|sample| sample.to_le_bytes())
.collect()
}
/// 转换为 WAV 格式字节
pub fn to_wav_bytes(&self) -> Vec<u8> {
let mut cursor = std::io::Cursor::new(Vec::new());
+8 -10
View File
@@ -1,20 +1,18 @@
//! AI Agent 集成模块
//!
//! 基于 aster-rust 框架实现 Agent 功能
//!
//! ## 架构设计
//! - aster_state - Aster Agent 状态管理
//! - aster_agent - Aster Agent 包装器
//! - event_converter - Aster 事件转换器
//! - credential_bridge - 凭证池桥接(连接 ProxyCast 凭证池与 Aster Provider)
//! - subagent_scheduler - SubAgent 调度器集成
//! 纯逻辑部分已迁移到 proxycast-agent crate,
//! 本模块保留深耦合部分(依赖 database, services, AppHandle)。
pub mod aster_agent;
pub mod aster_state;
pub mod credential_bridge;
pub mod event_converter;
pub mod subagent_scheduler;
// 从 proxycast-agent crate re-export
pub use proxycast_agent::event_converter;
pub use proxycast_agent::mcp_bridge;
pub use proxycast_agent::prompt;
// types 已迁移到 proxycast-core
pub use proxycast_core::agent::types;
@@ -23,7 +21,7 @@ pub use aster_state::AsterAgentState;
pub use credential_bridge::{
create_aster_provider, AsterProviderConfig, CredentialBridge, CredentialBridgeError,
};
pub use event_converter::{convert_agent_event, TauriAgentEvent};
pub use proxycast_agent::{convert_agent_event, convert_to_tauri_message, TauriAgentEvent};
pub use subagent_scheduler::{
ProxyCastScheduler, ProxyCastSubAgentExecutor, SubAgentProgressEvent,
};
-65
View File
@@ -1,65 +0,0 @@
# System Prompt 模块
<!-- 一旦我所属的文件夹有所变化,请更新我 -->
## 架构说明
为 Aster Agent 提供 System Prompt 配置,参考 claude-code-open 的设计。
### 设计决策
- **模块化模板**:将 System Prompt 拆分为多个独立模板,便于维护和定制
- **自动注入**:Agent 初始化时自动注入 System Prompt
- **环境感知**:自动添加当前日期、操作系统、工作目录等环境信息
## 文件索引
| 文件 | 说明 |
|------|------|
| `mod.rs` | 模块入口,导出公共类型 |
| `templates.rs` | 提示词模板定义 |
| `builder.rs` | 提示词构建器 |
## 模板内容
| 模板 | 说明 |
|------|------|
| `CORE_IDENTITY` | Agent 身份描述 |
| `TOOL_GUIDELINES` | 工具使用策略(read/write/edit/glob/grep/bash) |
| `CODING_GUIDELINES` | 代码编写指南 |
| `TASK_MANAGEMENT` | 任务管理(TodoWrite 使用) |
| `GIT_GUIDELINES` | Git 操作安全规则 |
| `OUTPUT_STYLE` | 输出风格指南 |
## 使用方式
### 基本使用
```rust
use crate::agent::prompt::SystemPromptBuilder;
let prompt = SystemPromptBuilder::new()
.working_dir("/path/to/project")
.build();
```
### 添加自定义指令
```rust
let prompt = SystemPromptBuilder::new()
.working_dir("/path/to/project")
.custom_instructions("额外的项目特定指令")
.build();
```
### 在 AsterAgentState 中的集成
System Prompt 在 `init_agent()` 时自动注入:
```rust
// 初始化时自动注入 System Prompt
state.init_agent().await?;
// 也可以动态添加自定义指令
state.add_custom_instructions("额外指令").await?;
```
+6 -3
View File
@@ -210,14 +210,17 @@ pub fn run() {
tracing::info!("[启动] GlobalConfigManager 事件发射器已设置");
}
// 设置 MCP Manager 的 AppHandle(用于发送 mcp:* 事件)
// 设置 MCP Manager 的事件发射器(用于发送 mcp:* 事件)
if let Some(mcp_manager) = app.try_state::<crate::mcp::McpManagerState>() {
let app_handle = app.handle().clone();
let emitter = proxycast_core::DynEmitter::new(
crate::app::TauriEventEmitter(app_handle),
);
tauri::async_runtime::block_on(async {
let mut manager = mcp_manager.lock().await;
manager.set_app_handle(app_handle);
manager.set_emitter(emitter);
});
tracing::info!("[启动] MCP Manager AppHandle 已设置");
tracing::info!("[启动] MCP Manager 事件发射器已设置");
}
// 初始化截图对话模块
+4 -236
View File
@@ -1,241 +1,9 @@
//! 请求处理器模块
//! 请求处理器模块(重导出层)
//!
//! 提供统一的请求处理管道,集成路由、容错、监控、插件等功能模块。
//!
//! # 架构
//!
//! 请求处理流程:
//! 1. 认证 (AuthStep)
//! 2. 参数注入 (InjectionStep)
//! 3. 路由解析 (RoutingStep)
//! 4. 插件前置钩子 (PluginPreStep)
//! 5. Provider 调用 (ProviderStep) - 包含重试和故障转移
//! 6. 插件后置钩子 (PluginPostStep)
//! 7. 统计记录 (TelemetryStep)
//! 核心逻辑已迁移到 `proxycast-processor` crate。
//! 本模块保留向后兼容路径和本地测试入口。
// context 和 error 已迁移到 proxycast-core
pub use proxycast_core::processor::RequestContext;
mod steps;
use crate::injection::Injector;
use crate::plugin::PluginManager;
use crate::resilience::{Failover, Retrier, TimeoutController};
use crate::router::{ModelMapper, Router};
use crate::services::provider_pool_service::ProviderPoolService;
use crate::telemetry::{StatsAggregator, TokenTracker};
use parking_lot::RwLock as ParkingLotRwLock;
use std::sync::Arc;
use tokio::sync::RwLock;
/// 统一的请求处理器
///
/// 集成所有功能模块,提供完整的请求处理管道
pub struct RequestProcessor {
/// 路由器
pub router: Arc<RwLock<Router>>,
/// 模型映射器
pub mapper: Arc<RwLock<ModelMapper>>,
/// 参数注入器
pub injector: Arc<RwLock<Injector>>,
/// 重试器
pub retrier: Arc<Retrier>,
/// 故障转移器
pub failover: Arc<Failover>,
/// 超时控制器
pub timeout: Arc<TimeoutController>,
/// 插件管理器
pub plugins: Arc<PluginManager>,
/// 统计聚合器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
pub stats: Arc<ParkingLotRwLock<StatsAggregator>>,
/// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
pub tokens: Arc<ParkingLotRwLock<TokenTracker>>,
/// 凭证池服务
pub pool_service: Arc<ProviderPoolService>,
/// 热重载协调锁(避免配置更新期间请求读取不一致的配置)
pub reload_lock: Arc<RwLock<()>>,
}
impl RequestProcessor {
/// 创建新的请求处理器
pub fn new(
router: Arc<RwLock<Router>>,
mapper: Arc<RwLock<ModelMapper>>,
injector: Arc<RwLock<Injector>>,
retrier: Arc<Retrier>,
failover: Arc<Failover>,
timeout: Arc<TimeoutController>,
plugins: Arc<PluginManager>,
stats: Arc<ParkingLotRwLock<StatsAggregator>>,
tokens: Arc<ParkingLotRwLock<TokenTracker>>,
pool_service: Arc<ProviderPoolService>,
) -> Self {
Self {
router,
mapper,
injector,
retrier,
failover,
timeout,
plugins,
stats,
tokens,
pool_service,
reload_lock: Arc::new(RwLock::new(())),
}
}
/// 使用默认配置创建请求处理器
pub fn with_defaults(pool_service: Arc<ProviderPoolService>) -> Self {
Self {
router: Arc::new(RwLock::new(Self::create_router_with_defaults())),
mapper: Arc::new(RwLock::new(ModelMapper::new())),
injector: Arc::new(RwLock::new(Injector::new())),
retrier: Arc::new(Retrier::with_defaults()),
failover: Arc::new(Failover::with_defaults()),
timeout: Arc::new(TimeoutController::with_defaults()),
plugins: Arc::new(PluginManager::with_defaults()),
stats: Arc::new(ParkingLotRwLock::new(StatsAggregator::with_defaults())),
tokens: Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults())),
pool_service,
reload_lock: Arc::new(RwLock::new(())),
}
}
/// 创建带默认路由规则的路由器
///
/// 注意:不再添加硬编码的路由规则,让用户设置的默认 Provider 生效
/// 用户可以通过 UI 或配置文件自定义路由规则
fn create_router_with_defaults() -> Router {
// 创建空的路由器,默认 Provider 会在启动时从配置中设置
// 不要硬编码任何 Provider,避免与用户配置冲突
let router = Router::new_empty();
tracing::info!("[ROUTER] 初始化空路由器,等待从配置加载默认 Provider");
router
}
/// 使用共享的统计和 Token 追踪器创建请求处理器
///
/// 这允许 RequestProcessor 与 TelemetryState 共享同一个 StatsAggregator 和 TokenTracker,
/// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。
pub fn with_shared_telemetry(
pool_service: Arc<ProviderPoolService>,
stats: Arc<ParkingLotRwLock<StatsAggregator>>,
tokens: Arc<ParkingLotRwLock<TokenTracker>>,
) -> Self {
Self {
router: Arc::new(RwLock::new(Self::create_router_with_defaults())),
mapper: Arc::new(RwLock::new(ModelMapper::new())),
injector: Arc::new(RwLock::new(Injector::new())),
retrier: Arc::new(Retrier::with_defaults()),
failover: Arc::new(Failover::with_defaults()),
timeout: Arc::new(TimeoutController::with_defaults()),
plugins: Arc::new(PluginManager::with_defaults()),
stats,
tokens,
pool_service,
reload_lock: Arc::new(RwLock::new(())),
}
}
/// 解析模型别名
///
/// 使用 ModelMapper 将模型别名解析为实际模型名称
///
/// # Arguments
/// * `model` - 原始模型名称(可能是别名)
///
/// # Returns
/// 解析后的实际模型名称
pub async fn resolve_model(&self, model: &str) -> String {
let mapper = self.mapper.read().await;
mapper.resolve(model)
}
/// 解析模型别名并更新请求上下文
///
/// # Arguments
/// * `ctx` - 请求上下文
///
/// # Returns
/// 解析后的模型名称
pub async fn resolve_model_for_context(&self, ctx: &mut RequestContext) -> String {
let resolved = self.resolve_model(&ctx.original_model).await;
ctx.set_resolved_model(resolved.clone());
tracing::debug!(
"[MAPPER] request_id={} original_model={} resolved_model={}",
ctx.request_id,
ctx.original_model,
resolved
);
resolved
}
/// 根据模型选择 Provider
///
/// 使用 Router 根据路由规则选择合适的 Provider
///
/// # Arguments
/// * `model` - 模型名称(应该是解析后的实际模型名)
///
/// # Returns
/// 选择的 Provider 类型(如果设置了)和是否使用默认 Provider
pub async fn route_model(&self, model: &str) -> (Option<crate::ProviderType>, bool) {
let router = self.router.read().await;
let result = router.route(model);
(result.provider, result.is_default)
}
/// 根据模型选择 Provider 并更新请求上下文
///
/// # Arguments
/// * `ctx` - 请求上下文
///
/// # Returns
/// 选择的 Provider 类型,如果未设置默认 Provider 则返回 None
pub async fn route_for_context(&self, ctx: &mut RequestContext) -> Option<crate::ProviderType> {
let (provider, is_default) = self.route_model(&ctx.resolved_model).await;
if let Some(p) = provider {
ctx.set_provider(p);
tracing::info!(
"[ROUTE] request_id={} model={} provider={} is_default={}",
ctx.request_id,
ctx.resolved_model,
p,
is_default
);
} else {
tracing::warn!(
"[ROUTE] request_id={} model={} 未设置默认 Provider",
ctx.request_id,
ctx.resolved_model
);
}
provider
}
/// 执行完整的路由解析流程
///
/// 包括模型别名解析和 Provider 选择
///
/// # Arguments
/// * `ctx` - 请求上下文
///
/// # Returns
/// 选择的 Provider 类型,如果未设置默认 Provider 则返回 None
pub async fn resolve_and_route(&self, ctx: &mut RequestContext) -> Option<crate::ProviderType> {
// 1. 解析模型别名
self.resolve_model_for_context(ctx).await;
// 2. 根据解析后的模型选择 Provider
self.route_for_context(ctx).await
}
}
pub use proxycast_processor::*;
#[cfg(test)]
mod tests;
+1
View File
@@ -3,6 +3,7 @@
use super::*;
use crate::services::provider_pool_service::ProviderPoolService;
use crate::ProviderType;
use std::sync::Arc;
#[test]
fn test_request_processor_new() {
+6 -175
View File
@@ -1,150 +1,25 @@
//! Tauri 执行回调实现
//!
//! 实现 aster-rust 的 ExecutionCallback trait,通过 Tauri 事件系统向前端发送进度更新。
//!
//! ## 事件类型
//! - `skill:step_start`: 步骤开始
//! - `skill:step_complete`: 步骤完成
//! - `skill:step_error`: 步骤错误
//! - `skill:complete`: 执行完成
//!
//! ## 使用示例
//! ```ignore
//! let callback = TauriExecutionCallback::new(app_handle, "exec-123".to_string());
//! callback.on_step_start("step-1", "数据处理", 1, 3);
//! ```
//! 通过 Tauri 事件系统向前端发送 Skill 执行进度更新。
use serde::Serialize;
use std::sync::atomic::{AtomicUsize, Ordering};
use tauri::{AppHandle, Emitter};
/// 步骤开始事件 Payload
#[derive(Debug, Clone, Serialize)]
pub struct StepStartPayload {
/// 执行 ID
pub execution_id: String,
/// 步骤 ID
pub step_id: String,
/// 步骤名称
pub step_name: String,
/// 当前步骤序号(从 1 开始)
pub current_step: usize,
/// 总步骤数
pub total_steps: usize,
}
/// 步骤完成事件 Payload
#[derive(Debug, Clone, Serialize)]
pub struct StepCompletePayload {
/// 执行 ID
pub execution_id: String,
/// 步骤 ID
pub step_id: String,
/// 步骤输出
pub output: String,
}
/// 步骤错误事件 Payload
#[derive(Debug, Clone, Serialize)]
pub struct StepErrorPayload {
/// 执行 ID
pub execution_id: String,
/// 步骤 ID
pub step_id: String,
/// 错误信息
pub error: String,
/// 是否会重试
pub will_retry: bool,
}
/// 执行完成事件 Payload
#[derive(Debug, Clone, Serialize)]
pub struct ExecutionCompletePayload {
/// 执行 ID
pub execution_id: String,
/// 是否成功
pub success: bool,
/// 最终输出(成功时)
pub output: Option<String>,
/// 错误信息(失败时)
pub error: Option<String>,
}
/// Tauri 事件名称常量
pub mod events {
/// 步骤开始事件
pub const STEP_START: &str = "skill:step_start";
/// 步骤完成事件
pub const STEP_COMPLETE: &str = "skill:step_complete";
/// 步骤错误事件
pub const STEP_ERROR: &str = "skill:step_error";
/// 执行完成事件
pub const COMPLETE: &str = "skill:complete";
}
/// ExecutionCallback Trait
///
/// 定义 Skill 执行过程中的回调接口。
/// 应用层需要实现此 trait 以接收执行进度更新。
pub trait ExecutionCallback: Send + Sync {
/// 步骤开始回调
///
/// # 参数
/// - `step_id`: 步骤 ID
/// - `step_name`: 步骤名称
/// - `current_step`: 当前步骤序号(从 1 开始)
/// - `total_steps`: 总步骤数
fn on_step_start(
&self,
step_id: &str,
step_name: &str,
current_step: usize,
total_steps: usize,
);
/// 步骤完成回调
///
/// # 参数
/// - `step_id`: 步骤 ID
/// - `output`: 步骤输出
fn on_step_complete(&self, step_id: &str, output: &str);
/// 步骤错误回调
///
/// # 参数
/// - `step_id`: 步骤 ID
/// - `error`: 错误信息
/// - `will_retry`: 是否会重试
fn on_step_error(&self, step_id: &str, error: &str, will_retry: bool);
/// 执行完成回调
///
/// # 参数
/// - `success`: 是否成功
/// - `final_output`: 最终输出(成功时)
/// - `error`: 错误信息(失败时)
fn on_complete(&self, success: bool, final_output: Option<&str>, error: Option<&str>);
}
use proxycast_skills::{
events, ExecutionCallback, ExecutionCompletePayload, StepCompletePayload, StepErrorPayload,
StepStartPayload,
};
/// Tauri 执行回调
///
/// 通过 Tauri 事件系统向前端发送 Skill 执行进度更新。
/// 实现 aster-rust 定义的 ExecutionCallback trait。
pub struct TauriExecutionCallback {
/// Tauri AppHandle
app_handle: AppHandle,
/// 执行 ID(用于区分多个并发执行)
execution_id: String,
/// 当前步骤计数器(用于跟踪步骤序号)
current_step: AtomicUsize,
}
impl TauriExecutionCallback {
/// 创建新的 TauriExecutionCallback 实例
///
/// # Arguments
/// * `app_handle` - Tauri AppHandle
/// * `execution_id` - 执行 ID,用于区分多个并发执行
pub fn new(app_handle: AppHandle, execution_id: String) -> Self {
Self {
app_handle,
@@ -153,33 +28,16 @@ impl TauriExecutionCallback {
}
}
/// 获取执行 ID
pub fn execution_id(&self) -> &str {
&self.execution_id
}
/// 获取当前步骤序号
pub fn current_step(&self) -> usize {
self.current_step.load(Ordering::SeqCst)
}
}
/// ExecutionCallback trait 实现
///
/// 通过 Tauri 事件系统向前端发送进度更新。
///
/// # Requirements
/// - 2.2: on_step_start 发送 "skill:step_start" 事件
/// - 2.3: on_step_complete 发送 "skill:step_complete" 事件
/// - 2.4: on_step_error 发送 "skill:step_error" 事件
/// - 2.5: on_complete 发送 "skill:complete" 事件
impl ExecutionCallback for TauriExecutionCallback {
/// 步骤开始回调
///
/// 发送 "skill:step_start" Tauri 事件到前端。
///
/// # Requirements
/// - 2.2: WHEN on_step_start is called, emit a "skill:step_start" Tauri event
fn on_step_start(
&self,
step_id: &str,
@@ -187,7 +45,6 @@ impl ExecutionCallback for TauriExecutionCallback {
current_step: usize,
total_steps: usize,
) {
// 更新当前步骤计数器
self.current_step.store(current_step, Ordering::SeqCst);
let payload = StepStartPayload {
@@ -216,12 +73,6 @@ impl ExecutionCallback for TauriExecutionCallback {
}
}
/// 步骤完成回调
///
/// 发送 "skill:step_complete" Tauri 事件到前端。
///
/// # Requirements
/// - 2.3: WHEN on_step_complete is called, emit a "skill:step_complete" Tauri event
fn on_step_complete(&self, step_id: &str, output: &str) {
let payload = StepCompletePayload {
execution_id: self.execution_id.clone(),
@@ -245,12 +96,6 @@ impl ExecutionCallback for TauriExecutionCallback {
}
}
/// 步骤错误回调
///
/// 发送 "skill:step_error" Tauri 事件到前端。
///
/// # Requirements
/// - 2.4: WHEN on_step_error is called, emit a "skill:step_error" Tauri event
fn on_step_error(&self, step_id: &str, error: &str, will_retry: bool) {
let payload = StepErrorPayload {
execution_id: self.execution_id.clone(),
@@ -261,10 +106,7 @@ impl ExecutionCallback for TauriExecutionCallback {
tracing::warn!(
"[TauriExecutionCallback] 步骤错误: execution_id={}, step_id={}, error={}, will_retry={}",
self.execution_id,
step_id,
error,
will_retry
self.execution_id, step_id, error, will_retry
);
if let Err(e) = self.app_handle.emit(events::STEP_ERROR, &payload) {
@@ -276,12 +118,6 @@ impl ExecutionCallback for TauriExecutionCallback {
}
}
/// 执行完成回调
///
/// 发送 "skill:complete" Tauri 事件到前端。
///
/// # Requirements
/// - 2.5: WHEN on_complete is called, emit a "skill:complete" Tauri event
fn on_complete(&self, success: bool, final_output: Option<&str>, error: Option<&str>) {
let payload = ExecutionCompletePayload {
execution_id: self.execution_id.clone(),
@@ -313,8 +149,3 @@ impl ExecutionCallback for TauriExecutionCallback {
}
}
}
#[cfg(test)]
mod tests {
// TODO: 在 Task 1.5 中添加属性测试
}
+3 -59
View File
@@ -1,20 +1,11 @@
//! ProxyCast LLM Provider 实现
//!
//! 实现 aster-rust 的 LlmProvider trait,使用 ProviderPoolService 选择凭证并调用 LLM API。
//!
//! ## 功能
//! - 通过 ProviderPoolService 选择可用凭证
//! - 支持指定 provider 类型和 model 参数
//! - 智能降级到 API Key Provider
//!
//! ## 依赖
//! - `ProviderPoolService`: 凭证池管理
//! - `ApiKeyProviderService`: API Key 服务(降级使用)
//! 使用 ProviderPoolService 选择凭证并调用 LLM API。
//! trait 定义(LlmProvider, SkillError)已迁移到 proxycast-skills crate。
use std::sync::Arc;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use crate::database::DbConnection;
use crate::models::anthropic::AnthropicMessagesRequest;
@@ -25,54 +16,7 @@ use crate::providers::{ClaudeCustomProvider, KiroProvider, OpenAICustomProvider}
use crate::services::api_key_provider_service::ApiKeyProviderService;
use crate::services::provider_pool_service::ProviderPoolService;
/// Skill 执行错误类型
///
/// 用于 LlmProvider trait 的错误返回
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum SkillError {
/// Provider 错误(凭证不可用、API 调用失败等)
ProviderError(String),
/// 执行错误(Skill 执行过程中的错误)
ExecutionError(String),
/// 配置错误
ConfigError(String),
}
impl std::fmt::Display for SkillError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SkillError::ProviderError(msg) => write!(f, "Provider error: {}", msg),
SkillError::ExecutionError(msg) => write!(f, "Execution error: {}", msg),
SkillError::ConfigError(msg) => write!(f, "Config error: {}", msg),
}
}
}
impl std::error::Error for SkillError {}
/// LLM Provider Trait
///
/// 定义 Skill 执行引擎调用 LLM 的接口。
/// 应用层需要实现此 trait 以提供 LLM 调用能力。
#[async_trait]
pub trait LlmProvider: Send + Sync {
/// 调用 LLM 进行对话
///
/// # 参数
/// - `system_prompt`: 系统提示词
/// - `user_message`: 用户消息
/// - `model`: 可选的模型名称
///
/// # 返回
/// - `Ok(String)`: LLM 的响应文本
/// - `Err(SkillError)`: 调用失败时的错误
async fn chat(
&self,
system_prompt: &str,
user_message: &str,
model: Option<&str>,
) -> Result<String, SkillError>;
}
use proxycast_skills::{LlmProvider, SkillError};
/// ProxyCast LLM Provider
///
+13 -24
View File
@@ -1,32 +1,21 @@
//! Skills 集成模块
//!
//! 本模块实现 aster-rust Skills 系统与 ProxyCast 的集成。
//!
//! ## 模块结构
//! - `llm_provider`: ProxyCastLlmProvider 实现,使用 ProviderPoolService 调用 LLM
//! - `execution_callback`: TauriExecutionCallback 实现,通过 Tauri 事件发送进度
//!
//! ## 使用示例
//! ```ignore
//! use proxycast::skills::{ProxyCastLlmProvider, TauriExecutionCallback};
//!
//! let provider = ProxyCastLlmProvider::new(pool_service, api_key_service, db);
//! let callback = TauriExecutionCallback::new(app_handle, execution_id);
//! ```
//! trait 定义和纯逻辑已迁移到 proxycast-skills crate,
//! 本模块保留 Tauri 相关的实现。
mod execution_callback;
mod llm_provider;
mod skill_loader;
pub use execution_callback::{
events, ExecutionCallback, ExecutionCompletePayload, StepCompletePayload, StepErrorPayload,
StepStartPayload, TauriExecutionCallback,
// 从 proxycast-skills crate re-export
pub use proxycast_skills::{
events, ExecutionCallback, ExecutionCompletePayload, LlmProvider, SkillError,
StepCompletePayload, StepErrorPayload, StepStartPayload,
};
pub use llm_provider::{LlmProvider, ProxyCastLlmProvider, SkillError};
pub(crate) use skill_loader::{
find_skill_by_name, get_proxycast_skills_dir, load_skills_from_directory,
};
#[cfg(test)]
pub(crate) use skill_loader::{
load_skill_from_file, parse_allowed_tools, parse_boolean, parse_skill_frontmatter,
pub use proxycast_skills::{
find_skill_by_name, get_proxycast_skills_dir, load_skill_from_file, load_skills_from_directory,
parse_allowed_tools, parse_boolean, parse_skill_frontmatter,
};
// Tauri 实现(留在主 crate)
pub use execution_callback::TauriExecutionCallback;
pub use llm_provider::ProxyCastLlmProvider;
+2 -2
View File
@@ -42,10 +42,10 @@
```
voice/
├── asr_service.rs ──→ voice-core (WhisperTranscriber, XunfeiClient)
├── asr_service.rs ──→ voice-core (WhisperTranscriber, AsrClient)
├── output_service.rs ──→ voice-core (OutputHandler)
├── recording_service.rs ──→ voice-core (threaded_recorder + Tauri State 包装)
├── processor.rs ──→ 本地 API 服务器 (LLM 润色)
├── recording_service.rs ──→ cpal (音频采集)
└── commands.rs ──→ 上述所有服务
```
+28 -192
View File
@@ -29,6 +29,8 @@ use std::path::PathBuf;
#[cfg(feature = "local-whisper")]
use crate::config::WhisperModelSize;
use crate::config::{load_config, AsrCredentialEntry, AsrProviderType};
use voice_core::asr_client::{AsrClient, BaiduClient, OpenAIWhisperClient, XunfeiClient};
use voice_core::types::AudioData;
/// ASR 服务
pub struct AsrService;
@@ -152,18 +154,7 @@ impl AsrService {
let model_path = Self::get_whisper_model_path(&whisper_config.model)?;
// 将 PCM 字节转换为 i16 采样
let samples: Vec<i16> = audio_data
.chunks_exact(2)
.map(|chunk| i16::from_le_bytes([chunk[0], chunk[1]]))
.collect();
// 检查音频数据是否有效
if samples.is_empty() {
return Err("音频数据为空".to_string());
}
// 创建 AudioData
let audio = voice_core::types::AudioData::new(samples, sample_rate, 1);
let audio = Self::build_audio_data(audio_data, sample_rate)?;
// 检查录音时长
if !audio.is_valid() {
@@ -240,80 +231,26 @@ impl AsrService {
}
/// OpenAI Whisper API 识别
///
/// 使用手动构建 multipart/form-data 请求
async fn transcribe_openai(
credential: &AsrCredentialEntry,
audio_data: &[u8],
sample_rate: u32,
) -> Result<String, String> {
let config = credential.openai_config.as_ref().ok_or("OpenAI 配置缺失")?;
let audio = Self::build_audio_data(audio_data, sample_rate)?;
// 构建 WAV 文件
let wav_data = Self::build_wav(audio_data, sample_rate, 1)?;
// 构建 multipart/form-data 请求体
let boundary = format!("----WebKitFormBoundary{}", uuid::Uuid::new_v4().simple());
let mut body = Vec::new();
// 添加 file 字段
body.extend_from_slice(format!("--{boundary}\r\n").as_bytes());
body.extend_from_slice(
b"Content-Disposition: form-data; name=\"file\"; filename=\"audio.wav\"\r\n",
);
body.extend_from_slice(b"Content-Type: audio/wav\r\n\r\n");
body.extend_from_slice(&wav_data);
body.extend_from_slice(b"\r\n");
// 添加 model 字段
body.extend_from_slice(format!("--{boundary}\r\n").as_bytes());
body.extend_from_slice(b"Content-Disposition: form-data; name=\"model\"\r\n\r\n");
body.extend_from_slice(b"whisper-1\r\n");
// 添加 language 字段
body.extend_from_slice(format!("--{boundary}\r\n").as_bytes());
body.extend_from_slice(b"Content-Disposition: form-data; name=\"language\"\r\n\r\n");
body.extend_from_slice(credential.language.as_bytes());
body.extend_from_slice(b"\r\n");
// 结束边界
body.extend_from_slice(format!("--{boundary}--\r\n").as_bytes());
// 构建请求
let base_url = config
.base_url
.as_deref()
.unwrap_or("https://api.openai.com");
let url = format!("{base_url}/v1/audio/transcriptions");
let client = reqwest::Client::new();
let response = client
.post(&url)
.header("Authorization", format!("Bearer {}", config.api_key))
.header(
"Content-Type",
format!("multipart/form-data; boundary={boundary}"),
)
.body(body)
.send()
.await
.map_err(|e| format!("请求失败: {e}"))?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
return Err(format!("OpenAI API 错误: {status} - {body}"));
let mut client = OpenAIWhisperClient::new(config.api_key.clone());
if let Some(base_url) = config.base_url.clone() {
client = client.with_host(base_url);
}
if !credential.language.is_empty() {
client = client.with_language(credential.language.clone());
}
#[derive(serde::Deserialize)]
struct WhisperResponse {
text: String,
}
let result: WhisperResponse = response
.json()
let result = client
.transcribe(&audio)
.await
.map_err(|e| format!("解析响应失败: {e}"))?;
.map_err(|e| format!("OpenAI Whisper 识别失败: {e}"))?;
Ok(result.text)
}
@@ -325,83 +262,15 @@ impl AsrService {
sample_rate: u32,
) -> Result<String, String> {
let config = credential.baidu_config.as_ref().ok_or("百度配置缺失")?;
let audio = Self::build_audio_data(audio_data, sample_rate)?;
// 获取 Access Token
let token_url = format!(
"https://aip.baidubce.com/oauth/2.0/token?grant_type=client_credentials&client_id={}&client_secret={}",
config.api_key, config.secret_key
);
let client = reqwest::Client::new();
let token_resp = client
.post(&token_url)
.send()
let client = BaiduClient::new(config.api_key.clone(), config.secret_key.clone());
let result = client
.transcribe(&audio)
.await
.map_err(|e| format!("获取 Token 失败: {e}"))?;
.map_err(|e| format!("百度识别失败: {e}"))?;
#[derive(serde::Deserialize)]
struct TokenResponse {
access_token: String,
}
let token: TokenResponse = token_resp
.json()
.await
.map_err(|e| format!("解析 Token 失败: {e}"))?;
// 构建 WAV 并 Base64 编码
let wav_data = Self::build_wav(audio_data, sample_rate, 1)?;
let speech = base64::Engine::encode(&base64::engine::general_purpose::STANDARD, &wav_data);
#[derive(serde::Serialize)]
struct AsrRequest {
format: String,
rate: u32,
channel: u16,
cuid: String,
token: String,
speech: String,
len: usize,
}
let request = AsrRequest {
format: "wav".to_string(),
rate: sample_rate,
channel: 1,
cuid: "proxycast".to_string(),
token: token.access_token,
speech,
len: wav_data.len(),
};
let response = client
.post("https://vop.baidu.com/server_api")
.json(&request)
.send()
.await
.map_err(|e| format!("请求失败: {e}"))?;
#[derive(serde::Deserialize)]
struct AsrResponse {
err_no: i32,
err_msg: String,
#[serde(default)]
result: Vec<String>,
}
let result: AsrResponse = response
.json()
.await
.map_err(|e| format!("解析响应失败: {e}"))?;
if result.err_no != 0 {
return Err(format!(
"百度 ASR 错误: {} - {}",
result.err_no, result.err_msg
));
}
Ok(result.result.join(""))
Ok(result.text)
}
/// 讯飞语音识别
@@ -413,15 +282,7 @@ impl AsrService {
sample_rate: u32,
) -> Result<String, String> {
let config = credential.xunfei_config.as_ref().ok_or("讯飞配置缺失")?;
// 将 PCM 字节转换为 i16 采样
let samples: Vec<i16> = audio_data
.chunks_exact(2)
.map(|chunk| i16::from_le_bytes([chunk[0], chunk[1]]))
.collect();
// 创建 AudioData
let audio = voice_core::types::AudioData::new(samples, sample_rate, 1);
let audio = Self::build_audio_data(audio_data, sample_rate)?;
// 创建讯飞客户端
// 讯飞语言代码转换:zh -> zh_cn, en -> en_us
@@ -431,15 +292,13 @@ impl AsrService {
other => other.to_string(),
};
let client = voice_core::asr_client::XunfeiClient::new(
let client = XunfeiClient::new(
config.app_id.clone(),
config.api_key.clone(),
config.api_secret.clone(),
)
.with_language(xunfei_language);
// 调用识别
use voice_core::asr_client::AsrClient;
let result = client
.transcribe(&audio)
.await
@@ -448,36 +307,13 @@ impl AsrService {
Ok(result.text)
}
/// 构建 WAV 文件
fn build_wav(pcm_data: &[u8], sample_rate: u32, channels: u16) -> Result<Vec<u8>, String> {
let bits_per_sample: u16 = 16;
let byte_rate = sample_rate * u32::from(channels) * u32::from(bits_per_sample) / 8;
let block_align = channels * bits_per_sample / 8;
let data_size = pcm_data.len() as u32;
let file_size = 36 + data_size;
/// 将 PCM 字节构造成 voice-core 的 AudioData
fn build_audio_data(audio_data: &[u8], sample_rate: u32) -> Result<AudioData, String> {
let audio = AudioData::from_pcm16le_bytes(audio_data, sample_rate, 1);
if audio.samples.is_empty() {
return Err("音频数据为空".to_string());
}
let mut wav = Vec::with_capacity(44 + pcm_data.len());
// RIFF header
wav.extend_from_slice(b"RIFF");
wav.extend_from_slice(&file_size.to_le_bytes());
wav.extend_from_slice(b"WAVE");
// fmt chunk
wav.extend_from_slice(b"fmt ");
wav.extend_from_slice(&16u32.to_le_bytes()); // chunk size
wav.extend_from_slice(&1u16.to_le_bytes()); // PCM format
wav.extend_from_slice(&channels.to_le_bytes());
wav.extend_from_slice(&sample_rate.to_le_bytes());
wav.extend_from_slice(&byte_rate.to_le_bytes());
wav.extend_from_slice(&block_align.to_le_bytes());
wav.extend_from_slice(&bits_per_sample.to_le_bytes());
// data chunk
wav.extend_from_slice(b"data");
wav.extend_from_slice(&data_size.to_le_bytes());
wav.extend_from_slice(pcm_data);
Ok(wav)
Ok(audio)
}
}
+1 -5
View File
@@ -328,11 +328,7 @@ pub async fn stop_recording(
);
// 将 i16 样本转换为字节(小端序)
let bytes: Vec<u8> = audio
.samples
.iter()
.flat_map(|&s| s.to_le_bytes())
.collect();
let bytes = audio.to_pcm16le_bytes();
Ok(StopRecordingResult {
audio_data: bytes,
+11 -34
View File
@@ -3,43 +3,20 @@
//! 提供模拟键盘输入和剪贴板输出功能
use crate::config::VoiceOutputMode;
use arboard::Clipboard;
use voice_core::{OutputHandler, OutputMode};
/// 输出文字到系统
///
/// 根据配置的输出模式,将文字输出到当前焦点应用
pub fn output_text(text: &str, mode: VoiceOutputMode) -> Result<(), String> {
match mode {
VoiceOutputMode::Type => type_text(text),
VoiceOutputMode::Clipboard => copy_to_clipboard(text),
VoiceOutputMode::Both => {
copy_to_clipboard(text)?;
type_text(text)
}
}
}
/// 模拟键盘输入文字
fn type_text(text: &str) -> Result<(), String> {
use enigo::{Enigo, Keyboard, Settings};
let mut enigo =
Enigo::new(&Settings::default()).map_err(|e| format!("初始化键盘模拟器失败: {e}"))?;
enigo.text(text).map_err(|e| format!("键盘输入失败: {e}"))?;
tracing::info!("[语音输出] 键盘输入完成: {} 字符", text.chars().count());
Ok(())
}
/// 复制到剪贴板
fn copy_to_clipboard(text: &str) -> Result<(), String> {
let mut clipboard = Clipboard::new().map_err(|e| format!("初始化剪贴板失败: {e}"))?;
clipboard
.set_text(text)
.map_err(|e| format!("复制到剪贴板失败: {e}"))?;
tracing::info!("[语音输出] 已复制到剪贴板: {} 字符", text.chars().count());
Ok(())
let output_mode = match mode {
VoiceOutputMode::Type => OutputMode::Type,
VoiceOutputMode::Clipboard => OutputMode::Clipboard,
VoiceOutputMode::Both => OutputMode::Both,
};
let mut handler = OutputHandler::new().map_err(|e| format!("初始化输出处理器失败: {e}"))?;
handler
.output(text, output_mode)
.map_err(|e| format!("输出文本失败: {e}"))
}
+5 -497
View File
@@ -1,508 +1,16 @@
//! 录音服务
//! 录音服务桥接层
//!
//! 管理录音状态,提供录音控制接口。
//!
//! ## 线程安全设计
//!
//! 由于 `cpal::Stream` 不实现 `Send` trait,无法直接在 Tauri 的 async 命令中使用。
//! 本模块采用**独立线程 + channel 通信**的方案:
//!
//! ```text
//! ┌─────────────────┐ Command ┌─────────────────┐
//! │ Tauri Command │ ───────────────> │ Recording │
//! │ (async) │ │ Thread │
//! │ │ <─────────────── │ (owns Stream) │
//! └─────────────────┘ Response └─────────────────┘
//! ```
//!
//! - 录音线程拥有 `cpal::Stream`,在独立线程中运行
//! - Tauri 命令通过 channel 发送控制指令
//! - 录音线程通过 channel 返回结果
//! 录音核心逻辑已迁移到 `voice-core` 的 `threaded_recorder` 模块。
//! 本模块保留 Tauri State 包装和向后兼容导出路径。
use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use std::sync::mpsc::{self, Receiver, Sender};
use std::sync::Arc;
use std::thread::{self, JoinHandle};
use std::time::Instant;
use voice_core::types::AudioData;
/// 麦克风设备信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AudioDeviceInfo {
/// 设备 ID(用于选择设备)
pub id: String,
/// 设备名称
pub name: String,
/// 是否为默认设备
pub is_default: bool,
}
pub use voice_core::{AudioDeviceInfo, RecordingCommand, RecordingResponse, RecordingService};
/// 获取所有可用的麦克风设备
pub fn list_audio_devices() -> Result<Vec<AudioDeviceInfo>, String> {
use cpal::traits::{DeviceTrait, HostTrait};
let host = cpal::default_host();
let default_device = host.default_input_device();
let default_name = default_device.as_ref().and_then(|d| d.name().ok());
let devices: Vec<AudioDeviceInfo> = host
.input_devices()
.map_err(|e| format!("无法枚举音频设备: {e}"))?
.filter_map(|device| {
let name = device.name().ok()?;
let is_default = default_name.as_ref().map(|n| n == &name).unwrap_or(false);
Some(AudioDeviceInfo {
id: name.clone(),
name,
is_default,
})
})
.collect();
Ok(devices)
}
/// 录音控制命令
#[derive(Debug)]
pub enum RecordingCommand {
/// 开始录音(可选指定设备 ID)
Start(Option<String>),
/// 停止录音
Stop,
/// 取消录音
Cancel,
/// 关闭录音线程
Shutdown,
}
/// 录音响应
#[derive(Debug)]
pub enum RecordingResponse {
/// 操作成功
Ok,
/// 停止录音成功,返回音频数据
AudioData(AudioData),
/// 操作失败
Error(String),
}
/// 录音服务
///
/// 使用独立线程管理 cpal::Stream,通过 channel 与 Tauri 命令通信
pub struct RecordingService {
/// 命令发送端
command_tx: Option<Sender<RecordingCommand>>,
/// 响应接收端
response_rx: Option<Receiver<RecordingResponse>>,
/// 录音线程句柄
thread_handle: Option<JoinHandle<()>>,
/// 是否正在录音(共享状态,用于快速查询)
is_recording: Arc<AtomicBool>,
/// 当前音量级别(共享状态,用于快速查询)
volume_level: Arc<AtomicU32>,
/// 录音开始时间(共享状态)
start_time: Arc<Mutex<Option<Instant>>>,
}
impl RecordingService {
/// 创建新的录音服务
pub fn new() -> Self {
Self {
command_tx: None,
response_rx: None,
thread_handle: None,
is_recording: Arc::new(AtomicBool::new(false)),
volume_level: Arc::new(AtomicU32::new(0)),
start_time: Arc::new(Mutex::new(None)),
}
}
/// 确保录音线程已启动
fn ensure_thread_started(&mut self) {
if self.command_tx.is_some() {
return;
}
let (cmd_tx, cmd_rx) = mpsc::channel::<RecordingCommand>();
let (resp_tx, resp_rx) = mpsc::channel::<RecordingResponse>();
let is_recording = Arc::clone(&self.is_recording);
let volume_level = Arc::clone(&self.volume_level);
let start_time = Arc::clone(&self.start_time);
let handle = thread::spawn(move || {
recording_thread_main(cmd_rx, resp_tx, is_recording, volume_level, start_time);
});
self.command_tx = Some(cmd_tx);
self.response_rx = Some(resp_rx);
self.thread_handle = Some(handle);
tracing::info!("[录音服务] 录音线程已启动");
}
/// 开始录音(可选指定设备 ID)
pub fn start(&mut self, device_id: Option<String>) -> Result<(), String> {
self.ensure_thread_started();
let tx = self.command_tx.as_ref().ok_or("录音线程未启动")?;
let rx = self.response_rx.as_ref().ok_or("录音线程未启动")?;
tx.send(RecordingCommand::Start(device_id))
.map_err(|e| format!("发送命令失败: {e}"))?;
match rx.recv() {
Ok(RecordingResponse::Ok) => {
tracing::info!("[录音服务] 开始录音");
Ok(())
}
Ok(RecordingResponse::Error(e)) => Err(e),
Ok(_) => Err("意外的响应".to_string()),
Err(e) => Err(format!("接收响应失败: {e}")),
}
}
/// 停止录音并返回音频数据
pub fn stop(&mut self) -> Result<AudioData, String> {
let tx = self.command_tx.as_ref().ok_or("录音线程未启动")?;
let rx = self.response_rx.as_ref().ok_or("录音线程未启动")?;
tx.send(RecordingCommand::Stop)
.map_err(|e| format!("发送命令失败: {e}"))?;
match rx.recv() {
Ok(RecordingResponse::AudioData(audio)) => {
tracing::info!("[录音服务] 停止录音,时长: {:.2}s", audio.duration_secs);
Ok(audio)
}
Ok(RecordingResponse::Error(e)) => Err(e),
Ok(_) => Err("意外的响应".to_string()),
Err(e) => Err(format!("接收响应失败: {e}")),
}
}
/// 取消录音
pub fn cancel(&mut self) {
if let Some(tx) = &self.command_tx {
let _ = tx.send(RecordingCommand::Cancel);
// 使用 try_recv 避免阻塞,或者设置超时
if let Some(rx) = &self.response_rx {
// 尝试接收响应,但不阻塞太久
use std::time::Duration;
match rx.recv_timeout(Duration::from_millis(500)) {
Ok(_) => tracing::info!("[录音服务] 取消录音成功"),
Err(std::sync::mpsc::RecvTimeoutError::Timeout) => {
tracing::warn!("[录音服务] 取消录音超时,强制继续");
}
Err(e) => {
tracing::warn!("[录音服务] 取消录音响应错误: {}", e);
}
}
}
}
// 无论如何都重置状态
self.is_recording.store(false, Ordering::SeqCst);
self.volume_level.store(0, Ordering::SeqCst);
*self.start_time.lock() = None;
}
/// 获取当前音量级别(0-100)
pub fn get_volume(&self) -> u32 {
self.volume_level.load(Ordering::SeqCst)
}
/// 获取录音时长(秒)
pub fn get_duration(&self) -> f32 {
self.start_time
.lock()
.map(|t| t.elapsed().as_secs_f32())
.unwrap_or(0.0)
}
/// 是否正在录音
pub fn is_recording(&self) -> bool {
self.is_recording.load(Ordering::SeqCst)
}
/// 关闭录音服务
pub fn shutdown(&mut self) {
if let Some(tx) = self.command_tx.take() {
let _ = tx.send(RecordingCommand::Shutdown);
}
if let Some(handle) = self.thread_handle.take() {
let _ = handle.join();
}
self.response_rx = None;
tracing::info!("[录音服务] 已关闭");
}
}
impl Default for RecordingService {
fn default() -> Self {
Self::new()
}
}
impl Drop for RecordingService {
fn drop(&mut self) {
self.shutdown();
}
}
/// 录音线程主函数
///
/// 在独立线程中运行,拥有 cpal::Stream
fn recording_thread_main(
cmd_rx: Receiver<RecordingCommand>,
resp_tx: Sender<RecordingResponse>,
is_recording: Arc<AtomicBool>,
volume_level: Arc<AtomicU32>,
start_time: Arc<Mutex<Option<Instant>>>,
) {
use cpal::traits::{DeviceTrait, HostTrait, StreamTrait};
// 录音数据缓冲区
let samples: Arc<Mutex<Vec<i16>>> = Arc::new(Mutex::new(Vec::new()));
// 当前活跃的音频流
let mut active_stream: Option<cpal::Stream> = None;
// 实际使用的采样率和声道数
let mut actual_sample_rate: u32 = 16000;
#[allow(unused_assignments)]
let mut actual_channels: u16 = 1;
tracing::debug!("[录音线程] 开始运行");
loop {
match cmd_rx.recv() {
Ok(RecordingCommand::Start(device_id)) => {
// 如果已在录音,返回错误
if is_recording.load(Ordering::SeqCst) {
let _ = resp_tx.send(RecordingResponse::Error("已在录音中".to_string()));
continue;
}
// 清空缓冲区
samples.lock().clear();
// 获取输入设备
let host = cpal::default_host();
let device = if let Some(ref id) = device_id {
// 查找指定设备
host.input_devices()
.ok()
.and_then(|mut devices| {
devices.find(|d| d.name().ok().as_ref() == Some(id))
})
.or_else(|| {
tracing::warn!("[录音线程] 未找到指定设备 {},使用默认设备", id);
host.default_input_device()
})
} else {
host.default_input_device()
};
let device = match device {
Some(d) => d,
None => {
let _ =
resp_tx.send(RecordingResponse::Error("未找到麦克风设备".to_string()));
continue;
}
};
tracing::info!("[录音线程] 使用麦克风: {:?}", device.name());
// 获取设备支持的配置
let supported_config = match device.default_input_config() {
Ok(c) => c,
Err(e) => {
let _ = resp_tx
.send(RecordingResponse::Error(format!("获取音频配置失败: {e}")));
continue;
}
};
tracing::info!(
"[录音线程] 设备支持配置: 采样率={}, 声道={}",
supported_config.sample_rate().0,
supported_config.channels()
);
// 使用设备默认配置
actual_sample_rate = supported_config.sample_rate().0;
actual_channels = supported_config.channels();
let config = cpal::StreamConfig {
channels: actual_channels,
sample_rate: supported_config.sample_rate(),
buffer_size: cpal::BufferSize::Default,
};
// 创建共享状态的克隆
let samples_clone = Arc::clone(&samples);
let volume_clone = Arc::clone(&volume_level);
let is_rec_clone = Arc::clone(&is_recording);
let channels = actual_channels;
// 回调计数器(用于调试)
let callback_count = Arc::new(AtomicU32::new(0));
let callback_count_clone = Arc::clone(&callback_count);
// 创建输入流
let stream = match device.build_input_stream(
&config,
move |data: &[f32], _: &cpal::InputCallbackInfo| {
if !is_rec_clone.load(Ordering::SeqCst) {
return;
}
// 增加回调计数
let count = callback_count_clone.fetch_add(1, Ordering::SeqCst);
if count == 0 {
tracing::info!("[录音线程] 首次收到音频数据,数据长度: {}", data.len());
} else if count % 100 == 0 {
tracing::debug!("[录音线程] 已收到 {} 次音频回调", count);
}
// 计算音量级别(使用 RMS 均方根,更准确反映音量)
let sum_sq: f32 = data.iter().map(|s| s * s).sum();
let rms = (sum_sq / data.len() as f32).sqrt();
// 将 RMS 值映射到 0-100 范围
// 静音时 RMS 约 0.001-0.01,说话时约 0.02-0.1
// 使用更高的系数来提高灵敏度
let level = ((rms * 1500.0).min(100.0)) as u32;
// 每 50 次回调打印一次音量(用于调试)
if count % 50 == 0 {
tracing::debug!("[录音线程] RMS: {:.6}, 音量: {}%", rms, level);
}
volume_clone.store(level, Ordering::SeqCst);
// 如果是多声道,转换为单声道
let mono_data: Vec<f32> = if channels > 1 {
data.chunks(channels as usize)
.map(|chunk| chunk.iter().sum::<f32>() / channels as f32)
.collect()
} else {
data.to_vec()
};
// 转换为 i16 并存储
let i16_samples: Vec<i16> = mono_data
.iter()
.map(|&s| (s.clamp(-1.0, 1.0) * i16::MAX as f32) as i16)
.collect();
samples_clone.lock().extend(i16_samples);
},
|err| {
tracing::error!("[录音线程] 录音流错误: {}", err);
},
None,
) {
Ok(s) => s,
Err(e) => {
let _ =
resp_tx.send(RecordingResponse::Error(format!("创建音频流失败: {e}")));
continue;
}
};
// 开始播放(录音)
if let Err(e) = stream.play() {
let _ = resp_tx.send(RecordingResponse::Error(format!("启动录音失败: {e}")));
continue;
}
tracing::info!("[录音线程] stream.play() 成功,等待音频数据...");
// 保存流和状态
active_stream = Some(stream);
is_recording.store(true, Ordering::SeqCst);
*start_time.lock() = Some(Instant::now());
let _ = resp_tx.send(RecordingResponse::Ok);
tracing::info!(
"[录音线程] 开始录音,采样率: {}, 声道: {}",
actual_sample_rate,
actual_channels
);
}
Ok(RecordingCommand::Stop) => {
if !is_recording.load(Ordering::SeqCst) {
let _ = resp_tx.send(RecordingResponse::Error("未在录音中".to_string()));
continue;
}
// 停止录音
is_recording.store(false, Ordering::SeqCst);
// 停止并释放流
if let Some(stream) = active_stream.take() {
drop(stream);
}
// 获取录音数据(已转换为单声道)
let audio_samples = samples.lock().clone();
let audio = AudioData::new(audio_samples, actual_sample_rate, 1);
// 重置开始时间
*start_time.lock() = None;
volume_level.store(0, Ordering::SeqCst);
// 检查录音时长
if !audio.is_valid() {
let _ = resp_tx.send(RecordingResponse::Error(
"录音时间过短(需要至少 0.5 秒)".to_string(),
));
continue;
}
let _ = resp_tx.send(RecordingResponse::AudioData(audio));
tracing::info!("[录音线程] 停止录音");
}
Ok(RecordingCommand::Cancel) => {
// 停止录音
is_recording.store(false, Ordering::SeqCst);
// 停止并释放流
if let Some(stream) = active_stream.take() {
drop(stream);
}
// 清空缓冲区
samples.lock().clear();
// 重置状态
*start_time.lock() = None;
volume_level.store(0, Ordering::SeqCst);
let _ = resp_tx.send(RecordingResponse::Ok);
tracing::info!("[录音线程] 取消录音");
}
Ok(RecordingCommand::Shutdown) => {
// 清理资源
is_recording.store(false, Ordering::SeqCst);
if let Some(stream) = active_stream.take() {
drop(stream);
}
tracing::info!("[录音线程] 收到关闭命令,退出");
break;
}
Err(_) => {
// channel 已关闭,退出线程
tracing::info!("[录音线程] channel 已关闭,退出");
break;
}
}
}
voice_core::list_audio_devices().map_err(|e| e.to_string())
}
/// 全局录音服务状态(Tauri State 包装)
@@ -1,13 +0,0 @@
# Seeds for failure cases proptest has generated in the past. It is
# automatically read and these particular cases re-run before any
# novel cases are generated.
#
# It is recommended to check this file in to source control so that
# everyone who runs the test benefits from these saved cases.
cc 399023994da7d3a7d7407cbcc210e9c88f626c7158b717607e01a3f767d1e0b6 # shrinks to num_providers = 2
cc 9598ca62cf55f54ee97f00786a0e6029df29a5ede07c2f6ccd731a92ac2f1d6e # shrinks to num_keys = 2
cc 8ecfedd6b97ec094a400ca1af4c6c011f39a60688dd76327247ca8a54ca3240c # shrinks to num_keys = 2
cc d9e6f7a966ae7126d118843e3c99009616930f30348a08dfedaeab933fe9877b # shrinks to num_errors = 1
cc dfc5e61afb3ab4ec5b6283e3b92fa88ae5149458c321170000e222b95bd499e4 # shrinks to name = "aaa", api_host = "https://aaa.aa/"
cc b35b5acac2443a80f05fd96b8f46cd2b80e38a54e73b4a09aaaf5c3b68af319c # shrinks to api_key = "a0a0___a-0a-aA_-A---"
cc 05448979dc0877ad4bffe94f37f10f79ba6243d3289e0bf7629616ac8901c292 # shrinks to api_key = "-A0a_aa0-A0a_-a-_Aaa", alias = None
-852
View File
@@ -1,852 +0,0 @@
//! API Key Provider 属性测试
//!
//! 使用 proptest 进行属性测试,验证 API Key Provider 服务的正确性。
//!
//! **Feature: provider-ui-refactor**
use proptest::prelude::*;
use std::collections::HashSet;
use std::sync::Arc;
use tempfile::TempDir;
use proxycast_lib::database::dao::api_key_provider::{
ApiKeyEntry, ApiKeyProvider, ApiKeyProviderDao, ApiProviderType, ProviderGroup,
};
use proxycast_lib::database::DbConnection;
use proxycast_lib::services::api_key_provider_service::ApiKeyProviderService;
use rusqlite::Connection;
/// 测试上下文
#[allow(dead_code)]
struct TestContext {
pub temp_dir: TempDir,
pub db: DbConnection,
pub service: ApiKeyProviderService,
}
impl TestContext {
/// 创建测试上下文
pub fn new() -> Result<Self, Box<dyn std::error::Error>> {
let temp_dir = TempDir::new()?;
let db_path = temp_dir.path().join("test.db");
let conn = Connection::open(&db_path)?;
// 创建表结构
conn.execute(
"CREATE TABLE IF NOT EXISTS api_key_providers (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
type TEXT NOT NULL,
api_host TEXT NOT NULL,
is_system INTEGER NOT NULL DEFAULT 0,
group_name TEXT NOT NULL,
enabled INTEGER NOT NULL DEFAULT 0,
sort_order INTEGER NOT NULL DEFAULT 0,
api_version TEXT,
project TEXT,
location TEXT,
region TEXT,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
)",
[],
)?;
conn.execute(
"CREATE TABLE IF NOT EXISTS api_keys (
id TEXT PRIMARY KEY,
provider_id TEXT NOT NULL,
api_key_encrypted TEXT NOT NULL,
alias TEXT,
enabled INTEGER NOT NULL DEFAULT 1,
usage_count INTEGER NOT NULL DEFAULT 0,
error_count INTEGER NOT NULL DEFAULT 0,
last_used_at TEXT,
created_at TEXT NOT NULL,
FOREIGN KEY (provider_id) REFERENCES api_key_providers(id) ON DELETE CASCADE
)",
[],
)?;
conn.execute(
"CREATE TABLE IF NOT EXISTS provider_ui_state (
key TEXT PRIMARY KEY,
value TEXT NOT NULL
)",
[],
)?;
let db = Arc::new(std::sync::Mutex::new(conn));
let service = ApiKeyProviderService::new();
Ok(Self {
temp_dir,
db,
service,
})
}
/// 创建测试 Provider
pub fn create_test_provider(&self, id: &str) -> Result<ApiKeyProvider, String> {
let now = chrono::Utc::now();
let provider = ApiKeyProvider {
id: id.to_string(),
name: format!("Test Provider {id}"),
provider_type: ApiProviderType::Openai,
api_host: "https://api.test.com".to_string(),
is_system: false,
group: ProviderGroup::Custom,
enabled: true,
sort_order: 0,
api_version: None,
project: None,
location: None,
region: None,
custom_models: vec![],
created_at: now,
updated_at: now,
};
let conn = self.db.lock().map_err(|e| e.to_string())?;
ApiKeyProviderDao::insert_provider(&conn, &provider).map_err(|e| e.to_string())?;
Ok(provider)
}
/// 添加测试 API Key
pub fn add_test_api_key(
&self,
provider_id: &str,
api_key: &str,
) -> Result<ApiKeyEntry, String> {
self.service
.add_api_key(&self.db, provider_id, api_key, None)
}
}
// ============================================================================
// Property 12: 轮询负载均衡正确性
// **Validates: Requirements 7.3**
// ============================================================================
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// Property 12: 轮询负载均衡正确性
///
/// *对于任意* 拥有 N 个启用的 API Key 的 Provider,连续 N 次获取 API Key 应各返回不同的 Key
///
/// **Feature: provider-ui-refactor, Property 12: 轮询负载均衡正确性**
/// **Validates: Requirements 7.3**
#[test]
fn test_round_robin_load_balancing(num_keys in 2usize..10) {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建测试 Provider
let provider_id = format!("test-provider-{}", uuid::Uuid::new_v4());
ctx.create_test_provider(&provider_id).expect("Failed to create provider");
// 添加 N 个 API Keys
let mut expected_keys = Vec::new();
for i in 0..num_keys {
let api_key = format!("sk-test-key-{provider_id}-{i}");
ctx.add_test_api_key(&provider_id, &api_key).expect("Failed to add API key");
expected_keys.push(api_key);
}
// 连续获取 N 次 API Key
let mut retrieved_keys = Vec::new();
for _ in 0..num_keys {
let key = ctx.service
.get_next_api_key(&ctx.db, &provider_id)
.expect("Failed to get next API key")
.expect("No API key returned");
retrieved_keys.push(key);
}
// 验证:连续 N 次获取应返回 N 个不同的 Key
let unique_keys: HashSet<_> = retrieved_keys.iter().collect();
prop_assert_eq!(
unique_keys.len(),
num_keys,
"Expected {} unique keys, but got {}. Keys: {:?}",
num_keys,
unique_keys.len(),
retrieved_keys
);
// 验证:所有返回的 Key 都在预期列表中
for key in &retrieved_keys {
prop_assert!(
expected_keys.contains(key),
"Unexpected key returned: {}",
key
);
}
}
/// Property 12 补充测试:轮询循环性
///
/// *对于任意* 拥有 N 个启用的 API Key 的 Provider,获取 2N 次应该循环使用所有 Key
///
/// **Feature: provider-ui-refactor, Property 12: 轮询负载均衡正确性**
/// **Validates: Requirements 7.3**
#[test]
fn test_round_robin_cycling(num_keys in 2usize..8) {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建测试 Provider
let provider_id = format!("test-provider-cycle-{}", uuid::Uuid::new_v4());
ctx.create_test_provider(&provider_id).expect("Failed to create provider");
// 添加 N 个 API Keys
for i in 0..num_keys {
let api_key = format!("sk-cycle-key-{provider_id}-{i}");
ctx.add_test_api_key(&provider_id, &api_key).expect("Failed to add API key");
}
// 获取 2N 次 API Key
let mut first_cycle = Vec::new();
let mut second_cycle = Vec::new();
for i in 0..(num_keys * 2) {
let key = ctx.service
.get_next_api_key(&ctx.db, &provider_id)
.expect("Failed to get next API key")
.expect("No API key returned");
if i < num_keys {
first_cycle.push(key);
} else {
second_cycle.push(key);
}
}
// 验证:第一轮和第二轮应该返回相同的 Key 序列
prop_assert_eq!(
first_cycle,
second_cycle,
"Round robin should cycle through keys in the same order"
);
}
}
// ============================================================================
// Property 13: API Key 使用统计正确性
// **Validates: Requirements 7.4**
// ============================================================================
proptest! {
#![proptest_config(ProptestConfig::with_cases(20))]
/// Property 13: API Key 使用统计正确性
///
/// *对于任意* API Key 使用记录操作,使用次数应正确递增
///
/// **Feature: provider-ui-refactor, Property 13: API Key 使用统计正确性**
/// **Validates: Requirements 7.4**
#[test]
fn test_usage_count_increment(num_usages in 1usize..10) {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建测试 Provider
let provider_id = format!("test-provider-usage-{}", uuid::Uuid::new_v4());
ctx.create_test_provider(&provider_id).expect("Failed to create provider");
// 添加 API Key
let api_key = format!("sk-usage-test-{provider_id}");
let entry = ctx.add_test_api_key(&provider_id, &api_key)
.expect("Failed to add API key");
// 初始使用次数应为 0
prop_assert_eq!(entry.usage_count, 0, "Initial usage count should be 0");
// 记录 N 次使用
for _ in 0..num_usages {
ctx.service.record_usage(&ctx.db, &entry.id)
.expect("Failed to record usage");
}
// 获取更新后的 API Key
let conn = ctx.db.lock().expect("Failed to lock db");
let updated = ApiKeyProviderDao::get_api_key_by_id(&conn, &entry.id)
.expect("Failed to get API key")
.expect("API key not found");
// 验证:使用次数应等于记录次数
prop_assert_eq!(
updated.usage_count as usize,
num_usages,
"Usage count should equal number of record_usage calls"
);
// 验证:最后使用时间应被更新
prop_assert!(
updated.last_used_at.is_some(),
"last_used_at should be set after usage"
);
}
/// Property 13 补充测试:错误次数递增
///
/// *对于任意* API Key 错误记录操作,错误次数应正确递增
///
/// **Feature: provider-ui-refactor, Property 13: API Key 使用统计正确性**
/// **Validates: Requirements 7.4**
#[test]
fn test_error_count_increment(num_errors in 1usize..10) {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建测试 Provider
let provider_id = format!("test-provider-error-{}", uuid::Uuid::new_v4());
ctx.create_test_provider(&provider_id).expect("Failed to create provider");
// 添加 API Key
let api_key = format!("sk-error-test-{provider_id}");
let entry = ctx.add_test_api_key(&provider_id, &api_key)
.expect("Failed to add API key");
// 初始错误次数应为 0
prop_assert_eq!(entry.error_count, 0, "Initial error count should be 0");
// 记录 N 次错误
for _ in 0..num_errors {
ctx.service.record_error(&ctx.db, &entry.id)
.expect("Failed to record error");
}
// 获取更新后的 API Key
let conn = ctx.db.lock().expect("Failed to lock db");
let updated = ApiKeyProviderDao::get_api_key_by_id(&conn, &entry.id)
.expect("Failed to get API key")
.expect("API key not found");
// 验证:错误次数应等于记录次数
prop_assert_eq!(
updated.error_count as usize,
num_errors,
"Error count should equal number of record_error calls"
);
}
/// Property 13 补充测试:使用和错误统计独立
///
/// *对于任意* API Key,使用次数和错误次数应独立递增
///
/// **Feature: provider-ui-refactor, Property 13: API Key 使用统计正确性**
/// **Validates: Requirements 7.4**
#[test]
fn test_usage_and_error_independent(
num_usages in 1usize..5,
num_errors in 1usize..5
) {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建测试 Provider
let provider_id = format!("test-provider-mixed-{}", uuid::Uuid::new_v4());
ctx.create_test_provider(&provider_id).expect("Failed to create provider");
// 添加 API Key
let api_key = format!("sk-mixed-test-{provider_id}");
let entry = ctx.add_test_api_key(&provider_id, &api_key)
.expect("Failed to add API key");
// 交替记录使用和错误
for i in 0..(num_usages + num_errors) {
if i < num_usages {
ctx.service.record_usage(&ctx.db, &entry.id)
.expect("Failed to record usage");
}
if i < num_errors {
ctx.service.record_error(&ctx.db, &entry.id)
.expect("Failed to record error");
}
}
// 获取更新后的 API Key
let conn = ctx.db.lock().expect("Failed to lock db");
let updated = ApiKeyProviderDao::get_api_key_by_id(&conn, &entry.id)
.expect("Failed to get API key")
.expect("API key not found");
// 验证:使用次数和错误次数应独立
prop_assert_eq!(
updated.usage_count as usize,
num_usages,
"Usage count should equal number of record_usage calls"
);
prop_assert_eq!(
updated.error_count as usize,
num_errors,
"Error count should equal number of record_error calls"
);
}
}
// ============================================================================
// Property 16: 数据持久化 Round-Trip
// **Validates: Requirements 9.1**
// ============================================================================
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// Property 16: 数据持久化 Round-Trip
///
/// *对于任意* Provider 配置,保存后重新加载应得到等价的配置数据
///
/// **Feature: provider-ui-refactor, Property 16: 数据持久化 Round-Trip**
/// **Validates: Requirements 9.1**
#[test]
fn test_provider_persistence_round_trip(
name in "[a-zA-Z0-9 ]{3,30}",
api_host in "https://[a-z]{3,10}\\.[a-z]{2,5}/[a-z]{0,10}"
) {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建 Provider
let provider = ctx.service
.add_custom_provider(
&ctx.db,
name.clone(),
ApiProviderType::Openai,
api_host.clone(),
None,
None,
None,
None,
)
.expect("Failed to create provider");
// 重新加载 Provider
let loaded = ctx.service
.get_provider(&ctx.db, &provider.id)
.expect("Failed to get provider")
.expect("Provider not found");
// 验证:加载的数据应与保存的数据等价
prop_assert_eq!(&loaded.provider.id, &provider.id, "ID should match");
prop_assert_eq!(&loaded.provider.name, &name, "Name should match");
prop_assert_eq!(&loaded.provider.api_host, &api_host, "API host should match");
prop_assert_eq!(loaded.provider.is_system, false, "Should not be system provider");
prop_assert_eq!(loaded.provider.group, ProviderGroup::Custom, "Group should be Custom");
}
/// Property 16 补充测试:UI 状态持久化 Round-Trip
///
/// *对于任意* UI 状态键值对,保存后重新加载应得到相同的值
///
/// **Feature: provider-ui-refactor, Property 16: 数据持久化 Round-Trip**
/// **Validates: Requirements 9.1, 8.4**
#[test]
fn test_ui_state_persistence_round_trip(
key in "[a-z_]{3,20}",
value in "[a-zA-Z0-9_,\\[\\]\"{}:]{1,100}"
) {
let ctx = TestContext::new().expect("Failed to create test context");
// 保存 UI 状态
ctx.service
.set_ui_state(&ctx.db, &key, &value)
.expect("Failed to set UI state");
// 重新加载 UI 状态
let loaded = ctx.service
.get_ui_state(&ctx.db, &key)
.expect("Failed to get UI state")
.expect("UI state not found");
// 验证:加载的值应与保存的值相同
prop_assert_eq!(&loaded, &value, "UI state value should match");
}
/// Property 16 补充测试:Provider 排序持久化 Round-Trip
///
/// *对于任意* Provider 排序顺序,保存后重新加载应保持相同的顺序
///
/// **Feature: provider-ui-refactor, Property 16: 数据持久化 Round-Trip**
/// **Validates: Requirements 9.1, 8.4**
#[test]
fn test_provider_sort_order_persistence(num_providers in 2usize..6) {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建多个 Provider
let mut provider_ids = Vec::new();
for i in 0..num_providers {
let provider = ctx.service
.add_custom_provider(
&ctx.db,
format!("Provider {i}"),
ApiProviderType::Openai,
format!("https://api{i}.test.com"),
None,
None,
None,
None,
)
.expect("Failed to create provider");
provider_ids.push(provider.id);
}
// 反转排序顺序
let reversed_ids: Vec<_> = provider_ids.iter().rev().cloned().collect();
let sort_orders: Vec<(String, i32)> = reversed_ids
.iter()
.enumerate()
.map(|(i, id)| (id.clone(), i as i32))
.collect();
// 更新排序顺序
ctx.service
.update_provider_sort_orders(&ctx.db, sort_orders)
.expect("Failed to update sort orders");
// 重新加载所有 Provider
let loaded = ctx.service
.get_all_providers(&ctx.db)
.expect("Failed to get providers");
// 过滤出我们创建的 Provider
let our_providers: Vec<_> = loaded
.iter()
.filter(|p| provider_ids.contains(&p.provider.id))
.collect();
// 验证:排序顺序应与更新后的顺序一致
for (i, expected_id) in reversed_ids.iter().enumerate() {
let provider = our_providers
.iter()
.find(|p| &p.provider.id == expected_id)
.expect("Provider not found");
prop_assert_eq!(
provider.provider.sort_order,
i as i32,
"Sort order should match for provider {}",
expected_id
);
}
}
/// Property 16 补充测试:API Key 持久化 Round-Trip
///
/// *对于任意* API Key,保存后重新加载应得到等价的数据(除了加密的 key)
///
/// **Feature: provider-ui-refactor, Property 16: 数据持久化 Round-Trip**
/// **Validates: Requirements 9.1**
#[test]
fn test_api_key_persistence_round_trip(
api_key in "[a-zA-Z0-9_-]{20,50}",
alias in proptest::option::of("[a-zA-Z0-9 ]{3,20}")
) {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建 Provider
let provider_id = format!("test-provider-key-rt-{}", uuid::Uuid::new_v4());
ctx.create_test_provider(&provider_id).expect("Failed to create provider");
// 添加 API Key
let entry = ctx.service
.add_api_key(&ctx.db, &provider_id, &api_key, alias.clone())
.expect("Failed to add API key");
// 重新加载 Provider(包含 API Keys)
let loaded = ctx.service
.get_provider(&ctx.db, &provider_id)
.expect("Failed to get provider")
.expect("Provider not found");
// 找到我们添加的 API Key
let loaded_key = loaded.api_keys
.iter()
.find(|k| k.id == entry.id)
.expect("API Key not found");
// 验证:加载的数据应与保存的数据等价
prop_assert_eq!(&loaded_key.id, &entry.id, "ID should match");
prop_assert_eq!(&loaded_key.provider_id, &provider_id, "Provider ID should match");
prop_assert_eq!(&loaded_key.alias, &alias, "Alias should match");
prop_assert_eq!(loaded_key.enabled, true, "Should be enabled by default");
prop_assert_eq!(loaded_key.usage_count, 0, "Usage count should be 0");
prop_assert_eq!(loaded_key.error_count, 0, "Error count should be 0");
// 验证:解密后的 API Key 应与原始值相同
let decrypted = ctx.service
.decrypt_api_key(&loaded_key.api_key_encrypted)
.expect("Failed to decrypt");
prop_assert_eq!(&decrypted, &api_key, "Decrypted API key should match original");
}
}
// ============================================================================
// Property 17: API Key 加密存储
// **Validates: Requirements 9.2**
// ============================================================================
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// Property 17: API Key 加密存储
///
/// *对于任意* 存储的 API Key,数据库中的值不应为明文
///
/// **Feature: provider-ui-refactor, Property 17: API Key 加密存储**
/// **Validates: Requirements 9.2**
#[test]
fn test_api_key_encryption(api_key in "[a-zA-Z0-9_-]{20,50}") {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建测试 Provider
let provider_id = format!("test-provider-enc-{}", uuid::Uuid::new_v4());
ctx.create_test_provider(&provider_id).expect("Failed to create provider");
// 添加 API Key
let entry = ctx.add_test_api_key(&provider_id, &api_key)
.expect("Failed to add API key");
// 验证:存储的值不是明文
prop_assert_ne!(
&entry.api_key_encrypted,
&api_key,
"API Key should be encrypted, not stored as plaintext"
);
// 验证:加密后的值看起来像 Base64
prop_assert!(
entry.api_key_encrypted.chars().all(|c| c.is_alphanumeric() || c == '+' || c == '/' || c == '='),
"Encrypted value should be Base64 encoded"
);
// 验证:可以正确解密
let decrypted = ctx.service.decrypt_api_key(&entry.api_key_encrypted)
.expect("Failed to decrypt API key");
prop_assert_eq!(
&decrypted,
&api_key,
"Decrypted key should match original"
);
}
/// Property 17 补充测试:加密 Round-Trip
///
/// *对于任意* API Key,加密后解密应得到原始值
///
/// **Feature: provider-ui-refactor, Property 17: API Key 加密存储**
/// **Validates: Requirements 9.2**
#[test]
fn test_encryption_round_trip(api_key in "[a-zA-Z0-9_-]{10,100}") {
let service = ApiKeyProviderService::new();
// 加密
let encrypted = service.encrypt_api_key(&api_key);
// 验证:加密后不等于原文
prop_assert_ne!(
&encrypted,
&api_key,
"Encrypted value should differ from original"
);
// 解密
let decrypted = service.decrypt_api_key(&encrypted)
.expect("Failed to decrypt");
// 验证:解密后等于原文
prop_assert_eq!(
&decrypted,
&api_key,
"Decrypted value should match original"
);
}
}
#[cfg(test)]
mod unit_tests {
use super::*;
/// 单元测试:基本的 Provider CRUD 操作
#[test]
fn test_provider_crud() {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建 Provider
let provider = ctx
.service
.add_custom_provider(
&ctx.db,
"Test Provider".to_string(),
ApiProviderType::Openai,
"https://api.test.com".to_string(),
None,
None,
None,
None,
)
.expect("Failed to create provider");
assert!(provider.id.starts_with("custom-"));
assert_eq!(provider.name, "Test Provider");
assert!(!provider.is_system);
// 获取 Provider
let retrieved = ctx
.service
.get_provider(&ctx.db, &provider.id)
.expect("Failed to get provider")
.expect("Provider not found");
assert_eq!(retrieved.provider.id, provider.id);
// 更新 Provider
let updated = ctx
.service
.update_provider(
&ctx.db,
&provider.id,
Some("Updated Name".to_string()),
None, // provider_type
None, // api_host
Some(false), // enabled
None, // sort_order
None, // api_version
None, // project
None, // location
None, // region
None, // custom_models
)
.expect("Failed to update provider");
assert_eq!(updated.name, "Updated Name");
assert!(!updated.enabled);
// 删除 Provider
let deleted = ctx
.service
.delete_custom_provider(&ctx.db, &provider.id)
.expect("Failed to delete provider");
assert!(deleted);
// 验证已删除
let not_found = ctx
.service
.get_provider(&ctx.db, &provider.id)
.expect("Failed to get provider");
assert!(not_found.is_none());
}
/// 单元测试:API Key CRUD 操作
#[test]
fn test_api_key_crud() {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建 Provider
let provider_id = "test-provider-key-crud";
ctx.create_test_provider(provider_id)
.expect("Failed to create provider");
// 添加 API Key
let key = ctx
.add_test_api_key(provider_id, "sk-test-key-123")
.expect("Failed to add API key");
assert!(!key.id.is_empty());
assert_eq!(key.provider_id, provider_id);
assert!(key.enabled);
// 切换启用状态
let toggled = ctx
.service
.toggle_api_key(&ctx.db, &key.id, false)
.expect("Failed to toggle API key");
assert!(!toggled.enabled);
// 更新别名
let aliased = ctx
.service
.update_api_key_alias(&ctx.db, &key.id, Some("My Key".to_string()))
.expect("Failed to update alias");
assert_eq!(aliased.alias, Some("My Key".to_string()));
// 删除 API Key
let deleted = ctx
.service
.delete_api_key(&ctx.db, &key.id)
.expect("Failed to delete API key");
assert!(deleted);
}
/// 单元测试:重复 API Key 检测
/// 验证修复:第一次添加 API Key 无法保存的问题
#[test]
fn test_duplicate_api_key_detection() {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建 Provider
let provider_id = "test-provider-duplicate";
ctx.create_test_provider(provider_id)
.expect("Failed to create provider");
// 第一次添加 API Key 应该成功
let api_key = "sk-duplicate-test-123";
let first_result = ctx.add_test_api_key(provider_id, api_key);
assert!(first_result.is_ok(), "第一次添加应该成功");
// 第二次添加相同的 API Key 应该失败
let second_result = ctx.add_test_api_key(provider_id, api_key);
assert!(second_result.is_err(), "第二次添加相同 API Key 应该失败");
assert!(
second_result.unwrap_err().contains("该 API Key 已存在"),
"错误信息应该提示 API Key 已存在"
);
// 验证 Provider 中只有一个 API Key
let provider = ctx
.service
.get_provider(&ctx.db, provider_id)
.expect("Failed to get provider")
.expect("Provider not found");
assert_eq!(provider.api_keys.len(), 1, "应该只有一个 API Key");
}
/// 单元测试:系统 Provider 不能删除
#[test]
fn test_system_provider_cannot_be_deleted() {
let ctx = TestContext::new().expect("Failed to create test context");
// 创建系统 Provider
let now = chrono::Utc::now();
let provider = ApiKeyProvider {
id: "system-openai".to_string(),
name: "OpenAI".to_string(),
provider_type: ApiProviderType::Openai,
api_host: "https://api.openai.com".to_string(),
is_system: true, // 系统 Provider
group: ProviderGroup::Mainstream,
enabled: true,
sort_order: 1,
api_version: None,
project: None,
location: None,
region: None,
custom_models: vec![],
created_at: now,
updated_at: now,
};
{
let conn = ctx.db.lock().expect("Failed to lock db");
ApiKeyProviderDao::insert_provider(&conn, &provider).expect("Failed to insert");
}
// 尝试删除系统 Provider
let result = ctx.service.delete_custom_provider(&ctx.db, "system-openai");
assert!(result.is_err());
assert!(result.unwrap_err().contains("不允许删除系统 Provider"));
}
}
@@ -1 +0,0 @@