diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index f1b21f172..33cfe7773 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -10428,6 +10428,7 @@ dependencies = [ "futures-util", "hmac", "hound", + "parking_lot", "reqwest 0.12.28", "serde", "serde_json", diff --git a/src-tauri/capabilities/default.json b/src-tauri/capabilities/default.json new file mode 100644 index 000000000..867f9f09e --- /dev/null +++ b/src-tauri/capabilities/default.json @@ -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" + ] +} \ No newline at end of file diff --git a/src-tauri/crates/agent/Cargo.toml b/src-tauri/crates/agent/Cargo.toml new file mode 100644 index 000000000..e5a1ff791 --- /dev/null +++ b/src-tauri/crates/agent/Cargo.toml @@ -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 diff --git a/src-tauri/src/agent/event_converter.rs b/src-tauri/crates/agent/src/event_converter.rs similarity index 100% rename from src-tauri/src/agent/event_converter.rs rename to src-tauri/crates/agent/src/event_converter.rs diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs new file mode 100644 index 000000000..4094ba3b8 --- /dev/null +++ b/src-tauri/crates/agent/src/lib.rs @@ -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; diff --git a/src-tauri/src/agent/mcp_bridge.rs b/src-tauri/crates/agent/src/mcp_bridge.rs similarity index 88% rename from src-tauri/src/agent/mcp_bridge.rs rename to src-tauri/crates/agent/src/mcp_bridge.rs index 9447ff94e..305f885bf 100644 --- a/src-tauri/src/agent/mcp_bridge.rs +++ b/src-tauri/crates/agent/src/mcp_bridge.rs @@ -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 桥接客户端 /// diff --git a/src-tauri/src/agent/prompt/builder.rs b/src-tauri/crates/agent/src/prompt/builder.rs similarity index 97% rename from src-tauri/src/agent/prompt/builder.rs rename to src-tauri/crates/agent/src/prompt/builder.rs index a1f25bbe2..6b8426a46 100644 --- a/src-tauri/src/agent/prompt/builder.rs +++ b/src-tauri/crates/agent/src/prompt/builder.rs @@ -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")); } } diff --git a/src-tauri/src/agent/prompt/mod.rs b/src-tauri/crates/agent/src/prompt/mod.rs similarity index 100% rename from src-tauri/src/agent/prompt/mod.rs rename to src-tauri/crates/agent/src/prompt/mod.rs index 0145759df..687fa8ed7 100644 --- a/src-tauri/src/agent/prompt/mod.rs +++ b/src-tauri/crates/agent/src/prompt/mod.rs @@ -7,8 +7,8 @@ //! - templates - 提示词模板定义 //! - builder - 提示词构建器 -pub mod templates; pub mod builder; +pub mod templates; pub use builder::SystemPromptBuilder; pub use templates::*; diff --git a/src-tauri/src/agent/prompt/templates.rs b/src-tauri/crates/agent/src/prompt/templates.rs similarity index 99% rename from src-tauri/src/agent/prompt/templates.rs rename to src-tauri/crates/agent/src/prompt/templates.rs index ba9136bca..b538512a7 100644 --- a/src-tauri/src/agent/prompt/templates.rs +++ b/src-tauri/crates/agent/src/prompt/templates.rs @@ -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#"# 输出风格 diff --git a/src-tauri/crates/skills/Cargo.toml b/src-tauri/crates/skills/Cargo.toml new file mode 100644 index 000000000..fd0a1c7fe --- /dev/null +++ b/src-tauri/crates/skills/Cargo.toml @@ -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 diff --git a/src-tauri/crates/skills/src/execution_callback.rs b/src-tauri/crates/skills/src/execution_callback.rs new file mode 100644 index 000000000..f8f2854d4 --- /dev/null +++ b/src-tauri/crates/skills/src/execution_callback.rs @@ -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, + pub error: Option, +} + +/// 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>); +} diff --git a/src-tauri/crates/skills/src/lib.rs b/src-tauri/crates/skills/src/lib.rs new file mode 100644 index 000000000..a0d3366b6 --- /dev/null +++ b/src-tauri/crates/skills/src/lib.rs @@ -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, +}; diff --git a/src-tauri/crates/skills/src/llm_provider.rs b/src-tauri/crates/skills/src/llm_provider.rs new file mode 100644 index 000000000..614b35670 --- /dev/null +++ b/src-tauri/crates/skills/src/llm_provider.rs @@ -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; +} diff --git a/src-tauri/src/skills/skill_loader.rs b/src-tauri/crates/skills/src/skill_loader.rs similarity index 84% rename from src-tauri/src/skills/skill_loader.rs rename to src-tauri/crates/skills/src/skill_loader.rs index ac5cab1ae..1f2312a04 100644 --- a/src-tauri/src/skills/skill_loader.rs +++ b/src-tauri/crates/skills/src/skill_loader.rs @@ -1,7 +1,6 @@ //! Skill 定义加载器 //! //! 负责从 `~/.proxycast/skills//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, - /// Skill 描述 pub description: Option, - /// 允许的工具 #[serde(rename = "allowed-tools")] pub allowed_tools: Option, - /// 参数提示 #[serde(rename = "argument-hint")] pub argument_hint: Option, - /// 使用场景 #[serde(rename = "when-to-use")] pub when_to_use: Option, - /// 版本 pub version: Option, - /// 偏好模型 pub model: Option, - /// 偏好 Provider pub provider: Option, - /// 是否禁用模型调用 #[serde(rename = "disable-model-invocation")] pub disable_model_invocation: Option, - /// 执行模式 #[serde(rename = "execution-mode")] pub execution_mode: Option, } /// 内部 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>, - /// 参数提示 pub argument_hint: Option, - /// 使用场景 pub when_to_use: Option, - /// 偏好模型 pub model: Option, - /// 偏好 Provider pub provider: Option, - /// 是否禁用模型调用 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> { +pub fn parse_allowed_tools(value: Option<&str>) -> Option> { value.and_then(|v| { if v.is_empty() { return None; @@ -130,7 +108,7 @@ pub(crate) fn parse_allowed_tools(value: Option<&str>) -> Option> { } /// 解析布尔值字段 -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 { @@ -178,12 +156,12 @@ pub(crate) fn load_skill_from_file( } /// 获取 ProxyCast Skills 目录 -pub(crate) fn get_proxycast_skills_dir() -> Option { +pub fn get_proxycast_skills_dir() -> Option { dirs::home_dir().map(|home| home.join(".proxycast").join("skills")) } /// 从目录加载所有 Skills -pub(crate) fn load_skills_from_directory(dir_path: &Path) -> Vec { +pub fn load_skills_from_directory(dir_path: &Path) -> Vec { let mut results = Vec::new(); if !dir_path.exists() { @@ -216,7 +194,7 @@ pub(crate) fn load_skills_from_directory(dir_path: &Path) -> Vec Result { +pub fn find_skill_by_name(skill_name: &str) -> Result { let skills_dir = get_proxycast_skills_dir().ok_or_else(|| "无法获取 Skills 目录".to_string())?; diff --git a/src-tauri/crates/voice-core/Cargo.toml b/src-tauri/crates/voice-core/Cargo.toml index 6ce8e8c2f..c11503347 100644 --- a/src-tauri/crates/voice-core/Cargo.toml +++ b/src-tauri/crates/voice-core/Cargo.toml @@ -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"] } diff --git a/src-tauri/crates/voice-core/README.md b/src-tauri/crates/voice-core/README.md index a73a95760..390250b4f 100644 --- a/src-tauri/crates/voice-core/README.md +++ b/src-tauri/crates/voice-core/README.md @@ -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 diff --git a/src-tauri/crates/voice-core/src/device.rs b/src-tauri/crates/voice-core/src/device.rs new file mode 100644 index 000000000..b8b70127d --- /dev/null +++ b/src-tauri/crates/voice-core/src/device.rs @@ -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> { + 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) +} diff --git a/src-tauri/crates/voice-core/src/lib.rs b/src-tauri/crates/voice-core/src/lib.rs index f8f68c3a7..8d69b3ff8 100644 --- a/src-tauri/crates/voice-core/src/lib.rs +++ b/src-tauri/crates/voice-core/src/lib.rs @@ -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::*; diff --git a/src-tauri/crates/voice-core/src/threaded_recorder.rs b/src-tauri/crates/voice-core/src/threaded_recorder.rs new file mode 100644 index 000000000..2ebc5389e --- /dev/null +++ b/src-tauri/crates/voice-core/src/threaded_recorder.rs @@ -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), + /// 停止录音 + Stop, + /// 取消录音 + Cancel, + /// 关闭录音线程 + Shutdown, +} + +/// 录音响应 +#[derive(Debug)] +pub enum RecordingResponse { + /// 操作成功 + Ok, + /// 停止录音成功,返回音频数据 + AudioData(AudioData), + /// 操作失败 + Error(String), +} + +/// 录音服务 +/// +/// 使用独立线程管理 cpal::Stream,通过 channel 与 Tauri 命令通信 +pub struct RecordingService { + /// 命令发送端 + command_tx: Option>, + /// 响应接收端 + response_rx: Option>, + /// 录音线程句柄 + thread_handle: Option>, + /// 是否正在录音(共享状态,用于快速查询) + is_recording: Arc, + /// 当前音量级别(共享状态,用于快速查询) + volume_level: Arc, + /// 录音开始时间(共享状态) + start_time: Arc>>, +} + +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::(); + let (resp_tx, resp_rx) = mpsc::channel::(); + + 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) -> 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 { + 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, + resp_tx: Sender, + is_recording: Arc, + volume_level: Arc, + start_time: Arc>>, +) { + use cpal::traits::{DeviceTrait, HostTrait, StreamTrait}; + + // 录音数据缓冲区 + let samples: Arc>> = Arc::new(Mutex::new(Vec::new())); + // 当前活跃的音频流 + let mut active_stream: Option = 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 = if channels > 1 { + data.chunks(channels as usize) + .map(|chunk| chunk.iter().sum::() / channels as f32) + .collect() + } else { + data.to_vec() + }; + + // 转换为 i16 并存储 + let i16_samples: Vec = 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; + } + } + } +} diff --git a/src-tauri/crates/voice-core/src/types.rs b/src-tauri/crates/voice-core/src/types.rs index 67fc37a17..f64405d7d 100644 --- a/src-tauri/crates/voice-core/src/types.rs +++ b/src-tauri/crates/voice-core/src/types.rs @@ -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 { + self.samples + .iter() + .flat_map(|sample| sample.to_le_bytes()) + .collect() + } + /// 转换为 WAV 格式字节 pub fn to_wav_bytes(&self) -> Vec { let mut cursor = std::io::Cursor::new(Vec::new()); diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index 80d596a03..ab27e8be5 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -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, }; diff --git a/src-tauri/src/agent/prompt/README.md b/src-tauri/src/agent/prompt/README.md deleted file mode 100644 index ef4113a1f..000000000 --- a/src-tauri/src/agent/prompt/README.md +++ /dev/null @@ -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?; -``` diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 13f0c3ab2..ad84a92af 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -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::() { 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 事件发射器已设置"); } // 初始化截图对话模块 diff --git a/src-tauri/src/processor/mod.rs b/src-tauri/src/processor/mod.rs index 659d34f09..f59942270 100644 --- a/src-tauri/src/processor/mod.rs +++ b/src-tauri/src/processor/mod.rs @@ -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>, - /// 模型映射器 - pub mapper: Arc>, - /// 参数注入器 - pub injector: Arc>, - /// 重试器 - pub retrier: Arc, - /// 故障转移器 - pub failover: Arc, - /// 超时控制器 - pub timeout: Arc, - /// 插件管理器 - pub plugins: Arc, - /// 统计聚合器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) - pub stats: Arc>, - /// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) - pub tokens: Arc>, - /// 凭证池服务 - pub pool_service: Arc, - /// 热重载协调锁(避免配置更新期间请求读取不一致的配置) - pub reload_lock: Arc>, -} - -impl RequestProcessor { - /// 创建新的请求处理器 - pub fn new( - router: Arc>, - mapper: Arc>, - injector: Arc>, - retrier: Arc, - failover: Arc, - timeout: Arc, - plugins: Arc, - stats: Arc>, - tokens: Arc>, - pool_service: Arc, - ) -> 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) -> 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, - stats: Arc>, - tokens: Arc>, - ) -> 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, 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 { - 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 { - // 1. 解析模型别名 - self.resolve_model_for_context(ctx).await; - - // 2. 根据解析后的模型选择 Provider - self.route_for_context(ctx).await - } -} +pub use proxycast_processor::*; #[cfg(test)] mod tests; diff --git a/src-tauri/src/processor/tests.rs b/src-tauri/src/processor/tests.rs index 95527d081..4b33d86c2 100644 --- a/src-tauri/src/processor/tests.rs +++ b/src-tauri/src/processor/tests.rs @@ -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() { diff --git a/src-tauri/src/skills/execution_callback.rs b/src-tauri/src/skills/execution_callback.rs index 4da4a535a..bdedd458e 100644 --- a/src-tauri/src/skills/execution_callback.rs +++ b/src-tauri/src/skills/execution_callback.rs @@ -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, - /// 错误信息(失败时) - pub error: Option, -} - -/// 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 中添加属性测试 -} diff --git a/src-tauri/src/skills/llm_provider.rs b/src-tauri/src/skills/llm_provider.rs index 57f68cc41..36bd86106 100644 --- a/src-tauri/src/skills/llm_provider.rs +++ b/src-tauri/src/skills/llm_provider.rs @@ -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; -} +use proxycast_skills::{LlmProvider, SkillError}; /// ProxyCast LLM Provider /// diff --git a/src-tauri/src/skills/mod.rs b/src-tauri/src/skills/mod.rs index 387695195..dc6c6d6f4 100644 --- a/src-tauri/src/skills/mod.rs +++ b/src-tauri/src/skills/mod.rs @@ -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; diff --git a/src-tauri/src/voice/README.md b/src-tauri/src/voice/README.md index 846d7155e..8ee394ca0 100644 --- a/src-tauri/src/voice/README.md +++ b/src-tauri/src/voice/README.md @@ -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 ──→ 上述所有服务 ``` diff --git a/src-tauri/src/voice/asr_service.rs b/src-tauri/src/voice/asr_service.rs index c7eb8052a..543e99784 100644 --- a/src-tauri/src/voice/asr_service.rs +++ b/src-tauri/src/voice/asr_service.rs @@ -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 = 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 { 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 { 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, - } - - 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 { let config = credential.xunfei_config.as_ref().ok_or("讯飞配置缺失")?; - - // 将 PCM 字节转换为 i16 采样 - let samples: Vec = 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, 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 { + 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) } } diff --git a/src-tauri/src/voice/commands.rs b/src-tauri/src/voice/commands.rs index 9426466a5..26b2567d3 100644 --- a/src-tauri/src/voice/commands.rs +++ b/src-tauri/src/voice/commands.rs @@ -328,11 +328,7 @@ pub async fn stop_recording( ); // 将 i16 样本转换为字节(小端序) - let bytes: Vec = audio - .samples - .iter() - .flat_map(|&s| s.to_le_bytes()) - .collect(); + let bytes = audio.to_pcm16le_bytes(); Ok(StopRecordingResult { audio_data: bytes, diff --git a/src-tauri/src/voice/output_service.rs b/src-tauri/src/voice/output_service.rs index ec175dc86..f1d720f1b 100644 --- a/src-tauri/src/voice/output_service.rs +++ b/src-tauri/src/voice/output_service.rs @@ -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}")) } diff --git a/src-tauri/src/voice/recording_service.rs b/src-tauri/src/voice/recording_service.rs index 6521a5ce4..a44175a2f 100644 --- a/src-tauri/src/voice/recording_service.rs +++ b/src-tauri/src/voice/recording_service.rs @@ -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, 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 = 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), - /// 停止录音 - Stop, - /// 取消录音 - Cancel, - /// 关闭录音线程 - Shutdown, -} - -/// 录音响应 -#[derive(Debug)] -pub enum RecordingResponse { - /// 操作成功 - Ok, - /// 停止录音成功,返回音频数据 - AudioData(AudioData), - /// 操作失败 - Error(String), -} - -/// 录音服务 -/// -/// 使用独立线程管理 cpal::Stream,通过 channel 与 Tauri 命令通信 -pub struct RecordingService { - /// 命令发送端 - command_tx: Option>, - /// 响应接收端 - response_rx: Option>, - /// 录音线程句柄 - thread_handle: Option>, - /// 是否正在录音(共享状态,用于快速查询) - is_recording: Arc, - /// 当前音量级别(共享状态,用于快速查询) - volume_level: Arc, - /// 录音开始时间(共享状态) - start_time: Arc>>, -} - -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::(); - let (resp_tx, resp_rx) = mpsc::channel::(); - - 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) -> 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 { - 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, - resp_tx: Sender, - is_recording: Arc, - volume_level: Arc, - start_time: Arc>>, -) { - use cpal::traits::{DeviceTrait, HostTrait, StreamTrait}; - - // 录音数据缓冲区 - let samples: Arc>> = Arc::new(Mutex::new(Vec::new())); - // 当前活跃的音频流 - let mut active_stream: Option = 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 = if channels > 1 { - data.chunks(channels as usize) - .map(|chunk| chunk.iter().sum::() / channels as f32) - .collect() - } else { - data.to_vec() - }; - - // 转换为 i16 并存储 - let i16_samples: Vec = 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 包装) diff --git a/src-tauri/tests/api_key_provider_tests.proptest-regressions b/src-tauri/tests/api_key_provider_tests.proptest-regressions deleted file mode 100644 index 53a9b07e7..000000000 --- a/src-tauri/tests/api_key_provider_tests.proptest-regressions +++ /dev/null @@ -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 diff --git a/src-tauri/tests/api_key_provider_tests.rs b/src-tauri/tests/api_key_provider_tests.rs deleted file mode 100644 index fc7ec5a2f..000000000 --- a/src-tauri/tests/api_key_provider_tests.rs +++ /dev/null @@ -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> { - 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 { - 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 { - 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")); - } -} diff --git a/src-tauri/tests/filter_expression_tests.rs b/src-tauri/tests/filter_expression_tests.rs deleted file mode 100644 index 8b1378917..000000000 --- a/src-tauri/tests/filter_expression_tests.rs +++ /dev/null @@ -1 +0,0 @@ -