From bf0c1a05fbab23d75d95d176da8554e07c9072be Mon Sep 17 00:00:00 2001 From: coso Date: Mon, 9 Feb 2026 02:08:53 +0800 Subject: [PATCH] chore: bump version to 0.61.0 --- package.json | 2 +- src-tauri/Cargo.lock | 40 +- src-tauri/Cargo.toml | 4 +- src-tauri/crates/agent/Cargo.toml | 7 + src-tauri/crates/agent/src/aster_state.rs | 602 ++++++++++++++ .../crates/agent/src/aster_state_support.rs | 163 ++++ .../crates/agent/src/credential_bridge.rs | 563 +++++++++++++ src-tauri/crates/agent/src/lib.rs | 22 +- src-tauri/crates/agent/src/mcp_bridge.rs | 10 +- src-tauri/crates/agent/src/session_store.rs | 134 ++++ .../crates/agent/src/subagent_scheduler.rs | 246 ++++++ src-tauri/crates/core/src/app_bootstrap.rs | 68 ++ src-tauri/crates/core/src/app_utils.rs | 120 +++ src-tauri/crates/core/src/lib.rs | 6 + src-tauri/crates/core/src/logger.rs | 5 + src-tauri/crates/core/src/tray_format.rs | 116 +++ src-tauri/crates/core/src/tray_menu_meta.rs | 201 +++++ src-tauri/crates/core/src/tray_state.rs | 288 +++++++ src-tauri/crates/credential/src/quota.rs | 8 - src-tauri/crates/server-utils/src/lib.rs | 51 ++ src-tauri/crates/server/src/handlers/mod.rs | 7 +- src-tauri/crates/services/Cargo.toml | 5 + .../services/src/file_browser_service.rs | 730 +++++++++++++++++ src-tauri/crates/services/src/lib.rs | 19 +- .../src/screenshot_capture_service.rs | 244 ++++++ .../services/src/screenshot_image_service.rs | 31 + .../crates/services/src/sysinfo_service.rs | 180 +++++ .../crates/services/src/voice_asr_service.rs | 310 +++++++ .../services/src/voice_command_service.rs | 155 ++++ .../services/src/voice_config_service.rs | 132 +++ .../services/src/voice_output_service.rs | 22 + .../services/src/voice_processor_service.rs | 49 ++ .../services/src/voice_recording_service.rs | 35 + src-tauri/crates/skills/Cargo.toml | 5 + src-tauri/crates/skills/src/lib.rs | 4 +- .../skills/src/proxycast_llm_provider.rs | 543 +++++++++++++ src-tauri/crates/voice-core/README.md | 1 + src-tauri/crates/voice-core/src/lib.rs | 1 + .../crates/voice-core/src/text_polish.rs | 133 +++ .../voice-core/src/threaded_recorder.rs | 2 +- src-tauri/src/agent/README.md | 12 +- src-tauri/src/agent/aster_agent.rs | 175 +--- src-tauri/src/agent/aster_state.rs | 756 +----------------- src-tauri/src/agent/credential_bridge.rs | 563 +------------ src-tauri/src/agent/mod.rs | 4 +- src-tauri/src/agent/subagent_scheduler.rs | 256 +----- src-tauri/src/app/bootstrap.rs | 88 +- src-tauri/src/app/commands/config.rs | 2 +- src-tauri/src/app/commands/server.rs | 2 +- src-tauri/src/app/runner.rs | 4 +- src-tauri/src/app/setup.rs | 6 +- src-tauri/src/app/state.rs | 19 +- src-tauri/src/app/types.rs | 4 +- src-tauri/src/app/utils.rs | 134 +--- .../src/commands/api_key_provider_cmd.rs | 2 +- src-tauri/src/commands/aster_agent_cmd.rs | 6 +- src-tauri/src/commands/config_cmd.rs | 2 +- src-tauri/src/commands/connection_cmd.rs | 2 +- src-tauri/src/commands/context_memory.rs | 2 +- src-tauri/src/commands/general_chat_cmd.rs | 2 +- src-tauri/src/commands/kiro_local.rs | 2 +- src-tauri/src/commands/machine_id_cmd.rs | 2 +- src-tauri/src/commands/material_cmd.rs | 2 +- src-tauri/src/commands/mcp_cmd.rs | 4 +- src-tauri/src/commands/model_cmd.rs | 2 +- src-tauri/src/commands/model_registry_cmd.rs | 2 +- src-tauri/src/commands/network_cmd.rs | 4 +- src-tauri/src/commands/orchestrator_cmd.rs | 2 +- src-tauri/src/commands/persona_cmd.rs | 2 +- src-tauri/src/commands/plugin_cmd.rs | 4 +- src-tauri/src/commands/plugin_install_cmd.rs | 4 +- src-tauri/src/commands/prompt_cmd.rs | 4 +- src-tauri/src/commands/provider_pool_cmd.rs | 8 +- src-tauri/src/commands/screenshot_cmd.rs | 38 +- src-tauri/src/commands/skill_cmd.rs | 5 +- src-tauri/src/commands/skill_exec_cmd.rs | 8 +- src-tauri/src/commands/switch_cmd.rs | 7 +- src-tauri/src/commands/template_cmd.rs | 2 +- src-tauri/src/commands/terminal_cmd.rs | 2 +- src-tauri/src/commands/tool_hooks.rs | 2 +- src-tauri/src/commands/unified_chat_cmd.rs | 2 +- src-tauri/src/commands/update_cmd.rs | 4 +- src-tauri/src/commands/usage_cmd.rs | 2 +- src-tauri/src/commands/websocket_cmd.rs | 2 +- src-tauri/src/commands/workspace_cmd.rs | 2 +- src-tauri/src/config/mod.rs | 15 +- src-tauri/src/config/observer/mod.rs | 21 +- .../src/config/observer/tauri_observer.rs | 1 + src-tauri/src/config/tests.rs | 78 +- src-tauri/src/content/mod.rs | 3 - src-tauri/src/credential/mod.rs | 29 - src-tauri/src/database/mod.rs | 5 - src-tauri/src/dev_bridge.rs | 2 +- src-tauri/src/dev_bridge/dispatcher.rs | 63 +- src-tauri/src/lib.rs | 65 +- src-tauri/src/logger.rs | 12 +- src-tauri/src/mcp/README.md | 42 - src-tauri/src/mcp/mod.rs | 19 - src-tauri/src/memory/mod.rs | 3 - src-tauri/src/middleware/mod.rs | 5 - src-tauri/src/models/mod.rs | 39 - src-tauri/src/plugin/mod.rs | 21 +- src-tauri/src/processor/mod.rs | 9 - src-tauri/src/processor/steps/mod.rs | 4 - src-tauri/src/router/mod.rs | 3 - src-tauri/src/screenshot/README.md | 6 +- src-tauri/src/screenshot/capture.rs | 252 +----- src-tauri/src/screenshot/config.rs | 3 +- src-tauri/src/screenshot/shortcut.rs | 2 +- src-tauri/src/server/mod.rs | 5 - src-tauri/src/server_utils.rs | 5 - src-tauri/src/services/README.md | 22 +- .../src/services/file_browser_service.rs | 700 +--------------- src-tauri/src/services/mod.rs | 32 +- src-tauri/src/services/sysinfo_service.rs | 190 +---- .../src/services/update_check_service.rs | 8 + src-tauri/src/session/mod.rs | 21 - src-tauri/src/skills/README.md | 11 +- src-tauri/src/skills/llm_provider.rs | 543 +------------ src-tauri/src/skills/mod.rs | 16 +- src-tauri/src/terminal/mod.rs | 25 +- .../tests.rs => tests/credential_tests.rs} | 23 +- src-tauri/src/tests/mod.rs | 2 + .../tests.rs => tests/processor_tests.rs} | 4 +- src-tauri/src/tray/format.rs | 124 +-- src-tauri/src/tray/menu.rs | 82 +- src-tauri/src/tray/state.rs | 218 +---- src-tauri/src/tray/sync.rs | 48 -- src-tauri/src/voice/README.md | 19 +- src-tauri/src/voice/asr_service.rs | 321 +------- src-tauri/src/voice/commands.rs | 231 +----- src-tauri/src/voice/config.rs | 48 +- src-tauri/src/voice/output_service.rs | 24 +- src-tauri/src/voice/processor.rs | 170 +--- src-tauri/src/voice/recording_service.rs | 38 +- src-tauri/src/websocket/mod.rs | 19 - src-tauri/src/workspace/mod.rs | 3 - src-tauri/tauri.conf.json | 2 +- 138 files changed, 5629 insertions(+), 5450 deletions(-) create mode 100644 src-tauri/crates/agent/src/aster_state.rs create mode 100644 src-tauri/crates/agent/src/aster_state_support.rs create mode 100644 src-tauri/crates/agent/src/credential_bridge.rs create mode 100644 src-tauri/crates/agent/src/session_store.rs create mode 100644 src-tauri/crates/agent/src/subagent_scheduler.rs create mode 100644 src-tauri/crates/core/src/app_bootstrap.rs create mode 100644 src-tauri/crates/core/src/app_utils.rs create mode 100644 src-tauri/crates/core/src/tray_format.rs create mode 100644 src-tauri/crates/core/src/tray_menu_meta.rs create mode 100644 src-tauri/crates/core/src/tray_state.rs create mode 100644 src-tauri/crates/services/src/file_browser_service.rs create mode 100644 src-tauri/crates/services/src/screenshot_capture_service.rs create mode 100644 src-tauri/crates/services/src/screenshot_image_service.rs create mode 100644 src-tauri/crates/services/src/sysinfo_service.rs create mode 100644 src-tauri/crates/services/src/voice_asr_service.rs create mode 100644 src-tauri/crates/services/src/voice_command_service.rs create mode 100644 src-tauri/crates/services/src/voice_config_service.rs create mode 100644 src-tauri/crates/services/src/voice_output_service.rs create mode 100644 src-tauri/crates/services/src/voice_processor_service.rs create mode 100644 src-tauri/crates/services/src/voice_recording_service.rs create mode 100644 src-tauri/crates/skills/src/proxycast_llm_provider.rs create mode 100644 src-tauri/crates/voice-core/src/text_polish.rs delete mode 100644 src-tauri/src/content/mod.rs delete mode 100644 src-tauri/src/credential/mod.rs delete mode 100644 src-tauri/src/database/mod.rs delete mode 100644 src-tauri/src/mcp/README.md delete mode 100644 src-tauri/src/mcp/mod.rs delete mode 100644 src-tauri/src/memory/mod.rs delete mode 100644 src-tauri/src/middleware/mod.rs delete mode 100644 src-tauri/src/models/mod.rs delete mode 100644 src-tauri/src/processor/mod.rs delete mode 100644 src-tauri/src/processor/steps/mod.rs delete mode 100644 src-tauri/src/router/mod.rs delete mode 100644 src-tauri/src/server/mod.rs delete mode 100644 src-tauri/src/server_utils.rs create mode 100644 src-tauri/src/services/update_check_service.rs delete mode 100644 src-tauri/src/session/mod.rs rename src-tauri/src/{credential/tests.rs => tests/credential_tests.rs} (99%) create mode 100644 src-tauri/src/tests/mod.rs rename src-tauri/src/{processor/tests.rs => tests/processor_tests.rs} (99%) delete mode 100644 src-tauri/src/websocket/mod.rs delete mode 100644 src-tauri/src/workspace/mod.rs diff --git a/package.json b/package.json index 47c31eeec..1aceaab57 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.60.0", + "version": "0.61.0", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 33cfe7773..4804d1177 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -6635,7 +6635,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.60.0" +version = "0.61.0" dependencies = [ "anyhow", "arboard", @@ -6730,24 +6730,29 @@ dependencies = [ [[package]] name = "proxycast-agent" -version = "0.60.0" +version = "0.61.0" dependencies = [ "aster", "async-trait", "chrono", + "dirs 5.0.1", "proxycast-core", "proxycast-mcp", + "proxycast-providers", + "proxycast-services", "rmcp 0.6.4", "serde", "serde_json", + "tempfile", "tokio", "tokio-util", "tracing", + "uuid", ] [[package]] name = "proxycast-config" -version = "0.60.0" +version = "0.61.0" dependencies = [ "async-trait", "parking_lot", @@ -6763,7 +6768,7 @@ dependencies = [ [[package]] name = "proxycast-core" -version = "0.60.0" +version = "0.61.0" dependencies = [ "async-trait", "axum 0.7.9", @@ -6802,7 +6807,7 @@ dependencies = [ [[package]] name = "proxycast-credential" -version = "0.60.0" +version = "0.61.0" dependencies = [ "axum 0.7.9", "chrono", @@ -6820,7 +6825,7 @@ dependencies = [ [[package]] name = "proxycast-infra" -version = "0.60.0" +version = "0.61.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -6840,7 +6845,7 @@ dependencies = [ [[package]] name = "proxycast-mcp" -version = "0.60.0" +version = "0.61.0" dependencies = [ "async-trait", "glob", @@ -6855,7 +6860,7 @@ dependencies = [ [[package]] name = "proxycast-processor" -version = "0.60.0" +version = "0.61.0" dependencies = [ "async-trait", "parking_lot", @@ -6874,7 +6879,7 @@ dependencies = [ [[package]] name = "proxycast-providers" -version = "0.60.0" +version = "0.61.0" dependencies = [ "anyhow", "async-stream", @@ -6908,7 +6913,7 @@ dependencies = [ [[package]] name = "proxycast-server" -version = "0.60.0" +version = "0.61.0" dependencies = [ "async-stream", "axum 0.7.9", @@ -6945,7 +6950,7 @@ dependencies = [ [[package]] name = "proxycast-server-utils" -version = "0.60.0" +version = "0.61.0" dependencies = [ "axum 0.7.9", "futures", @@ -6960,7 +6965,7 @@ dependencies = [ [[package]] name = "proxycast-services" -version = "0.60.0" +version = "0.61.0" dependencies = [ "anyhow", "aster", @@ -6992,16 +6997,21 @@ dependencies = [ "url", "urlencoding", "uuid", + "voice-core", "whoami", "zip", ] [[package]] name = "proxycast-skills" -version = "0.60.0" +version = "0.61.0" dependencies = [ "async-trait", "dirs 5.0.1", + "proxycast-core", + "proxycast-providers", + "proxycast-server-utils", + "proxycast-services", "regex", "serde", "serde_json", @@ -7010,7 +7020,7 @@ dependencies = [ [[package]] name = "proxycast-terminal" -version = "0.60.0" +version = "0.61.0" dependencies = [ "async-trait", "base64 0.22.1", @@ -7037,7 +7047,7 @@ dependencies = [ [[package]] name = "proxycast-websocket" -version = "0.60.0" +version = "0.61.0" dependencies = [ "axum 0.7.9", "dashmap 5.5.3", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index e43808baa..2e5c0499e 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.60.0" +version = "0.61.0" edition = "2021" authors = ["you"] repository = "https://github.com/aiclientproxy/proxycast" @@ -179,7 +179,7 @@ version = "2.4" [package] name = "proxycast" -version = "0.60.0" +version = "0.61.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/crates/agent/Cargo.toml b/src-tauri/crates/agent/Cargo.toml index e5a1ff791..aa8cdf1dd 100644 --- a/src-tauri/crates/agent/Cargo.toml +++ b/src-tauri/crates/agent/Cargo.toml @@ -8,6 +8,8 @@ repository.workspace = true [dependencies] proxycast-core.workspace = true proxycast-mcp.workspace = true +proxycast-services.workspace = true +proxycast-providers.workspace = true aster.workspace = true rmcp.workspace = true serde.workspace = true @@ -17,3 +19,8 @@ tokio-util.workspace = true async-trait.workspace = true tracing.workspace = true chrono.workspace = true +dirs.workspace = true +uuid.workspace = true + +[dev-dependencies] +tempfile.workspace = true diff --git a/src-tauri/crates/agent/src/aster_state.rs b/src-tauri/crates/agent/src/aster_state.rs new file mode 100644 index 000000000..bc554514a --- /dev/null +++ b/src-tauri/crates/agent/src/aster_state.rs @@ -0,0 +1,602 @@ +//! Aster Agent 状态管理 +//! +//! 管理 Aster Agent 实例和相关状态 +//! 提供 Tauri 应用与 Aster 框架的桥接 +//! 支持从 ProxyCast 凭证池自动选择凭证 +//! +//! ## 重要:SessionStore 注入 +//! +//! 为了让 Aster Agent 的消息存储到 ProxyCast 数据库,必须在创建 Agent 时 +//! 注入 `ProxyCastSessionStore`。使用 `init_agent_with_db()` 方法而不是 `init_agent()`。 +//! +//! ## Agent 身份配置 +//! +//! 通过 Aster 框架的 `AgentIdentity` API 设置 ProxyCast 专属的 Agent 身份, +//! 包括名称、语言偏好、产品描述等。这是架构层面的正确做法, +//! 而不是简单地追加提示词。 +//! +//! ## Skills 集成 +//! +//! Agent 初始化时会自动加载 `~/.proxycast/skills/` 目录下的 Skills 到 +//! aster-rust 的 global_registry,使 AI 能够自动发现和调用这些 Skills。 +//! +//! 参考文档:`docs/prd/chat-architecture-redesign.md` + +use aster::agents::{Agent, SessionConfig}; +use aster::model::ModelConfig; +#[cfg(test)] +use aster::skills::{global_registry, load_skills_from_directory, SkillSource}; +use std::sync::Arc; +use tokio::sync::RwLock; +use tokio_util::sync::CancellationToken; + +use crate::credential_bridge::{create_aster_provider, AsterProviderConfig, CredentialBridge}; +use proxycast_core::database::DbConnection; +use proxycast_services::aster_session_store::ProxyCastSessionStore; + +/// Provider 配置信息 +#[derive(Debug, Clone)] +pub struct ProviderConfig { + /// Provider 名称 (openai, anthropic, google, ollama 等) + pub provider_name: String, + /// 模型名称 + pub model_name: String, + /// API Key (可选,某些 provider 从环境变量读取) + pub api_key: Option, + /// Base URL (可选,用于自定义端点) + pub base_url: Option, + /// 凭证 UUID(来自凭证池,用于记录使用和健康状态) + pub credential_uuid: Option, +} + +/// Aster Agent 全局状态 +/// +/// 在 Tauri 应用中作为 managed state 使用 +pub struct AsterAgentState { + /// Aster Agent 实例 + agent: Arc>>, + /// 当前活跃的取消令牌(用于中止正在进行的对话) + cancel_tokens: Arc>>, + /// 当前 Provider 配置 + current_provider_config: Arc>>, + /// 凭证桥接器 + credential_bridge: CredentialBridge, +} + +impl Default for AsterAgentState { + fn default() -> Self { + Self::new() + } +} + +impl AsterAgentState { + /// 创建新的 Aster Agent 状态 + pub fn new() -> Self { + Self { + agent: Arc::new(RwLock::new(None)), + cancel_tokens: Arc::new(RwLock::new(std::collections::HashMap::new())), + current_provider_config: Arc::new(RwLock::new(None)), + credential_bridge: CredentialBridge::new(), + } + } + + /// 初始化 Agent(带数据库连接) + /// + /// 创建 Agent 并注入 ProxyCastSessionStore,确保消息存储到 ProxyCast 数据库。 + /// 同时设置 ProxyCast 专属的 Agent 身份(名称、语言、描述)。 + /// 自动加载 `~/.proxycast/skills/` 目录下的 Skills 到 aster-rust 的 global_registry。 + /// + /// **推荐使用此方法**而不是 `init_agent()`。 + /// + /// # 参数 + /// - `db`: 数据库连接,用于创建 SessionStore + pub async fn init_agent_with_db(&self, db: &DbConnection) -> Result<(), String> { + let mut agent_guard = self.agent.write().await; + if agent_guard.is_none() { + // 创建 SessionStore + let session_store = Arc::new(ProxyCastSessionStore::new(db.clone())); + tracing::info!("[AsterAgent] 创建 ProxyCastSessionStore 成功"); + + // 创建 Agent 并注入 SessionStore + let agent = Agent::new().with_session_store(session_store); + + // 验证 session_store 是否被正确设置 + let has_store = agent.session_store().is_some(); + tracing::info!( + "[AsterAgent] Agent 创建完成,session_store 已设置: {}", + has_store + ); + + // 使用异步方法设置 ProxyCast 专属身份 + let identity = crate::create_proxycast_identity(); + agent.set_identity(identity).await; + + // 加载 ProxyCast Skills 到 aster-rust 的 global_registry + crate::reload_proxycast_skills(); + + *agent_guard = Some(agent); + tracing::info!( + "[AsterAgent] Agent 初始化成功,已注入 ProxyCastSessionStore、ProxyCast 身份和 Skills" + ); + } else { + tracing::debug!("[AsterAgent] Agent 已初始化,跳过"); + } + Ok(()) + } + + /// 重新加载 ProxyCast Skills + /// + /// 当用户安装或卸载 Skills 后调用此方法刷新 registry。 + pub fn reload_proxycast_skills() { + crate::reload_proxycast_skills(); + } + + /// 初始化 Agent(无数据库版本) + /// + /// **警告**:此方法创建的 Agent 不会将消息存储到 ProxyCast 数据库, + /// 消息会存储到 Aster 默认的 `~/.aster/sessions.db`。 + /// + /// 建议使用 `init_agent_with_db()` 代替。 + #[deprecated( + since = "0.1.0", + note = "请使用 init_agent_with_db() 以确保消息存储到 ProxyCast 数据库" + )] + pub async fn init_agent(&self) -> Result<(), String> { + let mut agent_guard = self.agent.write().await; + if agent_guard.is_none() { + let agent = Agent::new(); + *agent_guard = Some(agent); + tracing::warn!( + "[AsterAgent] Agent 初始化(无 SessionStore),消息将存储到 Aster 默认数据库" + ); + } + Ok(()) + } + + /// 配置 Provider + /// + /// 根据配置创建并设置 Provider + /// + /// # 参数 + /// - `config`: Provider 配置 + /// - `session_id`: 会话 ID + /// - `db`: 数据库连接(用于初始化 Agent) + pub async fn configure_provider( + &self, + config: ProviderConfig, + session_id: &str, + db: &DbConnection, + ) -> Result<(), String> { + // 确保 Agent 已初始化(使用带数据库的版本) + self.init_agent_with_db(db).await?; + + // 设置环境变量(Aster 的 provider 从环境变量读取配置) + self.set_provider_env_vars(&config); + + // 创建 ModelConfig + let model_config = ModelConfig::new(&config.model_name) + .map_err(|e| format!("创建 ModelConfig 失败: {e}"))?; + + // 创建 Provider + let provider = aster::providers::create(&config.provider_name, model_config) + .await + .map_err(|e| format!("创建 Provider 失败: {e}"))?; + + // 更新 Agent 的 Provider + let agent_guard = self.agent.read().await; + if let Some(agent) = agent_guard.as_ref() { + agent + .update_provider(provider, session_id) + .await + .map_err(|e| format!("更新 Provider 失败: {e}"))?; + } + + // 保存当前配置 + let mut config_guard = self.current_provider_config.write().await; + *config_guard = Some(config.clone()); + + tracing::info!( + "[AsterAgent] Provider 配置成功: {} / {}", + config.provider_name, + config.model_name + ); + + Ok(()) + } + + /// 从凭证池配置 Provider + /// + /// 自动从 ProxyCast 凭证池选择可用凭证并配置 Aster Provider + /// + /// # 参数 + /// - `db`: 数据库连接 + /// - `provider_type`: Provider 类型 (openai, anthropic, kiro 等) + /// - `model`: 模型名称 + /// - `session_id`: 会话 ID + pub async fn configure_provider_from_pool( + &self, + db: &DbConnection, + provider_type: &str, + model: &str, + session_id: &str, + ) -> Result { + // 确保 Agent 已初始化(使用带数据库的版本) + self.init_agent_with_db(db).await?; + + // 从凭证池选择凭证并获取配置 + let aster_config = self + .credential_bridge + .select_and_configure(db, provider_type, model) + .await + .map_err(|e| format!("从凭证池选择凭证失败: {e}"))?; + + // 创建 Provider + let provider = create_aster_provider(&aster_config) + .await + .map_err(|e| format!("创建 Provider 失败: {e}"))?; + + // 更新 Agent 的 Provider + let agent_guard = self.agent.read().await; + if let Some(agent) = agent_guard.as_ref() { + agent + .update_provider(provider, session_id) + .await + .map_err(|e| format!("更新 Provider 失败: {e}"))?; + } + + // 保存当前配置 + let config = ProviderConfig { + provider_name: aster_config.provider_name.clone(), + model_name: aster_config.model_name.clone(), + api_key: aster_config.api_key.clone(), + base_url: aster_config.base_url.clone(), + credential_uuid: Some(aster_config.credential_uuid.clone()), + }; + let mut config_guard = self.current_provider_config.write().await; + *config_guard = Some(config); + + // 记录凭证使用 + if let Err(e) = self + .credential_bridge + .record_usage(db, &aster_config.credential_uuid) + { + tracing::warn!("[AsterAgent] 记录凭证使用失败: {}", e); + } + + tracing::info!( + "[AsterAgent] 从凭证池配置 Provider 成功: {} / {} (凭证: {})", + aster_config.provider_name, + aster_config.model_name, + aster_config.credential_uuid + ); + + Ok(aster_config) + } + + /// 标记当前凭证为健康 + pub fn mark_current_healthy(&self, db: &DbConnection, model: Option<&str>) { + if let Ok(config_guard) = self.current_provider_config.try_read() { + if let Some(config) = config_guard.as_ref() { + if let Some(uuid) = &config.credential_uuid { + if let Err(e) = self.credential_bridge.mark_healthy(db, uuid, model) { + tracing::warn!("[AsterAgent] 标记凭证健康失败: {}", e); + } + } + } + } + } + + /// 标记当前凭证为不健康 + pub fn mark_current_unhealthy(&self, db: &DbConnection, error: Option<&str>) { + if let Ok(config_guard) = self.current_provider_config.try_read() { + if let Some(config) = config_guard.as_ref() { + if let Some(uuid) = &config.credential_uuid { + if let Err(e) = self.credential_bridge.mark_unhealthy(db, uuid, error) { + tracing::warn!("[AsterAgent] 标记凭证不健康失败: {}", e); + } + } + } + } + } + + /// 设置 Provider 相关的环境变量 + fn set_provider_env_vars(&self, config: &ProviderConfig) { + tracing::info!( + "[AsterAgent] set_provider_env_vars: provider_name={}, model_name={}, has_api_key={}, base_url={:?}", + config.provider_name, + config.model_name, + config.api_key.is_some(), + config.base_url + ); + + // 根据 provider 类型设置对应的环境变量 + let env_key = match config.provider_name.as_str() { + "openai" => "OPENAI_API_KEY", + "anthropic" => "ANTHROPIC_API_KEY", + "google" => "GOOGLE_API_KEY", + "deepseek" | "custom_deepseek" => "OPENAI_API_KEY", // DeepSeek 使用 OpenAI 兼容 API + "groq" => "OPENAI_API_KEY", // Groq 使用 OpenAI 兼容 API + "mistral" => "OPENAI_API_KEY", // Mistral 使用 OpenAI 兼容 API + "openrouter" => "OPENROUTER_API_KEY", + "ollama" => return, // Ollama 不需要 API Key + _ => { + tracing::warn!( + "[AsterAgent] 未知的 provider_name: {}, 使用通用 OpenAI 格式", + config.provider_name + ); + // 通用 OpenAI 兼容格式 + if let Some(api_key) = &config.api_key { + std::env::set_var("OPENAI_API_KEY", api_key); + } + if let Some(base_url) = &config.base_url { + std::env::set_var("OPENAI_BASE_URL", base_url); + } + return; + } + }; + + tracing::info!("[AsterAgent] 设置环境变量: {}=***", env_key); + + if let Some(api_key) = &config.api_key { + std::env::set_var(env_key, api_key); + } + + if let Some(base_url) = &config.base_url { + let base_url_key = format!( + "{}_BASE_URL", + config.provider_name.to_uppercase().replace("_", "") + ); + std::env::set_var(base_url_key, base_url); + } + } + + /// 获取当前 Provider 配置 + pub async fn get_provider_config(&self) -> Option { + self.current_provider_config.read().await.clone() + } + + /// 清除当前 Provider 配置 + /// + /// 用于切换凭证后重置状态,下次对话时会重新从凭证池选择凭证 + pub async fn clear_provider_config(&self) { + let mut config_guard = self.current_provider_config.write().await; + *config_guard = None; + tracing::info!("[AsterAgent] Provider 配置已清除"); + } + + /// 检查 Provider 是否已配置 + pub async fn is_provider_configured(&self) -> bool { + self.current_provider_config.read().await.is_some() + } + + /// 获取 Agent 的只读引用并执行同步操作 + pub async fn with_agent(&self, f: F) -> Result + where + F: FnOnce(&Agent) -> R, + { + let guard = self.agent.read().await; + match guard.as_ref() { + Some(agent) => Ok(f(agent)), + None => Err("Agent not initialized".to_string()), + } + } + + /// 获取 Agent 的可变引用并执行同步操作 + pub async fn with_agent_mut(&self, f: F) -> Result + where + F: FnOnce(&mut Agent) -> R, + { + let mut guard = self.agent.write().await; + match guard.as_mut() { + Some(agent) => Ok(f(agent)), + None => Err("Agent not initialized".to_string()), + } + } + + /// 获取 Agent 的 Arc 引用 + /// + /// 用于需要长期持有 Agent 引用的场景 + pub fn get_agent_arc(&self) -> Arc>> { + self.agent.clone() + } + + /// 创建新的取消令牌 + pub async fn create_cancel_token(&self, session_id: &str) -> CancellationToken { + let token = CancellationToken::new(); + let mut tokens = self.cancel_tokens.write().await; + tokens.insert(session_id.to_string(), token.clone()); + token + } + + /// 取消指定会话的操作 + pub async fn cancel_session(&self, session_id: &str) -> bool { + let tokens = self.cancel_tokens.read().await; + if let Some(token) = tokens.get(session_id) { + token.cancel(); + true + } else { + false + } + } + + /// 移除取消令牌 + pub async fn remove_cancel_token(&self, session_id: &str) { + let mut tokens = self.cancel_tokens.write().await; + tokens.remove(session_id); + } + + // ------------------------------------------------------------------------ + // 项目上下文支持 + // ------------------------------------------------------------------------ + + /// 构建带项目上下文的 System Prompt + /// + /// 加载项目的人设、素材、模板配置,构建完整的 AI 提示词。 + /// + /// # 参数 + /// - `db`: 数据库连接 + /// - `project_id`: 项目 ID + /// + /// # 返回 + /// - 成功返回构建好的 System Prompt + /// - 失败返回错误信息 + pub fn build_project_system_prompt( + db: &DbConnection, + project_id: &str, + ) -> Result { + crate::build_project_system_prompt(db, project_id) + } + + /// 创建带项目上下文的会话配置 + /// + /// 自动加载项目配置并构建 SessionConfig。 + /// + /// # 参数 + /// - `db`: 数据库连接 + /// - `session_id`: 会话 ID + /// - `project_id`: 项目 ID(可选,如果为 None 则不注入项目上下文) + /// + /// # 返回 + /// - 构建好的 SessionConfig + pub fn create_session_config_with_project( + db: &DbConnection, + session_id: &str, + project_id: Option<&str>, + ) -> SessionConfig { + crate::create_session_config_with_project(db, session_id, project_id) + } + + /// 检查 Agent 是否已初始化 + pub async fn is_initialized(&self) -> bool { + self.agent.read().await.is_some() + } +} + +pub use crate::aster_state_support::{message_helpers, SessionConfigBuilder}; + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + use tempfile::TempDir; + + #[tokio::test] + async fn test_aster_state_init() { + let state = AsterAgentState::new(); + assert!(!state.is_initialized().await); + + #[allow(deprecated)] + state.init_agent().await.unwrap(); + assert!(state.is_initialized().await); + } + + #[tokio::test] + async fn test_cancel_token() { + let state = AsterAgentState::new(); + let session_id = "test-session"; + + let token = state.create_cancel_token(session_id).await; + assert!(!token.is_cancelled()); + + assert!(state.cancel_session(session_id).await); + assert!(token.is_cancelled()); + + state.remove_cancel_token(session_id).await; + assert!(!state.cancel_session(session_id).await); + } + + // ========================================================================= + // Skills 集成测试 + // ========================================================================= + + /// 测试辅助函数:创建测试用的 Skill 目录 + fn create_test_skill(skills_dir: &std::path::Path, skill_name: &str, description: &str) { + let skill_path = skills_dir.join(skill_name); + fs::create_dir_all(&skill_path).unwrap(); + let skill_md = format!( + r#"--- +name: {} +description: {} +--- + +# {} + +这是一个测试 Skill。 +"#, + skill_name, description, skill_name + ); + fs::write(skill_path.join("SKILL.md"), skill_md).unwrap(); + } + + /// 测试:load_skills_from_directory 能正确加载 Skills + #[test] + fn test_load_skills_from_directory() { + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path(); + + // 创建测试 Skills + create_test_skill(skills_dir, "test-skill-1", "第一个测试技能"); + create_test_skill(skills_dir, "test-skill-2", "第二个测试技能"); + + // 加载 Skills + let skills = load_skills_from_directory(skills_dir, SkillSource::User); + + // 验证 + assert_eq!(skills.len(), 2); + let names: Vec<_> = skills.iter().map(|s| s.display_name.as_str()).collect(); + assert!(names.contains(&"test-skill-1")); + assert!(names.contains(&"test-skill-2")); + } + + /// 测试:空目录返回空列表 + #[test] + fn test_load_skills_empty_directory() { + let temp_dir = TempDir::new().unwrap(); + let skills = load_skills_from_directory(temp_dir.path(), SkillSource::User); + assert!(skills.is_empty()); + } + + /// 测试:不存在的目录返回空列表 + #[test] + fn test_load_skills_nonexistent_directory() { + let nonexistent = std::path::Path::new("/nonexistent/path/to/skills"); + let skills = load_skills_from_directory(nonexistent, SkillSource::User); + assert!(skills.is_empty()); + } + + /// 测试:global_registry 能正确注册和查找 Skills + #[test] + fn test_global_registry_register_and_find() { + let temp_dir = TempDir::new().unwrap(); + let skills_dir = temp_dir.path(); + + // 创建测试 Skill + create_test_skill(skills_dir, "registry-test-skill", "注册表测试技能"); + + // 加载并注册到 global_registry + let skills = load_skills_from_directory(skills_dir, SkillSource::User); + let registry = global_registry(); + + if let Ok(mut registry_guard) = registry.write() { + for skill in skills { + registry_guard.register(skill); + } + } + + // 验证能找到注册的 Skill + if let Ok(registry_guard) = registry.read() { + let found = registry_guard.find("registry-test-skill"); + assert!(found.is_some()); + assert_eq!(found.unwrap().display_name, "registry-test-skill"); + } + } + + /// 测试:reload_proxycast_skills 不会 panic(即使目录不存在) + #[test] + fn test_reload_proxycast_skills_no_panic() { + // 这个测试确保 reload_proxycast_skills 在各种情况下都不会 panic + // 即使 ~/.proxycast/skills/ 目录不存在 + AsterAgentState::reload_proxycast_skills(); + // 如果没有 panic,测试通过 + } +} diff --git a/src-tauri/crates/agent/src/aster_state_support.rs b/src-tauri/crates/agent/src/aster_state_support.rs new file mode 100644 index 000000000..9f03f27f9 --- /dev/null +++ b/src-tauri/crates/agent/src/aster_state_support.rs @@ -0,0 +1,163 @@ +//! Aster 状态支持模块 +//! +//! 提供可复用的会话配置构建、项目上下文 Prompt 构建、 +//! ProxyCast Skills 加载与 Agent 身份配置。 + +use aster::agents::{AgentIdentity, SessionConfig}; +use aster::skills::{global_registry, load_skills_from_directory, SkillSource}; +use proxycast_core::database::DbConnection; +use proxycast_services::project_context_builder::ProjectContextBuilder; + +/// 重新加载 ProxyCast Skills +pub fn reload_proxycast_skills() { + load_proxycast_skills(); +} + +/// 创建 ProxyCast 专属的 Agent 身份配置 +pub fn create_proxycast_identity() -> AgentIdentity { + AgentIdentity::new("ProxyCast 助手") + .with_language("Chinese") + .with_description( + "ProxyCast 是一个 AI 代理服务应用,帮助用户管理和使用各种 AI 模型的凭证。", + ) + .with_custom_prompt(PROXYCAST_IDENTITY_PROMPT.to_string()) +} + +/// 加载 ProxyCast Skills 到 aster-rust 的 global_registry +fn load_proxycast_skills() { + let home = match dirs::home_dir() { + Some(home_dir) => home_dir, + None => { + tracing::warn!("[AsterAgent] 无法获取 home 目录,跳过 Skills 加载"); + return; + } + }; + + let skills_dir = home.join(".proxycast").join("skills"); + if !skills_dir.exists() { + tracing::info!( + "[AsterAgent] ProxyCast Skills 目录不存在: {:?},跳过加载", + skills_dir + ); + return; + } + + let skills = load_skills_from_directory(&skills_dir, SkillSource::User); + let skill_count = skills.len(); + + if skill_count == 0 { + tracing::info!("[AsterAgent] ProxyCast Skills 目录为空,无 Skills 可加载"); + return; + } + + let registry = global_registry(); + if let Ok(mut registry_guard) = registry.write() { + for skill in skills { + let skill_name = skill.skill_name.clone(); + registry_guard.register(skill); + tracing::debug!("[AsterAgent] 已注册 Skill: {}", skill_name); + } + tracing::info!( + "[AsterAgent] 成功加载 {} 个 ProxyCast Skills 到 global_registry", + skill_count + ); + } else { + tracing::error!("[AsterAgent] 无法获取 global_registry 写锁,Skills 加载失败"); + } +} + +/// 构建带项目上下文的 System Prompt +pub fn build_project_system_prompt(db: &DbConnection, project_id: &str) -> Result { + let conn = db.lock().map_err(|e| format!("获取数据库连接失败: {e}"))?; + let context = ProjectContextBuilder::build_context(&conn, project_id) + .map_err(|e| format!("构建项目上下文失败: {e}"))?; + Ok(ProjectContextBuilder::build_system_prompt(&context)) +} + +/// 创建带项目上下文的会话配置 +pub fn create_session_config_with_project( + db: &DbConnection, + session_id: &str, + project_id: Option<&str>, +) -> SessionConfig { + let system_prompt = project_id.and_then(|pid| build_project_system_prompt(db, pid).ok()); + + SessionConfigBuilder::new(session_id) + .system_prompt(system_prompt.unwrap_or_default()) + .build() +} + +/// 会话配置构建器 +pub struct SessionConfigBuilder { + id: String, + max_turns: Option, + system_prompt: Option, +} + +impl SessionConfigBuilder { + pub fn new(id: impl Into) -> Self { + Self { + id: id.into(), + max_turns: None, + system_prompt: None, + } + } + + pub fn max_turns(mut self, turns: u32) -> Self { + self.max_turns = Some(turns); + self + } + + pub fn system_prompt(mut self, prompt: impl Into) -> Self { + self.system_prompt = Some(prompt.into()); + self + } + + pub fn build(self) -> SessionConfig { + SessionConfig { + id: self.id, + schedule_id: None, + max_turns: self.max_turns, + retry_config: None, + system_prompt: self.system_prompt, + } + } +} + +/// 消息构建辅助函数 +pub mod message_helpers { + use aster::conversation::message::Message; + + /// 创建用户文本消息 + pub fn user_text(text: impl Into) -> Message { + Message::user().with_text(text) + } + + /// 创建助手文本消息 + pub fn assistant_text(text: impl Into) -> Message { + Message::assistant().with_text(text) + } +} + +/// ProxyCast 专属的 Agent 身份提示词 +const PROXYCAST_IDENTITY_PROMPT: &str = r#"你是 ProxyCast 助手,一个专业、友好的 AI 技术伙伴。 + +## 关于 ProxyCast + +ProxyCast 是一个 AI 代理服务应用,帮助用户: +- 管理多个 AI 模型提供商的凭证(OpenAI、Claude、Gemini、Kiro 等) +- 通过统一的 API 接口访问不同的 AI 模型 +- 实现凭证池的负载均衡和健康检查 + +## 语言规范 + +1. **始终使用中文回复**:除非用户明确要求使用其他语言 +2. **代码注释使用中文**:生成代码时,注释应使用中文 +3. **技术术语保持原文**:API、JSON、HTTP、Token 等专业术语保持英文 + +## 交互风格 + +- 简洁专业,直接给出解决方案 +- 友好但不啰嗦,像经验丰富的技术伙伴 +- 遇到问题时,先分析原因再提供方案 +"#; diff --git a/src-tauri/crates/agent/src/credential_bridge.rs b/src-tauri/crates/agent/src/credential_bridge.rs new file mode 100644 index 000000000..6c3c33050 --- /dev/null +++ b/src-tauri/crates/agent/src/credential_bridge.rs @@ -0,0 +1,563 @@ +//! 凭证池桥接模块 +//! +//! 将 ProxyCast 凭证池与 Aster Provider 系统连接 +//! 支持从凭证池自动选择凭证并配置 Aster Provider +//! +//! ## 功能 +//! - 从凭证池选择可用凭证 +//! - 将凭证转换为 Aster Provider 配置 +//! - 支持 OAuth 和 API Key 两种凭证类型 +//! - 自动刷新过期的 OAuth Token +//! - 智能拆分 base_url 为 host + path,避免路径重复(如智谱 /v4/v1 问题) + +use aster::model::ModelConfig; +use aster::providers::base::Provider; +use proxycast_core::database::DbConnection; +use proxycast_core::models::provider_pool_model::{ + CredentialData, PoolProviderType, ProviderCredential, +}; +use proxycast_services::api_key_provider_service::ApiKeyProviderService; +use proxycast_services::provider_pool_service::ProviderPoolService; +use std::sync::Arc; + +/// 凭证桥接错误 +#[derive(Debug, Clone)] +pub enum CredentialBridgeError { + /// 没有可用凭证 + NoCredentials(String), + /// 凭证类型不支持 + UnsupportedCredentialType(String), + /// Provider 创建失败 + ProviderCreationFailed(String), + /// Token 刷新失败 + TokenRefreshFailed(String), + /// 数据库错误 + DatabaseError(String), +} + +impl std::fmt::Display for CredentialBridgeError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::NoCredentials(msg) => write!(f, "没有可用凭证: {msg}"), + Self::UnsupportedCredentialType(msg) => write!(f, "不支持的凭证类型: {msg}"), + Self::ProviderCreationFailed(msg) => write!(f, "Provider 创建失败: {msg}"), + Self::TokenRefreshFailed(msg) => write!(f, "Token 刷新失败: {msg}"), + Self::DatabaseError(msg) => write!(f, "数据库错误: {msg}"), + } + } +} + +impl std::error::Error for CredentialBridgeError {} + +/// Aster Provider 配置 +#[derive(Debug, Clone)] +pub struct AsterProviderConfig { + /// Provider 名称 (openai, anthropic, google 等) + pub provider_name: String, + /// 模型名称 + pub model_name: String, + /// API Key + pub api_key: Option, + /// Base URL + pub base_url: Option, + /// 凭证 UUID(用于记录使用和健康状态) + pub credential_uuid: String, +} + +/// 凭证池桥接器 +/// +/// 负责从 ProxyCast 凭证池选择凭证并转换为 Aster Provider 配置 +pub struct CredentialBridge { + pool_service: ProviderPoolService, + api_key_service: ApiKeyProviderService, +} + +impl Default for CredentialBridge { + fn default() -> Self { + Self::new() + } +} + +impl CredentialBridge { + pub fn new() -> Self { + Self { + pool_service: ProviderPoolService::new(), + api_key_service: ApiKeyProviderService::new(), + } + } + + /// 从凭证池选择凭证并创建 Aster Provider 配置 + /// + /// # 参数 + /// - `db`: 数据库连接 + /// - `provider_type`: Provider 类型 (openai, anthropic, kiro, deepseek 等) + /// - `model`: 模型名称 + /// + /// # 返回 + /// 成功时返回 AsterProviderConfig,失败时返回错误 + pub async fn select_and_configure( + &self, + db: &DbConnection, + provider_type: &str, + model: &str, + ) -> Result { + // 1. 从凭证池选择凭证 + // 将 provider_type 同时作为 provider_id_hint 传递,支持 60+ API Key Provider + // 例如 "deepseek", "moonshot", "qwen" 等 + let credential = self + .pool_service + .select_credential_with_fallback( + db, + &self.api_key_service, + provider_type, + Some(model), + Some(provider_type), // 传递 provider_id_hint 支持智能降级 + None, + ) + .await + .map_err(CredentialBridgeError::DatabaseError)? + .ok_or_else(|| { + CredentialBridgeError::NoCredentials(format!( + "没有找到 {provider_type} 类型的可用凭证" + )) + })?; + + // 2. 转换为 Aster Provider 配置,传递 provider_type 以便正确识别 Provider + self.credential_to_config(&credential, model, provider_type, db) + .await + } + + /// 将 ProxyCast 凭证转换为 Aster Provider 配置 + async fn credential_to_config( + &self, + credential: &ProviderCredential, + model: &str, + provider_type_hint: &str, + db: &DbConnection, + ) -> Result { + tracing::info!( + "[CredentialBridge] credential_to_config: provider_type_hint={}, credential_type={:?}", + provider_type_hint, + credential.provider_type + ); + + let (provider_name, api_key, base_url) = match &credential.credential { + // OpenAI API Key - 根据 provider_type_hint 确定实际的 Provider + CredentialData::OpenAIKey { api_key, base_url } => { + // 使用 provider_type_hint 来确定 aster provider 名称 + let provider = map_provider_type_to_aster(provider_type_hint); + tracing::info!( + "[CredentialBridge] OpenAIKey: provider_type_hint={} -> aster_provider={}", + provider_type_hint, + provider + ); + ( + provider.to_string(), + Some(api_key.clone()), + base_url.clone(), + ) + } + + // Claude/Anthropic API Key + CredentialData::ClaudeKey { api_key, base_url } + | CredentialData::AnthropicKey { api_key, base_url } => ( + "anthropic".to_string(), + Some(api_key.clone()), + base_url.clone(), + ), + + // Kiro OAuth - 需要获取 access_token + CredentialData::KiroOAuth { creds_file_path } => { + let token = self + .get_kiro_token(creds_file_path, db, &credential.uuid) + .await?; + // Kiro 使用 CodeWhisperer API,映射到 bedrock provider + ("bedrock".to_string(), Some(token), None) + } + + // Gemini OAuth + CredentialData::GeminiOAuth { + creds_file_path, .. + } => { + let token = self.get_oauth_token(creds_file_path).await?; + ("google".to_string(), Some(token), None) + } + + // Gemini API Key + CredentialData::GeminiApiKey { + api_key, base_url, .. + } => ( + "google".to_string(), + Some(api_key.clone()), + base_url.clone(), + ), + + // Vertex AI + CredentialData::VertexKey { + api_key, base_url, .. + } => ( + "gcpvertexai".to_string(), + Some(api_key.clone()), + base_url.clone(), + ), + + // Codex OAuth + CredentialData::CodexOAuth { + creds_file_path, + api_base_url, + } => { + let token = self.get_codex_token(creds_file_path).await?; + ("codex".to_string(), Some(token), api_base_url.clone()) + } + + // Claude OAuth + CredentialData::ClaudeOAuth { creds_file_path } => { + let token = self.get_oauth_token(creds_file_path).await?; + ("anthropic".to_string(), Some(token), None) + } + + // Antigravity OAuth + CredentialData::AntigravityOAuth { + creds_file_path, .. + } => { + let token = self.get_oauth_token(creds_file_path).await?; + ("google".to_string(), Some(token), None) + } + }; + + Ok(AsterProviderConfig { + provider_name, + model_name: model.to_string(), + api_key, + base_url, + credential_uuid: credential.uuid.clone(), + }) + } + + /// 获取 Kiro OAuth Token + async fn get_kiro_token( + &self, + creds_path: &str, + _db: &DbConnection, + _uuid: &str, + ) -> Result { + use proxycast_providers::providers::KiroProvider; + + let mut provider = KiroProvider::new(); + provider + .load_credentials_from_path(creds_path) + .await + .map_err(|e| { + CredentialBridgeError::TokenRefreshFailed(format!("加载 Kiro 凭证失败: {e}")) + })?; + + // 检查 token 是否过期,如果过期则刷新 + if provider.is_token_expired() { + tracing::info!("[CredentialBridge] Kiro token 已过期,尝试刷新"); + self.pool_service + .refresh_kiro_token(creds_path) + .await + .map_err(CredentialBridgeError::TokenRefreshFailed)?; + + // 重新加载凭证 + provider + .load_credentials_from_path(creds_path) + .await + .map_err(|e| { + CredentialBridgeError::TokenRefreshFailed(format!("重新加载凭证失败: {e}")) + })?; + } + + provider.credentials.access_token.ok_or_else(|| { + CredentialBridgeError::TokenRefreshFailed("缺少 access_token".to_string()) + }) + } + + /// 获取通用 OAuth Token + async fn get_oauth_token(&self, creds_path: &str) -> Result { + let content = std::fs::read_to_string(creds_path).map_err(|e| { + CredentialBridgeError::TokenRefreshFailed(format!("读取凭证文件失败: {e}")) + })?; + + let creds: serde_json::Value = serde_json::from_str(&content) + .map_err(|e| CredentialBridgeError::TokenRefreshFailed(format!("解析凭证失败: {e}")))?; + + creds["access_token"] + .as_str() + .map(String::from) + .ok_or_else(|| { + CredentialBridgeError::TokenRefreshFailed("凭证中缺少 access_token".to_string()) + }) + } + + /// 获取 Codex OAuth Token + async fn get_codex_token(&self, creds_path: &str) -> Result { + use proxycast_providers::providers::CodexProvider; + + let mut provider = CodexProvider::new(); + provider + .load_credentials_from_path(creds_path) + .await + .map_err(|e| { + CredentialBridgeError::TokenRefreshFailed(format!("加载 Codex 凭证失败: {e}")) + })?; + + provider.ensure_valid_token().await.map_err(|e| { + CredentialBridgeError::TokenRefreshFailed(format!("获取 Codex token 失败: {e}")) + }) + } + + /// 记录凭证使用 + pub fn record_usage(&self, db: &DbConnection, uuid: &str) -> Result<(), CredentialBridgeError> { + self.pool_service + .record_usage(db, uuid) + .map_err(CredentialBridgeError::DatabaseError) + } + + /// 标记凭证为健康 + pub fn mark_healthy( + &self, + db: &DbConnection, + uuid: &str, + model: Option<&str>, + ) -> Result<(), CredentialBridgeError> { + self.pool_service + .mark_healthy(db, uuid, model) + .map_err(CredentialBridgeError::DatabaseError) + } + + /// 标记凭证为不健康 + pub fn mark_unhealthy( + &self, + db: &DbConnection, + uuid: &str, + error: Option<&str>, + ) -> Result<(), CredentialBridgeError> { + self.pool_service + .mark_unhealthy(db, uuid, error) + .map_err(CredentialBridgeError::DatabaseError) + } +} + +/// 从 AsterProviderConfig 创建 Aster Provider +/// +/// 设置环境变量并调用 aster::providers::create +pub async fn create_aster_provider( + config: &AsterProviderConfig, +) -> Result, CredentialBridgeError> { + // 设置环境变量 + set_provider_env_vars(config); + + // 创建 ModelConfig + let model_config = ModelConfig::new(&config.model_name).map_err(|e| { + CredentialBridgeError::ProviderCreationFailed(format!("创建 ModelConfig 失败: {e}")) + })?; + + // 创建 Provider + aster::providers::create(&config.provider_name, model_config) + .await + .map_err(|e| { + CredentialBridgeError::ProviderCreationFailed(format!("创建 Provider 失败: {e}")) + }) +} + +/// 设置 Provider 环境变量 +/// 从 URL 中拆分 host(scheme+authority)和 path 部分 +/// +/// 例如: +/// - `https://api.openai.com` -> (`https://api.openai.com`, ``) +/// - `https://open.bigmodel.cn/api/paas/v4` -> (`https://open.bigmodel.cn`, `api/paas/v4`) +/// - `https://localhost:8080/v1` -> (`https://localhost:8080`, `v1`) +fn split_url_host_and_path(url: &str) -> (String, String) { + // 找到 scheme 之后的 authority 部分 + let after_scheme = if let Some(pos) = url.find("://") { + pos + 3 + } else { + return (url.to_string(), String::new()); + }; + + // 找到 authority 之后的第一个 /(即路径开始) + let path_start = url[after_scheme..].find('/').map(|p| p + after_scheme); + + match path_start { + Some(pos) => { + let host = url[..pos].to_string(); + let path = url[pos..].trim_matches('/').to_string(); + (host, path) + } + None => (url.to_string(), String::new()), + } +} + +fn set_provider_env_vars(config: &AsterProviderConfig) { + tracing::info!( + "[CredentialBridge] set_provider_env_vars: provider_name={}, has_api_key={}, base_url={:?}", + config.provider_name, + config.api_key.is_some(), + config.base_url + ); + + let env_key = match config.provider_name.as_str() { + "openai" => "OPENAI_API_KEY", + "anthropic" => "ANTHROPIC_API_KEY", + "google" => "GOOGLE_API_KEY", + "bedrock" => "AWS_ACCESS_KEY_ID", // Bedrock 使用 AWS 凭证 + "gcpvertexai" => "GOOGLE_API_KEY", + "codex" => "OPENAI_API_KEY", // Codex 兼容 OpenAI + "deepseek" | "custom_deepseek" => "OPENAI_API_KEY", // DeepSeek 使用 OpenAI 兼容 API + "groq" => "OPENAI_API_KEY", // Groq 使用 OpenAI 兼容 API + "mistral" => "OPENAI_API_KEY", // Mistral 使用 OpenAI 兼容 API + "openrouter" => "OPENROUTER_API_KEY", + _ => "OPENAI_API_KEY", // 默认使用 OpenAI 格式 + }; + + tracing::info!("[CredentialBridge] 设置环境变量: {}=***", env_key); + + if let Some(api_key) = &config.api_key { + std::env::set_var(env_key, api_key); + } + + // 设置 base_url + // Aster 的 OpenAI Provider 使用 OPENAI_HOST(仅 scheme+host+port)和 + // OPENAI_BASE_PATH(路径部分 + /chat/completions)环境变量 + if let Some(base_url) = &config.base_url { + match config.provider_name.as_str() { + "openai" => { + // 解析 base_url,将路径部分拆分到 OPENAI_BASE_PATH + // 例如 https://open.bigmodel.cn/api/paas/v4 + // -> OPENAI_HOST = https://open.bigmodel.cn + // -> OPENAI_BASE_PATH = api/paas/v4/chat/completions + let (host_part, path_part) = split_url_host_and_path(base_url); + if path_part.is_empty() { + // 无路径部分(如 https://api.openai.com),直接设置 + std::env::set_var("OPENAI_HOST", base_url); + // 清除可能残留的 OPENAI_BASE_PATH,使用 Aster 默认值 + std::env::remove_var("OPENAI_BASE_PATH"); + tracing::info!("[CredentialBridge] 设置 OPENAI_HOST={}", base_url); + } else { + // base_url 包含路径,需要拆分 + let base_path = format!("{}/chat/completions", path_part); + std::env::set_var("OPENAI_HOST", &host_part); + std::env::set_var("OPENAI_BASE_PATH", &base_path); + tracing::info!( + "[CredentialBridge] 设置 OPENAI_HOST={}, OPENAI_BASE_PATH={}", + host_part, + base_path + ); + } + } + "anthropic" => { + std::env::set_var("ANTHROPIC_BASE_URL", base_url); + } + _ => { + // 其他 Provider 使用通用格式 + let base_url_key = format!( + "{}_BASE_URL", + config.provider_name.to_uppercase().replace('-', "_") + ); + std::env::set_var(&base_url_key, base_url); + } + } + } +} + +/// Provider 类型映射 +/// +/// 将 ProxyCast PoolProviderType 映射到 Aster Provider 名称 +pub fn map_pool_type_to_aster(pool_type: &PoolProviderType) -> &'static str { + match pool_type { + PoolProviderType::Kiro => "bedrock", + PoolProviderType::Gemini => "google", + PoolProviderType::Antigravity => "google", + PoolProviderType::OpenAI => "openai", + PoolProviderType::Claude => "anthropic", + PoolProviderType::Anthropic => "anthropic", + PoolProviderType::AnthropicCompatible => "anthropic", + PoolProviderType::Vertex => "gcpvertexai", + PoolProviderType::GeminiApiKey => "google", + PoolProviderType::Codex => "codex", + PoolProviderType::ClaudeOAuth => "anthropic", + PoolProviderType::AzureOpenai => "azure", + PoolProviderType::AwsBedrock => "bedrock", + PoolProviderType::Ollama => "ollama", + } +} + +/// 将 provider_type 字符串映射到 Aster Provider 名称 +/// +/// 支持 60+ API Key Provider,包括 deepseek, moonshot, qwen 等 +fn map_provider_type_to_aster(provider_type: &str) -> &'static str { + match provider_type { + // 标准 Provider + "openai" => "openai", + "anthropic" | "claude" => "anthropic", + "google" | "gemini" => "google", + "bedrock" | "kiro" => "bedrock", + "gcpvertexai" | "vertex" => "gcpvertexai", + "codex" => "codex", + "azure" | "azure-openai" => "azure", + "ollama" => "ollama", + + // DeepSeek - 使用 openai 兼容 provider(Aster 会通过 alias 映射) + "deepseek" | "custom_deepseek" => "openai", + + // 其他 OpenAI 兼容 Provider - 使用 openai provider + // 这些 Provider 都使用 OpenAI 兼容 API,通过 base_url 区分 + "groq" => "openai", + "mistral" => "openai", + "openrouter" => "openrouter", + + // 默认使用 openai(OpenAI 兼容格式) + _ => "openai", + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_map_pool_type_to_aster() { + assert_eq!(map_pool_type_to_aster(&PoolProviderType::OpenAI), "openai"); + assert_eq!( + map_pool_type_to_aster(&PoolProviderType::Claude), + "anthropic" + ); + assert_eq!(map_pool_type_to_aster(&PoolProviderType::Gemini), "google"); + assert_eq!(map_pool_type_to_aster(&PoolProviderType::Kiro), "bedrock"); + } + + #[test] + fn test_credential_bridge_error_display() { + let err = CredentialBridgeError::NoCredentials("test".to_string()); + assert!(err.to_string().contains("没有可用凭证")); + } + + #[test] + fn test_split_url_host_and_path() { + // 无路径 + let (host, path) = split_url_host_and_path("https://api.openai.com"); + assert_eq!(host, "https://api.openai.com"); + assert_eq!(path, ""); + + // 带路径(智谱) + let (host, path) = split_url_host_and_path("https://open.bigmodel.cn/api/paas/v4"); + assert_eq!(host, "https://open.bigmodel.cn"); + assert_eq!(path, "api/paas/v4"); + + // 带端口 + let (host, path) = split_url_host_and_path("https://localhost:8080/v1"); + assert_eq!(host, "https://localhost:8080"); + assert_eq!(path, "v1"); + + // 尾部斜杠 + let (host, path) = split_url_host_and_path("https://api.deepseek.com/v1/"); + assert_eq!(host, "https://api.deepseek.com"); + assert_eq!(path, "v1"); + + // 仅根路径 + let (host, path) = split_url_host_and_path("https://api.openai.com/"); + assert_eq!(host, "https://api.openai.com"); + assert_eq!(path, ""); + } +} diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs index 4094ba3b8..df4561e61 100644 --- a/src-tauri/crates/agent/src/lib.rs +++ b/src-tauri/crates/agent/src/lib.rs @@ -1,12 +1,30 @@ //! ProxyCast Agent Crate //! //! 包含 Agent 模块中不依赖主 crate 内部模块的纯逻辑部分。 -//! 深耦合部分(aster_state、aster_agent、credential_bridge、subagent_scheduler) -//! 留在主 crate。 +//! 深耦合部分(aster_state、aster_agent 流式桥接)留在主 crate。 +pub mod aster_state; +pub mod aster_state_support; +pub mod credential_bridge; pub mod event_converter; pub mod mcp_bridge; pub mod prompt; +pub mod session_store; +pub mod subagent_scheduler; +pub use aster_state::{AsterAgentState, ProviderConfig}; +pub use aster_state_support::{ + build_project_system_prompt, create_proxycast_identity, create_session_config_with_project, + message_helpers, reload_proxycast_skills, SessionConfigBuilder, +}; +pub use credential_bridge::{ + create_aster_provider, AsterProviderConfig, CredentialBridge, CredentialBridgeError, +}; pub use event_converter::{convert_agent_event, convert_to_tauri_message, TauriAgentEvent}; pub use prompt::SystemPromptBuilder; +pub use session_store::{ + create_session_sync, get_session_sync, list_sessions_sync, SessionDetail, SessionInfo, +}; +pub use subagent_scheduler::{ + ProxyCastScheduler, ProxyCastSubAgentExecutor, SchedulerEventEmitter, SubAgentProgressEvent, +}; diff --git a/src-tauri/crates/agent/src/mcp_bridge.rs b/src-tauri/crates/agent/src/mcp_bridge.rs index 305f885bf..275647b52 100644 --- a/src-tauri/crates/agent/src/mcp_bridge.rs +++ b/src-tauri/crates/agent/src/mcp_bridge.rs @@ -3,17 +3,10 @@ //! 实现 Aster 的 McpClientTrait,将工具调用转发到 //! ProxyCast 已有的 MCP RunningService,避免重复启动进程。 -use aster::agents::mcp_client::{Error, McpClientTrait}; -use rmcp::model::{ - CallToolResult, GetPromptResult, InitializeResult, JsonObject, ListPromptsResult, - ListResourcesResult, ListToolsResult, ReadResourceResult, ServerNotification, -}; +use rmcp::model::InitializeResult; use rmcp::service::RunningService; use rmcp::RoleClient; -use serde_json::Value; use std::sync::Arc; -use tokio::sync::{mpsc, Mutex}; -use tokio_util::sync::CancellationToken; use proxycast_mcp::client::ProxyCastMcpClient; @@ -21,6 +14,7 @@ use proxycast_mcp::client::ProxyCastMcpClient; /// /// 持有 ProxyCast 的 RunningService 引用, /// 将 Aster 的工具调用转发到已有的 MCP 连接。 +#[allow(dead_code)] pub struct McpBridgeClient { /// 服务器名称 name: String, diff --git a/src-tauri/crates/agent/src/session_store.rs b/src-tauri/crates/agent/src/session_store.rs new file mode 100644 index 000000000..b1b7c5731 --- /dev/null +++ b/src-tauri/crates/agent/src/session_store.rs @@ -0,0 +1,134 @@ +//! Agent 会话存储服务 +//! +//! 提供会话创建、列表查询、详情查询能力。 +//! 数据来源为 ProxyCast 数据库(AgentDao)。 + +use chrono::Utc; +use proxycast_core::agent::types::{AgentMessage, AgentSession, ContentPart, MessageContent}; +use proxycast_core::database::dao::agent::AgentDao; +use proxycast_core::database::DbConnection; +use uuid::Uuid; + +use crate::event_converter::{TauriMessage, TauriMessageContent}; + +/// 会话信息(简化版) +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +pub struct SessionInfo { + pub id: String, + pub name: String, + pub created_at: i64, + pub updated_at: i64, + pub messages_count: usize, +} + +/// 会话详情(包含消息) +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +pub struct SessionDetail { + pub id: String, + pub name: String, + pub created_at: i64, + pub updated_at: i64, + pub messages: Vec, +} + +/// 创建新会话 +pub fn create_session_sync(db: &DbConnection, name: Option) -> Result { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; + let session_name = name.unwrap_or_else(|| "新对话".to_string()); + let session_id = Uuid::new_v4().to_string(); + let now = Utc::now().to_rfc3339(); + + let session = AgentSession { + id: session_id.clone(), + model: "agent:default".to_string(), + messages: Vec::new(), + system_prompt: None, + title: Some(session_name), + created_at: now.clone(), + updated_at: now, + }; + + AgentDao::create_session(&conn, &session).map_err(|e| format!("创建会话失败: {e}"))?; + + Ok(session_id) +} + +/// 列出所有会话 +pub fn list_sessions_sync(db: &DbConnection) -> Result, String> { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; + + let sessions = AgentDao::list_sessions(&conn).map_err(|e| format!("获取会话列表失败: {e}"))?; + + Ok(sessions + .into_iter() + .map(|session| { + let messages_count = AgentDao::get_message_count(&conn, &session.id).unwrap_or(0); + SessionInfo { + id: session.id, + name: session.title.unwrap_or_else(|| "未命名".to_string()), + created_at: chrono::DateTime::parse_from_rfc3339(&session.created_at) + .map(|dt| dt.timestamp()) + .unwrap_or(0), + updated_at: chrono::DateTime::parse_from_rfc3339(&session.updated_at) + .map(|dt| dt.timestamp()) + .unwrap_or(0), + messages_count, + } + }) + .collect()) +} + +/// 获取会话详情 +pub fn get_session_sync(db: &DbConnection, session_id: &str) -> Result { + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; + + let session = AgentDao::get_session(&conn, session_id) + .map_err(|e| format!("获取会话失败: {e}"))? + .ok_or_else(|| format!("会话不存在: {session_id}"))?; + + let messages = + AgentDao::get_messages(&conn, session_id).map_err(|e| format!("获取消息失败: {e}"))?; + + Ok(SessionDetail { + id: session.id, + name: session.title.unwrap_or_else(|| "未命名".to_string()), + created_at: chrono::DateTime::parse_from_rfc3339(&session.created_at) + .map(|dt| dt.timestamp()) + .unwrap_or(0), + updated_at: chrono::DateTime::parse_from_rfc3339(&session.updated_at) + .map(|dt| dt.timestamp()) + .unwrap_or(0), + messages: messages + .into_iter() + .map(|message| convert_agent_message(&message)) + .collect(), + }) +} + +/// 将 AgentMessage 转换为 TauriMessage +fn convert_agent_message(message: &AgentMessage) -> TauriMessage { + let content = match &message.content { + MessageContent::Text(text) => vec![TauriMessageContent::Text { text: text.clone() }], + MessageContent::Parts(parts) => parts + .iter() + .filter_map(|part| { + if let ContentPart::Text { text } = part { + Some(TauriMessageContent::Text { text: text.clone() }) + } else { + None + } + }) + .collect(), + }; + + let timestamp = chrono::DateTime::parse_from_rfc3339(&message.timestamp) + .map(|dt| dt.timestamp()) + .unwrap_or(0); + + TauriMessage { + id: None, + role: message.role.clone(), + content, + timestamp, + } +} diff --git a/src-tauri/crates/agent/src/subagent_scheduler.rs b/src-tauri/crates/agent/src/subagent_scheduler.rs new file mode 100644 index 000000000..000ba48b6 --- /dev/null +++ b/src-tauri/crates/agent/src/subagent_scheduler.rs @@ -0,0 +1,246 @@ +//! SubAgent 调度器集成 +//! +//! 将 aster-rust 的 SubAgent 调度器与 ProxyCast 凭证池集成。 +//! 纯逻辑位于此 crate,事件发送通过注入回调实现。 + +use std::collections::HashMap; +use std::sync::Arc; +use std::time::Duration; + +use aster::agents::context::AgentContext; +use aster::agents::subagent_scheduler::{ + SchedulerConfig, SchedulerError, SchedulerExecutionResult, SchedulerProgress, SchedulerResult, + SubAgentExecutor, SubAgentResult, SubAgentScheduler, SubAgentTask, + TokenUsage as SchedulerTokenUsage, +}; +use aster::conversation::message::Message; +use chrono::Utc; +use tokio::sync::RwLock; +use tracing::{debug, info, warn}; + +use crate::credential_bridge::{create_aster_provider, AsterProviderConfig, CredentialBridge}; +use proxycast_core::database::DbConnection; + +/// 调度器事件发射器 +pub type SchedulerEventEmitter = Arc; + +/// ProxyCast SubAgent 执行器 +/// +/// 实现 aster-rust 的 SubAgentExecutor trait, +/// 集成 ProxyCast 凭证池进行 LLM 调用。 +pub struct ProxyCastSubAgentExecutor { + /// 凭证桥接器 + credential_bridge: CredentialBridge, + /// 数据库连接 + db: DbConnection, + /// 默认模型 + default_model: String, + /// 默认 Provider 类型 + default_provider: String, +} + +impl ProxyCastSubAgentExecutor { + /// 创建新的执行器 + pub fn new(db: DbConnection) -> Self { + Self { + credential_bridge: CredentialBridge::new(), + db, + default_model: "claude-sonnet-4-20250514".to_string(), + default_provider: "anthropic".to_string(), + } + } + + /// 设置默认模型 + pub fn with_default_model(mut self, model: impl Into) -> Self { + self.default_model = model.into(); + self + } + + /// 设置默认 Provider + pub fn with_default_provider(mut self, provider: impl Into) -> Self { + self.default_provider = provider.into(); + self + } + + /// 从凭证池选择凭证 + async fn select_credential(&self, task: &SubAgentTask) -> SchedulerResult { + let model = task.model.as_deref().unwrap_or(&self.default_model); + let provider_type = &self.default_provider; + + let config = self + .credential_bridge + .select_and_configure(&self.db, provider_type, model) + .await + .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; + + Ok(config) + } + + /// 生成摘要 + fn generate_summary(&self, output: &str, task: &SubAgentTask) -> String { + let max_len = 500; + if output.chars().count() <= max_len { + format!("任务 {} 完成:\n{}", task.id, output) + } else { + let truncated: String = output.chars().take(max_len - 3).collect(); + format!("任务 {} 完成:\n{}...", task.id, truncated) + } + } +} + +#[async_trait::async_trait] +impl SubAgentExecutor for ProxyCastSubAgentExecutor { + async fn execute_task( + &self, + task: &SubAgentTask, + context: &AgentContext, + ) -> SchedulerResult { + let start_time = Utc::now(); + info!("执行 SubAgent 任务: {}", task.id); + + let provider_config = self.select_credential(task).await?; + debug!("使用凭证: {}", provider_config.credential_uuid); + + let provider = create_aster_provider(&provider_config) + .await + .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; + + let system_prompt = context.system_prompt.clone().unwrap_or_default(); + let user_message = Message::user().with_text(&task.prompt); + + let (response_msg, usage) = provider + .complete(&system_prompt, &[user_message], &[]) + .await + .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; + + let response = response_msg.as_concat_text(); + + let end_time = Utc::now(); + let duration = (end_time - start_time).to_std().unwrap_or(Duration::ZERO); + + let summary = if task.return_summary { + Some(self.generate_summary(&response, task)) + } else { + None + }; + + let token_usage = Some(SchedulerTokenUsage { + input_tokens: usage.usage.input_tokens.unwrap_or(0) as usize, + output_tokens: usage.usage.output_tokens.unwrap_or(0) as usize, + total_tokens: usage.usage.total_tokens.unwrap_or(0) as usize, + }); + + Ok(SubAgentResult { + task_id: task.id.clone(), + success: true, + output: Some(response), + summary, + error: None, + duration, + retries: 0, + started_at: start_time, + completed_at: end_time, + token_usage, + metadata: HashMap::new(), + }) + } +} + +/// ProxyCast SubAgent 调度器 +pub struct ProxyCastScheduler { + /// 内部调度器 + scheduler: Arc>>>, + /// 数据库连接 + db: DbConnection, +} + +impl ProxyCastScheduler { + /// 创建新的调度器 + pub fn new(db: DbConnection) -> Self { + Self { + scheduler: Arc::new(RwLock::new(None)), + db, + } + } + + /// 初始化调度器(不附带事件回调) + pub async fn init(&self, config: Option) { + self.init_with_event_emitter(config, None).await; + } + + /// 初始化调度器(可附带事件回调) + pub async fn init_with_event_emitter( + &self, + config: Option, + event_emitter: Option, + ) { + let executor = ProxyCastSubAgentExecutor::new(self.db.clone()); + let config = config.unwrap_or_default(); + + let scheduler = if let Some(emitter) = event_emitter { + SubAgentScheduler::new(config, executor).with_event_callback(move |event| { + match serde_json::to_value(&event) { + Ok(payload) => emitter(&payload), + Err(err) => warn!("序列化调度事件失败: {}", err), + } + }) + } else { + SubAgentScheduler::new(config, executor) + }; + + *self.scheduler.write().await = Some(scheduler); + info!("ProxyCast SubAgent 调度器初始化完成"); + } + + /// 执行任务 + pub async fn execute( + &self, + tasks: Vec, + parent_context: Option<&AgentContext>, + ) -> SchedulerResult { + let scheduler = self.scheduler.read().await; + let scheduler = scheduler + .as_ref() + .ok_or_else(|| SchedulerError::ContextError("调度器未初始化".to_string()))?; + + scheduler.execute(tasks, parent_context).await + } + + /// 取消执行 + pub async fn cancel(&self) { + if let Some(scheduler) = self.scheduler.read().await.as_ref() { + scheduler.cancel().await; + } + } +} + +/// Tauri 事件:SubAgent 进度 +#[derive(Debug, Clone, serde::Serialize)] +#[serde(rename_all = "camelCase")] +pub struct SubAgentProgressEvent { + /// 总任务数 + pub total: usize, + /// 已完成数 + pub completed: usize, + /// 失败数 + pub failed: usize, + /// 运行中数 + pub running: usize, + /// 进度百分比 + pub percentage: f64, + /// 当前任务 + pub current_tasks: Vec, +} + +impl From for SubAgentProgressEvent { + fn from(progress: SchedulerProgress) -> Self { + Self { + total: progress.total, + completed: progress.completed, + failed: progress.failed, + running: progress.running, + percentage: progress.percentage, + current_tasks: progress.current_tasks, + } + } +} diff --git a/src-tauri/crates/core/src/app_bootstrap.rs b/src-tauri/crates/core/src/app_bootstrap.rs new file mode 100644 index 000000000..6fdbb0a46 --- /dev/null +++ b/src-tauri/crates/core/src/app_bootstrap.rs @@ -0,0 +1,68 @@ +//! 应用启动前的配置校验逻辑 +//! +//! 该模块仅包含与 Tauri 无关的纯配置处理。 + +use crate::app_utils::{generate_api_key, is_valid_bind_host}; +use crate::config::{self, Config}; + +/// 配置验证错误 +#[derive(Debug)] +pub enum ConfigError { + LoadFailed(String), + SaveFailed(String), + InvalidHost, + DefaultApiKeyWithNonLocalBind, + TlsNotSupported, + RemoteManagementNotSupported, +} + +impl std::fmt::Display for ConfigError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + ConfigError::LoadFailed(e) => write!(f, "配置加载失败: {e}"), + ConfigError::SaveFailed(e) => write!(f, "配置保存失败: {e}"), + ConfigError::InvalidHost => { + write!( + f, + "无效的监听地址。允许的地址:127.0.0.1、localhost、::1、0.0.0.0、::" + ) + } + ConfigError::DefaultApiKeyWithNonLocalBind => write!( + f, + "监听所有网络接口 (0.0.0.0 或 ::) 时,必须设置非默认的 API Key" + ), + ConfigError::TlsNotSupported => write!(f, "当前版本尚未支持 TLS"), + ConfigError::RemoteManagementNotSupported => { + write!(f, "远程管理需要 TLS 支持,当前版本未启用") + } + } + } +} + +impl std::error::Error for ConfigError {} + +/// 加载并验证配置 +pub fn load_and_validate_config() -> Result { + let mut config = config::load_config().map_err(|e| ConfigError::LoadFailed(e.to_string()))?; + + if !is_valid_bind_host(&config.server.host) { + return Err(ConfigError::InvalidHost); + } + + if config.server.api_key == config::DEFAULT_API_KEY { + let new_key = generate_api_key(); + config.server.api_key = new_key; + config::save_config(&config).map_err(|e| ConfigError::SaveFailed(e.to_string()))?; + tracing::info!("检测到默认 API key,已自动生成并保存新密钥"); + } + + if config.server.tls.enable { + return Err(ConfigError::TlsNotSupported); + } + + if config.remote_management.allow_remote { + return Err(ConfigError::RemoteManagementNotSupported); + } + + Ok(config) +} diff --git a/src-tauri/crates/core/src/app_utils.rs b/src-tauri/crates/core/src/app_utils.rs new file mode 100644 index 000000000..a11f4552c --- /dev/null +++ b/src-tauri/crates/core/src/app_utils.rs @@ -0,0 +1,120 @@ +//! 应用通用工具函数 +//! +//! 包含 API Key 生成、绑定地址校验、Token 掩码等与 Tauri 无关的逻辑。 + +use crate::config; + +/// 生成安全的 API Key +pub fn generate_api_key() -> String { + config::generate_secure_api_key() +} + +/// 检查是否为回环地址 +pub fn is_loopback_host(host: &str) -> bool { + if host == "localhost" { + return true; + } + match host.parse::() { + Ok(address) => address.is_loopback(), + Err(_) => false, + } +} + +/// 检查是否为有效的绑定地址 +/// 允许回环地址、0.0.0.0 和私有网络地址。 +pub fn is_valid_bind_host(host: &str) -> bool { + if is_loopback_host(host) { + return true; + } + + if host == "0.0.0.0" || host == "::" { + return true; + } + + if let Ok(address) = host.parse::() { + if let std::net::IpAddr::V4(ipv4) = address { + let octets = ipv4.octets(); + return octets[0] == 10 + || (octets[0] == 172 && (octets[1] >= 16 && octets[1] <= 31)) + || (octets[0] == 192 && octets[1] == 168); + } + } + + false +} + +/// 检查是否为非本地绑定地址(需要强 API Key) +pub fn is_non_local_bind(host: &str) -> bool { + if host == "0.0.0.0" || host == "::" { + return true; + } + + if let Ok(address) = host.parse::() { + if let std::net::IpAddr::V4(ipv4) = address { + let octets = ipv4.octets(); + return octets[0] == 10 + || (octets[0] == 172 && (octets[1] >= 16 && octets[1] <= 31)) + || (octets[0] == 192 && octets[1] == 168); + } + } + + false +} + +/// 掩码敏感 Token +pub fn mask_token(token: &str) -> String { + let chars: Vec = token.chars().collect(); + if chars.len() <= 12 { + "****".to_string() + } else { + let prefix: String = chars[..6].iter().collect(); + let suffix: String = chars[chars.len() - 4..].iter().collect(); + format!("{prefix}****{suffix}") + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_is_loopback_host() { + assert!(is_loopback_host("localhost")); + assert!(is_loopback_host("127.0.0.1")); + assert!(is_loopback_host("::1")); + assert!(!is_loopback_host("0.0.0.0")); + assert!(!is_loopback_host("192.168.1.1")); + } + + #[test] + fn test_is_valid_bind_host() { + assert!(is_valid_bind_host("localhost")); + assert!(is_valid_bind_host("127.0.0.1")); + assert!(is_valid_bind_host("::1")); + assert!(is_valid_bind_host("0.0.0.0")); + assert!(is_valid_bind_host("::")); + assert!(is_valid_bind_host("192.168.1.1")); + assert!(is_valid_bind_host("10.0.0.1")); + assert!(is_valid_bind_host("172.16.0.1")); + assert!(is_valid_bind_host("172.31.255.255")); + assert!(!is_valid_bind_host("8.8.8.8")); + assert!(!is_valid_bind_host("1.1.1.1")); + } + + #[test] + fn test_is_non_local_bind() { + assert!(is_non_local_bind("0.0.0.0")); + assert!(is_non_local_bind("::")); + assert!(!is_non_local_bind("127.0.0.1")); + assert!(!is_non_local_bind("localhost")); + assert!(is_non_local_bind("192.168.1.1")); + assert!(is_non_local_bind("10.0.0.1")); + assert!(is_non_local_bind("172.16.0.1")); + } + + #[test] + fn test_mask_token() { + assert_eq!(mask_token("short"), "****"); + assert_eq!(mask_token("abcdefghijklmnop"), "abcdef****mnop"); + } +} diff --git a/src-tauri/crates/core/src/lib.rs b/src-tauri/crates/core/src/lib.rs index f81d7a2ff..be89a26f5 100644 --- a/src-tauri/crates/core/src/lib.rs +++ b/src-tauri/crates/core/src/lib.rs @@ -3,6 +3,7 @@ //! 包含纯数据类型、基础模块和无外部业务依赖的独立模块。 //! //! ## 模块结构 +//! - `app_utils`: 应用通用工具函数 //! - `models`: 核心数据模型定义 //! - `data`: 静态数据 //! - `logger`: 日志配置 @@ -16,9 +17,14 @@ //! - `session`: 会话管理(限速、粘性路由) //! - `session_files`: 会话文件存储 +pub mod app_bootstrap; +pub mod app_utils; pub mod data; pub mod logger; pub mod models; +pub mod tray_format; +pub mod tray_menu_meta; +pub mod tray_state; // 独立业务模块(无主 crate 依赖) pub mod backends; diff --git a/src-tauri/crates/core/src/logger.rs b/src-tauri/crates/core/src/logger.rs index 1c56b9917..8d607ac68 100644 --- a/src-tauri/crates/core/src/logger.rs +++ b/src-tauri/crates/core/src/logger.rs @@ -1,4 +1,5 @@ //! 日志管理模块 +use crate::config::LoggingConfig; use chrono::{Duration, Local, Utc}; use regex::Regex; use serde::{Deserialize, Serialize}; @@ -238,6 +239,10 @@ impl LogStore { pub type SharedLogStore = Arc>; +pub fn create_log_store_from_config(logging: &LoggingConfig) -> LogStore { + LogStore::with_custom_config(logging.retention_days, logging.enabled) +} + /// P2 安全修复:扩展日志脱敏规则,覆盖更多敏感字段 pub fn sanitize_log_message(message: &str) -> String { let patterns = [ diff --git a/src-tauri/crates/core/src/tray_format.rs b/src-tauri/crates/core/src/tray_format.rs new file mode 100644 index 000000000..e3426560f --- /dev/null +++ b/src-tauri/crates/core/src/tray_format.rs @@ -0,0 +1,116 @@ +//! 菜单文本格式化模块 +//! +//! 提供托盘菜单文本的格式化函数 + +/// 格式化服务器状态文本 +/// +/// # 示例输出 +/// - "🟢 API 服务器: 运行中 (127.0.0.1:8080)" +/// - "⚪ API 服务器: 已停止" +pub fn format_server_status(running: bool, host: &str, port: u16) -> String { + if running { + format!("🟢 API 服务器: 运行中 ({host}:{port})") + } else { + "⚪ API 服务器: 已停止".to_string() + } +} + +/// 格式化凭证状态文本 +/// +/// # 示例输出 +/// - "🔑 可用凭证: 3/5" +pub fn format_credential_status(available: usize, total: usize) -> String { + format!("🔑 可用凭证: {available}/{total}") +} + +/// 格式化请求统计文本 +/// +/// # 示例输出 +/// - "📊 今日请求: 128 次" +pub fn format_request_count(count: u64) -> String { + format!("📊 今日请求: {count} 次") +} + +/// 格式化 API 地址 +/// +/// # 示例输出 +/// - "http://127.0.0.1:8080" +pub fn format_api_address(host: &str, port: u16) -> String { + format!("http://{host}:{port}") +} + +#[cfg(test)] +mod tests { + use super::*; + use proptest::prelude::*; + + proptest! { + /// **Feature: system-tray, Property 2: 菜单内容格式化正确性** + /// **Validates: Requirements 2.2, 2.3, 2.4** + #[test] + fn prop_menu_content_formatting( + host in "[a-z0-9.]{1,50}", + port in 1024u16..65535, + available in 0usize..100, + total in 0usize..100, + requests in 0u64..1000000 + ) { + let running_status = format_server_status(true, &host, port); + prop_assert!(running_status.contains(&host), "运行状态应包含 host"); + prop_assert!(running_status.contains(&port.to_string()), "运行状态应包含 port"); + prop_assert!(running_status.contains("运行中"), "运行状态应包含'运行中'"); + + let stopped_status = format_server_status(false, &host, port); + prop_assert!(stopped_status.contains("已停止"), "停止状态应包含'已停止'"); + + let cred_status = format_credential_status(available, total); + prop_assert!(cred_status.contains(&available.to_string()), "凭证状态应包含可用数"); + prop_assert!(cred_status.contains(&total.to_string()), "凭证状态应包含总数"); + + let req_status = format_request_count(requests); + prop_assert!(req_status.contains(&requests.to_string()), "请求统计应包含请求次数"); + } + + /// **Feature: system-tray, Property 4: API 地址格式化正确性** + /// **Validates: Requirements 4.2** + #[test] + fn prop_api_address_formatting( + host in "[a-z0-9.]{1,50}", + port in 1024u16..65535 + ) { + let address = format_api_address(&host, port); + let expected = format!("http://{host}:{port}"); + prop_assert_eq!(address, expected, "API 地址格式应为 http://{{host}}:{{port}}"); + } + } + + #[test] + fn test_format_server_status_running() { + let status = format_server_status(true, "127.0.0.1", 8080); + assert_eq!(status, "🟢 API 服务器: 运行中 (127.0.0.1:8080)"); + } + + #[test] + fn test_format_server_status_stopped() { + let status = format_server_status(false, "127.0.0.1", 8080); + assert_eq!(status, "⚪ API 服务器: 已停止"); + } + + #[test] + fn test_format_credential_status() { + let status = format_credential_status(3, 5); + assert_eq!(status, "🔑 可用凭证: 3/5"); + } + + #[test] + fn test_format_request_count() { + let status = format_request_count(128); + assert_eq!(status, "📊 今日请求: 128 次"); + } + + #[test] + fn test_format_api_address() { + let address = format_api_address("127.0.0.1", 8080); + assert_eq!(address, "http://127.0.0.1:8080"); + } +} diff --git a/src-tauri/crates/core/src/tray_menu_meta.rs b/src-tauri/crates/core/src/tray_menu_meta.rs new file mode 100644 index 000000000..5bd3386ad --- /dev/null +++ b/src-tauri/crates/core/src/tray_menu_meta.rs @@ -0,0 +1,201 @@ +//! 托盘菜单元数据模块 +//! +//! 提供与 Tauri 无关的菜单 ID、地址解析和菜单元信息工具函数。 + +/// 菜单项 ID 常量 +pub mod menu_ids { + /// 状态信息 + pub const STATUS_INFO: &str = "status_info"; + /// 凭证信息 + pub const CREDENTIAL_INFO: &str = "credential_info"; + /// 请求信息 + pub const REQUEST_INFO: &str = "request_info"; + /// 分隔符 1 + pub const SEPARATOR_1: &str = "sep_1"; + /// 启动服务器 + pub const START_SERVER: &str = "start_server"; + /// 停止服务器 + pub const STOP_SERVER: &str = "stop_server"; + /// 刷新所有 Token + pub const REFRESH_TOKENS: &str = "refresh_tokens"; + /// 健康检查 + pub const HEALTH_CHECK: &str = "health_check"; + /// 分隔符 2 + pub const SEPARATOR_2: &str = "sep_2"; + /// 打开主窗口 + pub const OPEN_WINDOW: &str = "open_window"; + /// 复制 API 地址 + pub const COPY_API_ADDRESS: &str = "copy_api_address"; + /// 打开日志目录 + pub const OPEN_LOG_DIR: &str = "open_log_dir"; + /// 分隔符 3 + pub const SEPARATOR_3: &str = "sep_3"; + /// 开机自启 + pub const AUTO_START: &str = "auto_start"; + /// 分隔符 4 + pub const SEPARATOR_4: &str = "sep_4"; + /// 退出 + pub const QUIT: &str = "quit"; + + /// 获取所有必需的菜单项 ID 列表 + pub fn all_required_ids() -> Vec<&'static str> { + vec![ + STATUS_INFO, + CREDENTIAL_INFO, + REQUEST_INFO, + START_SERVER, + STOP_SERVER, + REFRESH_TOKENS, + HEALTH_CHECK, + OPEN_WINDOW, + COPY_API_ADDRESS, + OPEN_LOG_DIR, + AUTO_START, + QUIT, + ] + } +} + +/// 解析服务器地址字符串为 host 和 port +/// +/// 支持格式: +/// - "host:port" -> (host, port) +/// - "host" -> (host, 8080) +/// - "" -> ("127.0.0.1", 8080) +pub fn parse_server_address(address: &str) -> (String, u16) { + if address.is_empty() { + return ("127.0.0.1".to_string(), 8080); + } + + if let Some((host, port_str)) = address.rsplit_once(':') { + if let Ok(port) = port_str.parse::() { + return (host.to_string(), port); + } + } + + (address.to_string(), 8080) +} + +/// 获取菜单中包含的所有菜单项 ID +pub fn get_menu_item_ids() -> Vec<&'static str> { + menu_ids::all_required_ids() +} + +#[cfg(test)] +mod tests { + use super::*; + use proptest::prelude::*; + + #[test] + fn test_all_required_ids_not_empty() { + let ids = menu_ids::all_required_ids(); + assert!(!ids.is_empty(), "必需的菜单项 ID 列表不应为空"); + } + + #[test] + fn test_all_required_ids_unique() { + let ids = menu_ids::all_required_ids(); + let mut unique_ids = ids.clone(); + unique_ids.sort(); + unique_ids.dedup(); + assert_eq!(ids.len(), unique_ids.len(), "菜单项 ID 应该唯一"); + } + + #[test] + fn test_menu_ids_completeness() { + let ids = menu_ids::all_required_ids(); + + assert!(ids.contains(&menu_ids::STATUS_INFO), "应包含 STATUS_INFO"); + assert!( + ids.contains(&menu_ids::CREDENTIAL_INFO), + "应包含 CREDENTIAL_INFO" + ); + assert!(ids.contains(&menu_ids::REQUEST_INFO), "应包含 REQUEST_INFO"); + assert!(ids.contains(&menu_ids::START_SERVER), "应包含 START_SERVER"); + assert!(ids.contains(&menu_ids::STOP_SERVER), "应包含 STOP_SERVER"); + assert!( + ids.contains(&menu_ids::REFRESH_TOKENS), + "应包含 REFRESH_TOKENS" + ); + assert!(ids.contains(&menu_ids::HEALTH_CHECK), "应包含 HEALTH_CHECK"); + assert!(ids.contains(&menu_ids::OPEN_WINDOW), "应包含 OPEN_WINDOW"); + assert!( + ids.contains(&menu_ids::COPY_API_ADDRESS), + "应包含 COPY_API_ADDRESS" + ); + assert!(ids.contains(&menu_ids::OPEN_LOG_DIR), "应包含 OPEN_LOG_DIR"); + assert!(ids.contains(&menu_ids::AUTO_START), "应包含 AUTO_START"); + assert!(ids.contains(&menu_ids::QUIT), "应包含 QUIT"); + } + + #[test] + fn test_parse_server_address_with_port() { + let (host, port) = parse_server_address("127.0.0.1:8080"); + assert_eq!(host, "127.0.0.1"); + assert_eq!(port, 8080); + } + + #[test] + fn test_parse_server_address_without_port() { + let (host, port) = parse_server_address("localhost"); + assert_eq!(host, "localhost"); + assert_eq!(port, 8080); + } + + #[test] + fn test_parse_server_address_empty() { + let (host, port) = parse_server_address(""); + assert_eq!(host, "127.0.0.1"); + assert_eq!(port, 8080); + } + + #[test] + fn test_parse_server_address_ipv6() { + let (host, port) = parse_server_address("[::1]:9000"); + assert_eq!(host, "[::1]"); + assert_eq!(port, 9000); + } + + #[test] + fn test_get_menu_item_ids() { + let ids = get_menu_item_ids(); + assert_eq!(ids.len(), 12, "应有 12 个必需的菜单项"); + } + + proptest! { + #[test] + fn prop_menu_ids_completeness( + _server_running in any::(), + _available in 0usize..100, + _total in 0usize..100, + _requests in 0u64..1000000, + _auto_start in any::() + ) { + let ids = menu_ids::all_required_ids(); + + let required = vec![ + menu_ids::STATUS_INFO, + menu_ids::CREDENTIAL_INFO, + menu_ids::REQUEST_INFO, + menu_ids::START_SERVER, + menu_ids::STOP_SERVER, + menu_ids::REFRESH_TOKENS, + menu_ids::HEALTH_CHECK, + menu_ids::OPEN_WINDOW, + menu_ids::COPY_API_ADDRESS, + menu_ids::OPEN_LOG_DIR, + menu_ids::AUTO_START, + menu_ids::QUIT, + ]; + + for id in required { + prop_assert!(ids.contains(&id), "菜单项列表应包含 {}", id); + } + + let mut sorted_ids = ids.clone(); + sorted_ids.sort(); + sorted_ids.dedup(); + prop_assert_eq!(ids.len(), sorted_ids.len(), "菜单项 ID 应该唯一"); + } + } +} diff --git a/src-tauri/crates/core/src/tray_state.rs b/src-tauri/crates/core/src/tray_state.rs new file mode 100644 index 000000000..74069a852 --- /dev/null +++ b/src-tauri/crates/core/src/tray_state.rs @@ -0,0 +1,288 @@ +//! 托盘状态模块 +//! +//! 定义托盘图标状态和状态快照结构 + +use serde::{Deserialize, Serialize}; + +/// 托盘图标状态枚举 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum TrayIconStatus { + /// 正常运行(绿色)- 服务器运行且凭证健康 + Running, + /// 警告状态(黄色)- 有凭证即将过期或余额不足 + Warning, + /// 错误状态(红色)- 服务器停止或所有凭证无效 + Error, + /// 停止状态(灰色)- 服务器未启动 + Stopped, +} + +impl Default for TrayIconStatus { + fn default() -> Self { + Self::Stopped + } +} + +/// 凭证健康状态 +#[derive(Debug, Clone, Default)] +pub struct CredentialHealth { + /// 凭证是否有效 + pub is_valid: bool, + /// 是否即将过期 + pub is_expiring_soon: bool, + /// 是否余额不足 + pub is_low_balance: bool, +} + +impl CredentialHealth { + /// 创建健康的凭证状态 + pub fn healthy() -> Self { + Self { + is_valid: true, + is_expiring_soon: false, + is_low_balance: false, + } + } + + /// 创建无效的凭证状态 + pub fn invalid() -> Self { + Self { + is_valid: false, + is_expiring_soon: false, + is_low_balance: false, + } + } + + /// 检查凭证是否有警告 + pub fn has_warning(&self) -> bool { + self.is_valid && (self.is_expiring_soon || self.is_low_balance) + } +} + +/// 托盘状态快照 +#[derive(Debug, Clone, Serialize)] +pub struct TrayStateSnapshot { + /// 图标状态 + pub icon_status: TrayIconStatus, + /// 服务器是否运行 + pub server_running: bool, + /// 服务器地址 + pub server_address: String, + /// 可用凭证数 + pub available_credentials: usize, + /// 总凭证数 + pub total_credentials: usize, + /// 今日请求数 + pub today_requests: u64, + /// 是否开机自启 + pub auto_start_enabled: bool, +} + +impl Default for TrayStateSnapshot { + fn default() -> Self { + Self { + icon_status: TrayIconStatus::Stopped, + server_running: false, + server_address: String::new(), + available_credentials: 0, + total_credentials: 0, + today_requests: 0, + auto_start_enabled: false, + } + } +} + +/// 根据服务器状态和凭证健康状态计算托盘图标状态 +/// +/// # 规则 +/// - 服务器未运行 -> Stopped +/// - 服务器运行 + 所有凭证无效 -> Error +/// - 服务器运行 + 有凭证警告 -> Warning +/// - 服务器运行 + 所有凭证健康 -> Running +pub fn calculate_icon_status( + server_running: bool, + credentials: &[CredentialHealth], +) -> TrayIconStatus { + if !server_running { + return TrayIconStatus::Stopped; + } + + if credentials.is_empty() { + return TrayIconStatus::Error; + } + + let all_invalid = credentials.iter().all(|c| !c.is_valid); + if all_invalid { + return TrayIconStatus::Error; + } + + let has_warning = credentials.iter().any(|c| c.has_warning()); + if has_warning { + return TrayIconStatus::Warning; + } + + TrayIconStatus::Running +} + +/// 将凭证池健康数据转换为托盘所需的健康状态 +pub fn get_credential_health_from_pool( + pool_credentials: &[(String, bool, bool, bool)], // (id, is_valid, is_expiring_soon, is_low_balance) +) -> Vec { + pool_credentials + .iter() + .map( + |(_, is_valid, is_expiring_soon, is_low_balance)| CredentialHealth { + is_valid: *is_valid, + is_expiring_soon: *is_expiring_soon, + is_low_balance: *is_low_balance, + }, + ) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + use proptest::prelude::*; + + fn arb_credential_health() -> impl Strategy { + (any::(), any::(), any::()).prop_map( + |(is_valid, is_expiring_soon, is_low_balance)| CredentialHealth { + is_valid, + is_expiring_soon, + is_low_balance, + }, + ) + } + + proptest! { + /// **Feature: system-tray, Property 1: 状态到图标映射正确性** + /// **Validates: Requirements 1.1, 1.2, 1.3** + #[test] + fn prop_icon_status_mapping( + server_running in any::(), + credentials in prop::collection::vec(arb_credential_health(), 0..10) + ) { + let status = calculate_icon_status(server_running, &credentials); + + if !server_running { + prop_assert_eq!(status, TrayIconStatus::Stopped); + return Ok(()); + } + + if credentials.is_empty() { + prop_assert_eq!(status, TrayIconStatus::Error); + return Ok(()); + } + + let all_invalid = credentials.iter().all(|c| !c.is_valid); + if all_invalid { + prop_assert_eq!(status, TrayIconStatus::Error); + return Ok(()); + } + + let has_warning = credentials.iter().any(|c| c.has_warning()); + if has_warning { + prop_assert_eq!(status, TrayIconStatus::Warning); + return Ok(()); + } + + prop_assert_eq!(status, TrayIconStatus::Running); + } + } + + #[test] + fn test_credential_health_healthy() { + let health = CredentialHealth::healthy(); + assert!(health.is_valid); + assert!(!health.is_expiring_soon); + assert!(!health.is_low_balance); + assert!(!health.has_warning()); + } + + #[test] + fn test_credential_health_invalid() { + let health = CredentialHealth::invalid(); + assert!(!health.is_valid); + assert!(!health.has_warning()); + } + + #[test] + fn test_credential_health_warning() { + let mut health = CredentialHealth::healthy(); + health.is_expiring_soon = true; + assert!(health.has_warning()); + + let mut health2 = CredentialHealth::healthy(); + health2.is_low_balance = true; + assert!(health2.has_warning()); + } + + #[test] + fn test_calculate_icon_status_server_stopped() { + let credentials = vec![CredentialHealth::healthy()]; + let status = calculate_icon_status(false, &credentials); + assert_eq!(status, TrayIconStatus::Stopped); + } + + #[test] + fn test_calculate_icon_status_no_credentials() { + let credentials = vec![]; + let status = calculate_icon_status(true, &credentials); + assert_eq!(status, TrayIconStatus::Error); + } + + #[test] + fn test_calculate_icon_status_all_invalid() { + let credentials = vec![CredentialHealth::invalid(), CredentialHealth::invalid()]; + let status = calculate_icon_status(true, &credentials); + assert_eq!(status, TrayIconStatus::Error); + } + + #[test] + fn test_calculate_icon_status_with_warning() { + let credentials = vec![ + CredentialHealth::healthy(), + CredentialHealth { + is_valid: true, + is_expiring_soon: true, + is_low_balance: false, + }, + ]; + let status = calculate_icon_status(true, &credentials); + assert_eq!(status, TrayIconStatus::Warning); + } + + #[test] + fn test_calculate_icon_status_running() { + let credentials = vec![CredentialHealth::healthy(), CredentialHealth::healthy()]; + let status = calculate_icon_status(true, &credentials); + assert_eq!(status, TrayIconStatus::Running); + } + + #[test] + fn test_get_credential_health_from_pool() { + let pool_data = vec![ + ("cred1".to_string(), true, false, false), + ("cred2".to_string(), true, true, false), + ("cred3".to_string(), false, false, false), + ]; + + let health = get_credential_health_from_pool(&pool_data); + + assert_eq!(health.len(), 3); + assert!(health[0].is_valid); + assert!(!health[0].is_expiring_soon); + assert!(health[1].is_valid); + assert!(health[1].is_expiring_soon); + assert!(!health[2].is_valid); + } + + #[test] + fn test_get_credential_health_empty() { + let pool_data: Vec<(String, bool, bool, bool)> = vec![]; + let health = get_credential_health_from_pool(&pool_data); + assert!(health.is_empty()); + } +} diff --git a/src-tauri/crates/credential/src/quota.rs b/src-tauri/crates/credential/src/quota.rs index bdbba5e95..a177e2520 100644 --- a/src-tauri/crates/credential/src/quota.rs +++ b/src-tauri/crates/credential/src/quota.rs @@ -166,14 +166,6 @@ impl QuotaManager { self.exceeded_credentials.len() } - /// 手动设置凭证的冷却结束时间(仅用于测试) - #[cfg(test)] - pub fn set_cooldown_until(&self, credential_id: &str, until: DateTime) { - if let Some(mut record) = self.exceeded_credentials.get_mut(credential_id) { - record.cooldown_until = until; - } - } - /// 检查是否为配额超限错误 pub fn is_quota_exceeded_error(status_code: Option, error_message: &str) -> bool { if let Some(code) = status_code { diff --git a/src-tauri/crates/server-utils/src/lib.rs b/src-tauri/crates/server-utils/src/lib.rs index 0814d3d67..ab8646f32 100644 --- a/src-tauri/crates/server-utils/src/lib.rs +++ b/src-tauri/crates/server-utils/src/lib.rs @@ -1005,3 +1005,54 @@ mod property_tests { } } } + +/// 解析 models index.json 的 provider_id 列表 +pub fn load_model_registry_provider_ids_from_resources() -> Result, String> { + let index_path = + resolve_models_index_path().ok_or_else(|| "未找到 models index.json".to_string())?; + + let index_content = std::fs::read_to_string(&index_path) + .map_err(|e| format!("读取 models index.json 失败 ({index_path:?}): {e}"))?; + + let index_json = serde_json::from_str::(&index_content) + .map_err(|e| format!("解析 models index.json 失败: {e}"))?; + + let providers = index_json + .get("providers") + .and_then(|v| v.as_array()) + .ok_or_else(|| "models index.json 缺少 providers 数组".to_string())?; + + let mut provider_ids: Vec = providers + .iter() + .filter_map(|v| v.as_str()) + .map(|id| id.trim().to_string()) + .filter(|id| !id.is_empty()) + .collect(); + + provider_ids.sort(); + provider_ids.dedup(); + Ok(provider_ids) +} + +/// 定位 models index.json 路径 +pub fn resolve_models_index_path() -> Option { + let mut candidates: Vec = Vec::new(); + + if let Ok(current_dir) = std::env::current_dir() { + candidates.push(current_dir.join("src-tauri/resources/models/index.json")); + candidates.push(current_dir.join("resources/models/index.json")); + } + + if let Ok(exe_path) = std::env::current_exe() { + if let Some(parent) = exe_path.parent() { + candidates.push(parent.join("resources/models/index.json")); + candidates.push(parent.join("../../src-tauri/resources/models/index.json")); + candidates.push(parent.join("../../../src-tauri/resources/models/index.json")); + candidates.push(parent.join("../Resources/resources/models/index.json")); + candidates.push(parent.join("../../Resources/resources/models/index.json")); + candidates.push(parent.join("../../../Resources/resources/models/index.json")); + } + } + + candidates.into_iter().find(|path| path.exists()) +} diff --git a/src-tauri/crates/server/src/handlers/mod.rs b/src-tauri/crates/server/src/handlers/mod.rs index 996e0f029..e2e437859 100644 --- a/src-tauri/crates/server/src/handlers/mod.rs +++ b/src-tauri/crates/server/src/handlers/mod.rs @@ -13,7 +13,12 @@ pub mod websocket; pub use api::*; pub use credentials_api::*; pub use image_handler::*; -pub use kiro_credential::*; +// 避免 SelectCredentialRequest 歧义 glob re-export(credentials_api 和 kiro_credential 都定义了同名类型) +pub use kiro_credential::{ + get_available_credentials, get_credential_status, refresh_credential, select_credential, + AvailableCredential, AvailableCredentialsResponse, RefreshCredentialResponse, + SelectCredentialResponse, +}; pub use management::*; pub use provider_calls::*; pub use websocket::*; diff --git a/src-tauri/crates/services/Cargo.toml b/src-tauri/crates/services/Cargo.toml index 14d926eea..bbaca607d 100644 --- a/src-tauri/crates/services/Cargo.toml +++ b/src-tauri/crates/services/Cargo.toml @@ -5,10 +5,15 @@ edition.workspace = true authors.workspace = true repository.workspace = true +[features] +default = [] +local-whisper = ["voice-core/local-whisper"] + [dependencies] # 项目内 crate proxycast-core.workspace = true proxycast-providers.workspace = true +voice-core.workspace = true # 序列化 serde.workspace = true diff --git a/src-tauri/crates/services/src/file_browser_service.rs b/src-tauri/crates/services/src/file_browser_service.rs new file mode 100644 index 000000000..f60c33220 --- /dev/null +++ b/src-tauri/crates/services/src/file_browser_service.rs @@ -0,0 +1,730 @@ +//! 文件浏览器服务 +//! +//! 提供文件系统浏览功能 +//! 支持目录列表、文件预览等操作 +//! +//! # 功能 +//! - 列出目录内容 +//! - 读取文件预览 +//! - 获取文件元信息 +//! - 获取文件权限和 MIME 类型 + +use serde::{Deserialize, Serialize}; +use std::fs::{self, Metadata}; +#[cfg(unix)] +use std::os::unix::fs::PermissionsExt; +use std::path::{Path, PathBuf}; +use std::time::UNIX_EPOCH; +use tracing::{debug, error}; + +/// 文件条目 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FileEntry { + /// 文件名 + pub name: String, + /// 完整路径 + pub path: String, + /// 是否为目录 + #[serde(rename = "isDir")] + pub is_dir: bool, + /// 文件大小(字节) + pub size: u64, + /// 修改时间(Unix 时间戳毫秒) + #[serde(rename = "modifiedAt")] + pub modified_at: u64, + /// 文件类型/扩展名 + #[serde(rename = "fileType")] + pub file_type: Option, + /// 是否隐藏文件 + #[serde(rename = "isHidden")] + pub is_hidden: bool, + /// 文件权限字符串(如 -rw-r--r--) + #[serde(rename = "modeStr")] + pub mode_str: Option, + /// 文件权限数字(8进制) + pub mode: Option, + /// MIME 类型 + #[serde(rename = "mimeType")] + pub mime_type: Option, + /// 是否为符号链接 + #[serde(rename = "isSymlink")] + pub is_symlink: bool, +} + +/// 目录列表结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct DirectoryListing { + /// 当前路径 + pub path: String, + /// 父目录路径 + #[serde(rename = "parentPath")] + pub parent_path: Option, + /// 文件列表 + pub entries: Vec, + /// 错误信息 + pub error: Option, +} + +/// 文件预览结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FilePreview { + /// 文件路径 + pub path: String, + /// 文件内容(文本) + pub content: Option, + /// 是否为二进制文件 + #[serde(rename = "isBinary")] + pub is_binary: bool, + /// 文件大小 + pub size: u64, + /// 错误信息 + pub error: Option, +} + +/// 获取文件扩展名 +fn get_file_extension(path: &Path) -> Option { + path.extension() + .and_then(|ext| ext.to_str()) + .map(|s| s.to_lowercase()) +} + +/// 判断是否为隐藏文件 +fn is_hidden_file(name: &str) -> bool { + name.starts_with('.') +} + +/// 将 Unix 文件模式转换为权限字符串(如 -rw-r--r--) +#[cfg(unix)] +fn mode_to_string(mode: u32, is_dir: bool, is_symlink: bool) -> String { + let mut result = String::with_capacity(10); + + // 文件类型 + if is_symlink { + result.push('l'); + } else if is_dir { + result.push('d'); + } else { + result.push('-'); + } + + // 用户权限 + result.push(if mode & 0o400 != 0 { 'r' } else { '-' }); + result.push(if mode & 0o200 != 0 { 'w' } else { '-' }); + result.push(if mode & 0o100 != 0 { 'x' } else { '-' }); + + // 组权限 + result.push(if mode & 0o040 != 0 { 'r' } else { '-' }); + result.push(if mode & 0o020 != 0 { 'w' } else { '-' }); + result.push(if mode & 0o010 != 0 { 'x' } else { '-' }); + + // 其他用户权限 + result.push(if mode & 0o004 != 0 { 'r' } else { '-' }); + result.push(if mode & 0o002 != 0 { 'w' } else { '-' }); + result.push(if mode & 0o001 != 0 { 'x' } else { '-' }); + + result +} + +/// 根据文件扩展名和元数据获取 MIME 类型 +fn get_mime_type(path: &Path, metadata: &Metadata) -> String { + // 特殊类型检测 + if metadata.is_dir() { + return "directory".to_string(); + } + + #[cfg(unix)] + { + use std::os::unix::fs::FileTypeExt; + let ft = metadata.file_type(); + if ft.is_symlink() { + return "symlink".to_string(); + } + if ft.is_block_device() { + return "block-device".to_string(); + } + if ft.is_char_device() { + return "char-device".to_string(); + } + if ft.is_fifo() { + return "pipe".to_string(); + } + if ft.is_socket() { + return "socket".to_string(); + } + } + + // 基于扩展名的 MIME 类型映射 + let ext = path + .extension() + .and_then(|e| e.to_str()) + .map(|s| s.to_lowercase()); + + match ext.as_deref() { + // 文本文件 + Some("txt") => "text/plain", + Some("md" | "markdown") => "text/markdown", + Some("html" | "htm") => "text/html", + Some("css") => "text/css", + Some("xml") => "text/xml", + Some("csv") => "text/csv", + + // 代码文件 + Some("js" | "mjs" | "cjs") => "text/javascript", + Some("ts" | "mts" | "cts") => "text/typescript", + Some("tsx") => "text/tsx", + Some("jsx") => "text/jsx", + Some("json") => "application/json", + Some("yaml" | "yml") => "text/yaml", + Some("toml") => "text/toml", + Some("rs") => "text/x-rust", + Some("py") => "text/x-python", + Some("go") => "text/x-go", + Some("java") => "text/x-java", + Some("c") => "text/x-c", + Some("cpp" | "cc" | "cxx") => "text/x-c++", + Some("h" | "hpp") => "text/x-c-header", + Some("sh" | "bash" | "zsh") => "text/x-shellscript", + Some("sql") => "text/x-sql", + Some("vue") => "text/x-vue", + Some("svelte") => "text/x-svelte", + Some("swift") => "text/x-swift", + Some("kt" | "kts") => "text/x-kotlin", + Some("rb") => "text/x-ruby", + Some("php") => "text/x-php", + Some("lua") => "text/x-lua", + + // 图片 + Some("png") => "image/png", + Some("jpg" | "jpeg") => "image/jpeg", + Some("gif") => "image/gif", + Some("webp") => "image/webp", + Some("svg") => "image/svg+xml", + Some("ico") => "image/x-icon", + Some("bmp") => "image/bmp", + Some("tiff" | "tif") => "image/tiff", + + // 音频 + Some("mp3") => "audio/mpeg", + Some("wav") => "audio/wav", + Some("ogg") => "audio/ogg", + Some("flac") => "audio/flac", + Some("aac") => "audio/aac", + Some("m4a") => "audio/mp4", + + // 视频 + Some("mp4") => "video/mp4", + Some("webm") => "video/webm", + Some("avi") => "video/x-msvideo", + Some("mov") => "video/quicktime", + Some("mkv") => "video/x-matroska", + Some("wmv") => "video/x-ms-wmv", + + // 压缩文件 + Some("zip") => "application/zip", + Some("tar") => "application/x-tar", + Some("gz" | "gzip") => "application/gzip", + Some("bz2") => "application/x-bzip2", + Some("xz") => "application/x-xz", + Some("7z") => "application/x-7z-compressed", + Some("rar") => "application/vnd.rar", + + // 文档 + Some("pdf") => "application/pdf", + Some("doc") => "application/msword", + Some("docx") => "application/vnd.openxmlformats-officedocument.wordprocessingml.document", + Some("xls") => "application/vnd.ms-excel", + Some("xlsx") => "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", + Some("ppt") => "application/vnd.ms-powerpoint", + Some("pptx") => "application/vnd.openxmlformats-officedocument.presentationml.presentation", + + // 可执行文件 + Some("exe") => "application/x-msdownload", + Some("dmg") => "application/x-apple-diskimage", + Some("app") => "application/x-apple-application", + Some("deb") => "application/x-debian-package", + Some("rpm") => "application/x-rpm", + + // 字体 + Some("ttf") => "font/ttf", + Some("otf") => "font/otf", + Some("woff") => "font/woff", + Some("woff2") => "font/woff2", + + // 其他 + Some("wasm") => "application/wasm", + + _ => "application/octet-stream", + } + .to_string() +} + +/// 判断是否为文本文件(基于扩展名) +fn is_text_file(extension: Option<&str>) -> bool { + match extension { + Some(ext) => matches!( + ext, + "txt" + | "md" + | "json" + | "yaml" + | "yml" + | "toml" + | "xml" + | "html" + | "htm" + | "css" + | "js" + | "ts" + | "tsx" + | "jsx" + | "rs" + | "py" + | "go" + | "java" + | "c" + | "cpp" + | "h" + | "hpp" + | "sh" + | "bash" + | "zsh" + | "fish" + | "sql" + | "graphql" + | "vue" + | "svelte" + | "astro" + | "log" + | "env" + | "gitignore" + | "dockerignore" + | "editorconfig" + | "prettierrc" + | "eslintrc" + | "babelrc" + | "conf" + | "cfg" + | "ini" + | "properties" + ), + None => false, + } +} + +/// 列出目录内容 +pub fn list_directory(path: &str) -> DirectoryListing { + let path_buf = if path.is_empty() || path == "~" { + dirs::home_dir().unwrap_or_else(|| PathBuf::from("/")) + } else if path.starts_with('~') { + let home = dirs::home_dir().unwrap_or_else(|| PathBuf::from("/")); + home.join(&path[2..]) + } else { + PathBuf::from(path) + }; + + let canonical_path = match path_buf.canonicalize() { + Ok(p) => p, + Err(e) => { + error!("无法解析路径 {}: {}", path, e); + return DirectoryListing { + path: path.to_string(), + parent_path: None, + entries: vec![], + error: Some(format!("无法解析路径: {e}")), + }; + } + }; + + let parent_path = canonical_path + .parent() + .map(|p| p.to_string_lossy().to_string()); + + let entries = match fs::read_dir(&canonical_path) { + Ok(read_dir) => { + let mut entries: Vec = read_dir + .filter_map(|entry| { + let entry = entry.ok()?; + // 先获取符号链接信息 + let symlink_metadata = entry.metadata().ok(); + let is_symlink = entry + .file_type() + .ok() + .map(|ft| ft.is_symlink()) + .unwrap_or(false); + // 获取真实文件的元数据(解析符号链接) + let metadata = if is_symlink { + fs::metadata(entry.path()).ok().or(symlink_metadata)? + } else { + symlink_metadata? + }; + let name = entry.file_name().to_string_lossy().to_string(); + let path = entry.path(); + + let modified_at = metadata + .modified() + .ok() + .and_then(|t| t.duration_since(UNIX_EPOCH).ok()) + .map(|d| d.as_millis() as u64) + .unwrap_or(0); + + let file_type = if metadata.is_dir() { + Some("folder".to_string()) + } else { + get_file_extension(&path) + }; + + // 获取文件权限(仅 Unix) + #[cfg(unix)] + let (mode, mode_str) = { + let m = metadata.permissions().mode(); + ( + Some(m & 0o777), + Some(mode_to_string(m & 0o777, metadata.is_dir(), is_symlink)), + ) + }; + #[cfg(not(unix))] + let (mode, mode_str): (Option, Option) = (None, None); + + // 获取 MIME 类型 + let mime_type = get_mime_type(&path, &metadata); + + Some(FileEntry { + name: name.clone(), + path: path.to_string_lossy().to_string(), + is_dir: metadata.is_dir(), + size: metadata.len(), + modified_at, + file_type, + is_hidden: is_hidden_file(&name), + mode_str, + mode, + mime_type: Some(mime_type), + is_symlink, + }) + }) + .collect(); + + // 排序:目录在前,然后按名称排序 + entries.sort_by(|a, b| match (a.is_dir, b.is_dir) { + (true, false) => std::cmp::Ordering::Less, + (false, true) => std::cmp::Ordering::Greater, + _ => a.name.to_lowercase().cmp(&b.name.to_lowercase()), + }); + + entries + } + Err(e) => { + error!("无法读取目录 {}: {}", canonical_path.display(), e); + return DirectoryListing { + path: canonical_path.to_string_lossy().to_string(), + parent_path, + entries: vec![], + error: Some(format!("无法读取目录: {e}")), + }; + } + }; + + debug!( + "列出目录 {}: {} 个条目", + canonical_path.display(), + entries.len() + ); + + DirectoryListing { + path: canonical_path.to_string_lossy().to_string(), + parent_path, + entries, + error: None, + } +} + +/// 读取文件预览 +pub fn read_file_preview(path: &str, max_size: Option) -> FilePreview { + let max_size = max_size.unwrap_or(100 * 1024); // 默认 100KB + let path_buf = PathBuf::from(path); + + let metadata = match fs::metadata(&path_buf) { + Ok(m) => m, + Err(e) => { + return FilePreview { + path: path.to_string(), + content: None, + is_binary: false, + size: 0, + error: Some(format!("无法读取文件元信息: {e}")), + }; + } + }; + + if metadata.is_dir() { + return FilePreview { + path: path.to_string(), + content: None, + is_binary: false, + size: 0, + error: Some("不能预览目录".to_string()), + }; + } + + let size = metadata.len(); + let extension = get_file_extension(&path_buf); + let is_text = is_text_file(extension.as_deref()); + + if !is_text { + return FilePreview { + path: path.to_string(), + content: None, + is_binary: true, + size, + error: None, + }; + } + + // 读取文件内容 + let content = match fs::read(&path_buf) { + Ok(bytes) => { + let bytes_to_read = bytes.len().min(max_size); + match String::from_utf8(bytes[..bytes_to_read].to_vec()) { + Ok(s) => Some(s), + Err(_) => { + return FilePreview { + path: path.to_string(), + content: None, + is_binary: true, + size, + error: None, + }; + } + } + } + Err(e) => { + return FilePreview { + path: path.to_string(), + content: None, + is_binary: false, + size, + error: Some(format!("无法读取文件: {e}")), + }; + } + }; + + FilePreview { + path: path.to_string(), + content, + is_binary: false, + size, + error: None, + } +} + +/// 服务接口:列出目录 +pub async fn list_dir(path: String) -> Result { + Ok(list_directory(&path)) +} + +/// 服务接口:读取文件预览 +pub async fn read_file_preview_cmd( + path: String, + max_size: Option, +) -> Result { + Ok(read_file_preview(&path, max_size)) +} + +/// 服务接口:获取用户主目录 +pub async fn get_home_dir() -> Result { + dirs::home_dir() + .map(|p| p.to_string_lossy().to_string()) + .ok_or_else(|| "无法获取主目录".to_string()) +} + +/// 服务接口:创建新文件 +pub async fn create_file(path: String) -> Result<(), String> { + let path_buf = PathBuf::from(&path); + + // 检查文件是否已存在 + if path_buf.exists() { + return Err("文件已存在".to_string()); + } + + // 确保父目录存在 + if let Some(parent) = path_buf.parent() { + if !parent.exists() { + fs::create_dir_all(parent).map_err(|e| format!("无法创建父目录: {e}"))?; + } + } + + // 创建空文件 + fs::File::create(&path_buf).map_err(|e| format!("无法创建文件: {e}"))?; + + debug!("创建文件: {}", path); + Ok(()) +} + +/// 服务接口:创建新目录 +pub async fn create_directory(path: String) -> Result<(), String> { + let path_buf = PathBuf::from(&path); + + // 检查目录是否已存在 + if path_buf.exists() { + return Err("目录已存在".to_string()); + } + + fs::create_dir_all(&path_buf).map_err(|e| format!("无法创建目录: {e}"))?; + + debug!("创建目录: {}", path); + Ok(()) +} + +/// 服务接口:删除文件或目录 +pub async fn delete_file(path: String, recursive: bool) -> Result<(), String> { + let path_buf = PathBuf::from(&path); + + if !path_buf.exists() { + return Err("文件或目录不存在".to_string()); + } + + if path_buf.is_dir() { + if recursive { + fs::remove_dir_all(&path_buf).map_err(|e| format!("无法删除目录: {e}"))?; + } else { + fs::remove_dir(&path_buf) + .map_err(|e| format!("无法删除目录(目录非空,需要递归删除): {e}"))?; + } + debug!("删除目录: {}", path); + } else { + fs::remove_file(&path_buf).map_err(|e| format!("无法删除文件: {e}"))?; + debug!("删除文件: {}", path); + } + + Ok(()) +} + +/// 服务接口:重命名文件或目录 +pub async fn rename_file(old_path: String, new_path: String) -> Result<(), String> { + let old_path_buf = PathBuf::from(&old_path); + let new_path_buf = PathBuf::from(&new_path); + + if !old_path_buf.exists() { + return Err("源文件或目录不存在".to_string()); + } + + if new_path_buf.exists() { + return Err("目标文件或目录已存在".to_string()); + } + + fs::rename(&old_path_buf, &new_path_buf).map_err(|e| format!("无法重命名: {e}"))?; + + debug!("重命名: {} -> {}", old_path, new_path); + Ok(()) +} + +/// 服务接口:复制文件名到剪贴板(返回文件名供前端处理) +pub async fn get_file_name(path: String) -> Result { + let path_buf = PathBuf::from(&path); + path_buf + .file_name() + .and_then(|n| n.to_str()) + .map(|s| s.to_string()) + .ok_or_else(|| "无法获取文件名".to_string()) +} + +/// 服务接口:在 Finder 中显示文件 +pub async fn reveal_in_finder(path: String) -> Result<(), String> { + let path_buf = PathBuf::from(&path); + + if !path_buf.exists() { + return Err("文件或目录不存在".to_string()); + } + + #[cfg(target_os = "macos")] + { + std::process::Command::new("open") + .args(["-R", &path]) + .spawn() + .map_err(|e| format!("无法打开 Finder: {e}"))?; + } + + #[cfg(target_os = "windows")] + { + std::process::Command::new("explorer") + .args(["/select,", &path]) + .spawn() + .map_err(|e| format!("无法打开资源管理器: {}", e))?; + } + + #[cfg(target_os = "linux")] + { + // 尝试使用 xdg-open 打开父目录 + let parent = path_buf.parent().unwrap_or(&path_buf); + std::process::Command::new("xdg-open") + .arg(parent) + .spawn() + .map_err(|e| format!("无法打开文件管理器: {}", e))?; + } + + Ok(()) +} + +/// 服务接口:使用默认应用打开文件 +pub async fn open_with_default_app(path: String) -> Result<(), String> { + let path_buf = PathBuf::from(&path); + + if !path_buf.exists() { + return Err("文件不存在".to_string()); + } + + #[cfg(target_os = "macos")] + { + std::process::Command::new("open") + .arg(&path) + .spawn() + .map_err(|e| format!("无法打开文件: {e}"))?; + } + + #[cfg(target_os = "windows")] + { + std::process::Command::new("cmd") + .args(["/C", "start", "", &path]) + .spawn() + .map_err(|e| format!("无法打开文件: {}", e))?; + } + + #[cfg(target_os = "linux")] + { + std::process::Command::new("xdg-open") + .arg(&path) + .spawn() + .map_err(|e| format!("无法打开文件: {}", e))?; + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_list_home_directory() { + let result = list_directory("~"); + assert!(result.error.is_none()); + assert!(!result.entries.is_empty()); + } + + #[test] + fn test_is_hidden_file() { + assert!(is_hidden_file(".gitignore")); + assert!(is_hidden_file(".config")); + assert!(!is_hidden_file("readme.md")); + } + + #[test] + fn test_is_text_file() { + assert!(is_text_file(Some("txt"))); + assert!(is_text_file(Some("rs"))); + assert!(is_text_file(Some("json"))); + assert!(!is_text_file(Some("png"))); + assert!(!is_text_file(Some("exe"))); + assert!(!is_text_file(None)); + } +} diff --git a/src-tauri/crates/services/src/lib.rs b/src-tauri/crates/services/src/lib.rs index ee4160962..4078ae6b5 100644 --- a/src-tauri/crates/services/src/lib.rs +++ b/src-tauri/crates/services/src/lib.rs @@ -7,8 +7,15 @@ //! - `file_browser_service` - 文件浏览服务 //! - `sysinfo_service` - 系统信息服务 //! - `update_check_service` - 更新检查服务 -//! - `update_window` - 更新窗口 //! - `usage_service` - 使用统计服务 +//! - `voice_config_service` - 语音配置服务 +//! - `voice_processor_service` - 语音润色服务 +//! - `voice_output_service` - 语音输出服务 +//! - `voice_asr_service` - ASR 识别服务 +//! - `voice_command_service` - 语音命令业务服务 +//! - `voice_recording_service` - 录音状态与设备服务 +//! - `screenshot_capture_service` - 跨平台截图服务 +//! - `screenshot_image_service` - 截图图片编码服务 //! - `machine_id_service` - 机器 ID 服务 //! - `live_sync` - 实时同步 //! - `mcp_sync` - MCP 同步 @@ -36,8 +43,18 @@ // 无外部依赖的服务 pub mod context_memory_service; +pub mod file_browser_service; +pub mod screenshot_capture_service; +pub mod screenshot_image_service; +pub mod sysinfo_service; pub mod update_check_service; pub mod usage_service; +pub mod voice_asr_service; +pub mod voice_command_service; +pub mod voice_config_service; +pub mod voice_output_service; +pub mod voice_processor_service; +pub mod voice_recording_service; // 依赖 models 的服务 pub mod live_sync; diff --git a/src-tauri/crates/services/src/screenshot_capture_service.rs b/src-tauri/crates/services/src/screenshot_capture_service.rs new file mode 100644 index 000000000..4f746a64f --- /dev/null +++ b/src-tauri/crates/services/src/screenshot_capture_service.rs @@ -0,0 +1,244 @@ +//! 屏幕截图服务 +//! +//! 提供跨平台的屏幕截图功能,支持交互式区域选择 + +use std::path::PathBuf; +use tracing::{debug, error, info}; + +#[cfg(target_os = "macos")] +use tracing::warn; + +/// 截图错误类型 +#[derive(Debug, thiserror::Error)] +pub enum CaptureError { + #[error("用户取消截图")] + Cancelled, + #[error("权限不足: {0}")] + PermissionDenied(String), + #[error("系统错误: {0}")] + SystemError(String), + #[error("IO 错误: {0}")] + IoError(#[from] std::io::Error), + #[error("临时文件创建失败: {0}")] + TempFileError(String), +} + +/// 截图结果 +#[derive(Debug)] +pub enum CaptureResult { + /// 截图成功,包含图片路径 + Success(PathBuf), + /// 用户取消截图 + Cancelled, +} + +/// 启动交互式截图 +/// +/// 调用系统截图工具进行交互式区域选择 +/// +/// # 返回 +/// 成功返回图片路径,用户取消返回 Cancelled 错误,其他情况返回相应错误 +/// +/// # 平台支持 +/// - macOS: 使用 `screencapture -i -x` 命令 +/// - Windows: 使用 Windows API 或系统截图工具 +/// - Linux: 使用 `gnome-screenshot` 或 `scrot` +pub async fn start_capture() -> Result { + info!("启动交互式截图"); + + // 生成临时文件路径 + let temp_dir = std::env::temp_dir(); + let timestamp = chrono::Utc::now().format("%Y%m%d_%H%M%S_%3f"); + let filename = format!("proxycast_screenshot_{timestamp}.png"); + let temp_path = temp_dir.join(&filename); + + debug!("截图临时文件路径: {:?}", temp_path); + + // 根据平台调用不同的截图命令 + #[cfg(target_os = "macos")] + { + capture_macos(&temp_path).await?; + } + + #[cfg(target_os = "windows")] + { + capture_windows(&temp_path).await?; + } + + #[cfg(target_os = "linux")] + { + capture_linux(&temp_path).await?; + } + + // 检查文件是否存在(用户可能取消了截图) + if !temp_path.exists() { + info!("截图文件不存在,用户可能取消了截图"); + return Err(CaptureError::Cancelled); + } + + // 检查文件大小(空文件表示取消) + let metadata = std::fs::metadata(&temp_path)?; + if metadata.len() == 0 { + info!("截图文件为空,用户取消了截图"); + std::fs::remove_file(&temp_path)?; + return Err(CaptureError::Cancelled); + } + + info!("截图成功: {:?}", temp_path); + Ok(temp_path) +} + +/// macOS 截图实现 +#[cfg(target_os = "macos")] +async fn capture_macos(output_path: &PathBuf) -> Result<(), CaptureError> { + use std::process::Command; + + debug!("使用 macOS screencapture 命令"); + + // 先检查屏幕录制权限 + // 通过尝试执行一个快速的全屏截图到 /dev/null 来检测权限 + let permission_check = Command::new("screencapture") + .args(["-x", "-c"]) // -c 截图到剪贴板,快速检测权限 + .output(); + + if let Ok(output) = permission_check { + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr); + if stderr.contains("cannot") || stderr.contains("permission") { + warn!("屏幕录制权限不足: {}", stderr); + return Err(CaptureError::PermissionDenied( + "请在系统设置 → 隐私与安全性 → 录屏与系统录音 中授权 ProxyCast".to_string(), + )); + } + } + } + + // 使用 screencapture 命令 + // -i: 交互式选择区域 + // -x: 不播放截图声音 + let output = Command::new("screencapture") + .args(["-i", "-x", output_path.to_str().unwrap()]) + .output() + .map_err(|e| CaptureError::SystemError(format!("执行 screencapture 失败: {e}")))?; + + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr); + // 用户按 ESC 取消时,screencapture 返回非零状态但不创建文件 + if !output_path.exists() { + return Err(CaptureError::Cancelled); + } + error!("screencapture 命令失败: {}", stderr); + return Err(CaptureError::SystemError(format!( + "screencapture 失败: {stderr}" + ))); + } + + Ok(()) +} + +/// Windows 截图实现 +#[cfg(target_os = "windows")] +async fn capture_windows(output_path: &PathBuf) -> Result<(), CaptureError> { + use std::process::Command; + + debug!("使用 Windows 截图工具"); + + // 使用 PowerShell 调用截图功能 + // 注意:这是一个简化实现,实际可能需要使用 Windows API + let script = format!( + r#" + Add-Type -AssemblyName System.Windows.Forms + $screen = [System.Windows.Forms.Screen]::PrimaryScreen + $bitmap = New-Object System.Drawing.Bitmap($screen.Bounds.Width, $screen.Bounds.Height) + $graphics = [System.Drawing.Graphics]::FromImage($bitmap) + $graphics.CopyFromScreen($screen.Bounds.Location, [System.Drawing.Point]::Empty, $screen.Bounds.Size) + $bitmap.Save('{}') + "#, + output_path.to_str().unwrap().replace("\\", "\\\\") + ); + + let output = Command::new("powershell") + .args(["-Command", &script]) + .output() + .map_err(|e| CaptureError::SystemError(format!("执行 PowerShell 失败: {}", e)))?; + + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr); + error!("PowerShell 截图失败: {}", stderr); + return Err(CaptureError::SystemError(format!("截图失败: {}", stderr))); + } + + Ok(()) +} + +/// Linux 截图实现 +#[cfg(target_os = "linux")] +async fn capture_linux(output_path: &PathBuf) -> Result<(), CaptureError> { + use std::process::Command; + + debug!("使用 Linux 截图工具"); + + // 尝试使用 gnome-screenshot + let gnome_result = Command::new("gnome-screenshot") + .args(["-a", "-f", output_path.to_str().unwrap()]) + .output(); + + if let Ok(output) = gnome_result { + if output.status.success() { + return Ok(()); + } + } + + // 回退到 scrot + let scrot_result = Command::new("scrot") + .args(["-s", output_path.to_str().unwrap()]) + .output() + .map_err(|e| { + CaptureError::SystemError(format!( + "无法找到截图工具 (gnome-screenshot 或 scrot): {}", + e + )) + })?; + + if !scrot_result.status.success() { + let stderr = String::from_utf8_lossy(&scrot_result.stderr); + if !output_path.exists() { + return Err(CaptureError::Cancelled); + } + return Err(CaptureError::SystemError(format!("scrot 失败: {}", stderr))); + } + + Ok(()) +} + +/// 清理临时截图文件 +/// +/// # 参数 +/// - `path`: 要删除的文件路径 +pub fn cleanup_temp_file(path: &PathBuf) { + if path.exists() { + if let Err(e) = std::fs::remove_file(path) { + error!("删除临时截图文件失败: {}", e); + } else { + debug!("已删除临时截图文件: {:?}", path); + } + } +} + +#[cfg(test)] +mod tests { + + #[test] + fn test_temp_path_generation() { + let temp_dir = std::env::temp_dir(); + let timestamp = chrono::Utc::now().format("%Y%m%d_%H%M%S_%3f"); + let filename = format!("proxycast_screenshot_{timestamp}.png"); + let temp_path = temp_dir.join(&filename); + + assert!(temp_path + .to_str() + .unwrap() + .contains("proxycast_screenshot_")); + assert!(temp_path.to_str().unwrap().ends_with(".png")); + } +} diff --git a/src-tauri/crates/services/src/screenshot_image_service.rs b/src-tauri/crates/services/src/screenshot_image_service.rs new file mode 100644 index 000000000..59c8c1cb9 --- /dev/null +++ b/src-tauri/crates/services/src/screenshot_image_service.rs @@ -0,0 +1,31 @@ +//! 截图图片处理服务 +//! +//! 提供截图文件读取与 Base64 编码能力。 + +use base64::{engine::general_purpose::STANDARD, Engine}; +use std::path::Path; +use tokio::fs; + +/// 读取图片文件并转换为 Base64 +pub async fn read_image_as_base64(path: &str) -> Result { + tracing::debug!("读取图片为 Base64: {}", path); + + let file_path = Path::new(path); + + if !file_path.exists() { + return Err(format!("文件不存在: {}", file_path.display())); + } + + let bytes = fs::read(file_path) + .await + .map_err(|e| format!("读取文件失败: {e}"))?; + + if bytes.is_empty() { + return Err("文件为空".to_string()); + } + + let base64 = STANDARD.encode(&bytes); + + tracing::debug!("图片读取成功,大小: {} 字节", bytes.len()); + Ok(base64) +} diff --git a/src-tauri/crates/services/src/sysinfo_service.rs b/src-tauri/crates/services/src/sysinfo_service.rs new file mode 100644 index 000000000..459f22ed9 --- /dev/null +++ b/src-tauri/crates/services/src/sysinfo_service.rs @@ -0,0 +1,180 @@ +//! 系统信息服务 +//! +//! 提供 CPU、内存等系统信息的获取和订阅功能。 +//! 不依赖 Tauri,事件发送通过回调注入。 + +use serde::{Deserialize, Serialize}; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::Arc; +use sysinfo::{CpuRefreshKind, MemoryRefreshKind, RefreshKind, System}; +use tokio::sync::Mutex; +use tokio::time::{interval, Duration}; +use tracing::{debug, error, info}; + +/// 系统信息事件回调 +pub type SysinfoEmitter = Arc Result<(), String> + Send + Sync>; + +/// 系统信息数据点 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SysinfoData { + /// 时间戳(毫秒) + pub ts: u64, + /// 总 CPU 使用率(0-100) + pub cpu: f32, + /// 已用内存(GB) + #[serde(rename = "mem:used")] + pub mem_used: f64, + /// 总内存(GB) + #[serde(rename = "mem:total")] + pub mem_total: f64, + /// 各核心 CPU 使用率 + #[serde(flatten)] + pub cpu_cores: std::collections::HashMap, +} + +/// 系统信息服务状态 +pub struct SysinfoService { + /// 是否正在运行订阅 + running: AtomicBool, + /// 系统信息实例 + system: Mutex, +} + +impl SysinfoService { + /// 创建新的系统信息服务 + pub fn new() -> Self { + let system = System::new_with_specifics( + RefreshKind::new() + .with_cpu(CpuRefreshKind::everything()) + .with_memory(MemoryRefreshKind::everything()), + ); + Self { + running: AtomicBool::new(false), + system: Mutex::new(system), + } + } + + /// 获取当前系统信息快照 + pub async fn get_sysinfo(&self) -> SysinfoData { + let mut system = self.system.lock().await; + + system.refresh_cpu_usage(); + system.refresh_memory(); + + let ts = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_millis() as u64; + + let cpu = system.global_cpu_usage(); + let mem_total = system.total_memory() as f64 / 1024.0 / 1024.0 / 1024.0; + let mem_used = system.used_memory() as f64 / 1024.0 / 1024.0 / 1024.0; + + let mut cpu_cores = std::collections::HashMap::new(); + for (index, core) in system.cpus().iter().enumerate() { + cpu_cores.insert(format!("cpu:{index}"), core.cpu_usage()); + } + + SysinfoData { + ts, + cpu, + mem_used, + mem_total, + cpu_cores, + } + } + + /// 检查是否正在运行 + pub fn is_running(&self) -> bool { + self.running.load(Ordering::SeqCst) + } + + /// 停止订阅 + pub fn stop(&self) { + self.running.store(false, Ordering::SeqCst); + info!("系统信息订阅已停止"); + } +} + +impl Default for SysinfoService { + fn default() -> Self { + Self::new() + } +} + +/// 全局系统信息服务实例 +static SYSINFO_SERVICE: once_cell::sync::Lazy> = + once_cell::sync::Lazy::new(|| Arc::new(SysinfoService::new())); + +/// 获取全局系统信息服务 +pub fn get_sysinfo_service() -> Arc { + SYSINFO_SERVICE.clone() +} + +/// 获取当前系统信息 +pub async fn get_sysinfo() -> Result { + let service = get_sysinfo_service(); + Ok(service.get_sysinfo().await) +} + +/// 开始订阅系统信息(每秒回调一次) +pub async fn subscribe_sysinfo(emitter: SysinfoEmitter) -> Result<(), String> { + let service = get_sysinfo_service(); + + if service.is_running() { + debug!("系统信息订阅已在运行"); + return Ok(()); + } + + service.running.store(true, Ordering::SeqCst); + info!("开始系统信息订阅"); + + let service_clone = service.clone(); + tokio::spawn(async move { + let mut ticker = interval(Duration::from_secs(1)); + + ticker.tick().await; + + while service_clone.is_running() { + ticker.tick().await; + + if !service_clone.is_running() { + break; + } + + let data = service_clone.get_sysinfo().await; + if let Err(err) = emitter(&data) { + error!("发送系统信息事件失败: {}", err); + } + } + + info!("系统信息订阅任务结束"); + }); + + Ok(()) +} + +/// 停止订阅系统信息 +pub async fn unsubscribe_sysinfo() -> Result<(), String> { + let service = get_sysinfo_service(); + service.stop(); + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_get_sysinfo() { + let service = SysinfoService::new(); + tokio::time::sleep(Duration::from_millis(100)).await; + + let data = service.get_sysinfo().await; + assert!(data.ts > 0); + assert!(data.cpu >= 0.0 && data.cpu <= 100.0); + assert!(data.mem_total > 0.0); + assert!(data.mem_used >= 0.0); + assert!(data.mem_used <= data.mem_total); + } +} diff --git a/src-tauri/crates/services/src/voice_asr_service.rs b/src-tauri/crates/services/src/voice_asr_service.rs new file mode 100644 index 000000000..0097069cc --- /dev/null +++ b/src-tauri/crates/services/src/voice_asr_service.rs @@ -0,0 +1,310 @@ +//! ASR 服务 +//! +//! 统一管理语音识别服务,支持本地 Whisper 和云端 ASR。 +//! +//! ## 功能 +//! - 本地 Whisper 识别(离线、隐私) +//! - OpenAI Whisper API +//! - 百度语音识别 +//! - 讯飞语音识别(WebSocket 流式) +//! +//! ## 模型文件路径 +//! Whisper 模型文件存储在:`~/Library/Application Support/proxycast/models/whisper/` +//! +//! 支持的模型: +//! - `ggml-tiny.bin` (~75MB) +//! - `ggml-base.bin` (~142MB) +//! - `ggml-small.bin` (~466MB) +//! - `ggml-medium.bin` (~1.5GB) +//! +//! ## 使用示例 +//! ```rust,ignore +//! let credential = AsrService::get_default_credential()?.unwrap(); +//! let text = AsrService::transcribe(&credential, &audio_data, 16000).await?; +//! ``` + +#[cfg(feature = "local-whisper")] +use std::path::PathBuf; + +#[cfg(feature = "local-whisper")] +use proxycast_core::config::WhisperModelSize; +use proxycast_core::config::{AsrCredentialEntry, AsrProviderType}; + +use super::voice_config_service; +use voice_core::asr_client::{AsrClient, BaiduClient, OpenAIWhisperClient, XunfeiClient}; +use voice_core::types::AudioData; + +/// ASR 服务 +pub struct AsrService; + +impl AsrService { + /// 获取默认 ASR 凭证 + pub fn get_default_credential() -> Result, String> { + voice_config_service::get_default_asr_credential() + } + + /// 获取指定 ID 的 ASR 凭证 + pub fn get_credential(id: &str) -> Result, String> { + voice_config_service::get_asr_credential(id) + } + + /// 使用指定凭证进行语音识别 + /// + /// 当云端服务失败时,自动回退到本地 Whisper(需求 3.4) + pub async fn transcribe( + credential: &AsrCredentialEntry, + audio_data: &[u8], + sample_rate: u32, + ) -> Result { + // 如果是本地 Whisper,直接调用 + if matches!(credential.provider, AsrProviderType::WhisperLocal) { + return Self::transcribe_whisper_local(credential, audio_data, sample_rate).await; + } + + // 云端服务:先尝试云端,失败则回退到本地 Whisper + let cloud_result = match credential.provider { + AsrProviderType::OpenAI => { + Self::transcribe_openai(credential, audio_data, sample_rate).await + } + AsrProviderType::Baidu => { + Self::transcribe_baidu(credential, audio_data, sample_rate).await + } + AsrProviderType::Xunfei => { + Self::transcribe_xunfei(credential, audio_data, sample_rate).await + } + AsrProviderType::WhisperLocal => unreachable!(), // 已在上面处理 + }; + + // 云端成功,直接返回 + if cloud_result.is_ok() { + return cloud_result; + } + + // 云端失败,尝试回退到本地 Whisper + let cloud_error = cloud_result.unwrap_err(); + tracing::warn!( + "云端 ASR 服务 ({:?}) 失败: {},尝试回退到本地 Whisper", + credential.provider, + cloud_error + ); + + // 尝试获取本地 Whisper 凭证 + match Self::get_whisper_local_credential() { + Ok(Some(whisper_credential)) => { + tracing::info!("正在使用本地 Whisper 进行回退识别..."); + match Self::transcribe_whisper_local(&whisper_credential, audio_data, sample_rate) + .await + { + Ok(text) => { + tracing::info!("本地 Whisper 回退识别成功"); + Ok(text) + } + Err(whisper_error) => { + tracing::error!("本地 Whisper 回退也失败: {}", whisper_error); + // 返回原始云端错误,因为那是用户选择的服务 + Err(format!( + "云端服务失败: {cloud_error};本地 Whisper 回退也失败: {whisper_error}" + )) + } + } + } + Ok(None) => { + tracing::warn!("未找到本地 Whisper 凭证,无法回退"); + Err(format!( + "云端服务失败: {cloud_error};未配置本地 Whisper,无法回退" + )) + } + Err(e) => { + tracing::error!("获取本地 Whisper 凭证失败: {}", e); + Err(format!( + "云端服务失败: {cloud_error};获取本地 Whisper 凭证失败: {e}" + )) + } + } + } + + /// 获取本地 Whisper 凭证(用于回退) + fn get_whisper_local_credential() -> Result, String> { + voice_config_service::get_enabled_asr_credential_by_provider(AsrProviderType::WhisperLocal) + } + + /// 本地 Whisper 识别 + #[cfg(feature = "local-whisper")] + async fn transcribe_whisper_local( + credential: &AsrCredentialEntry, + audio_data: &[u8], + sample_rate: u32, + ) -> Result { + // 获取 Whisper 配置 + let whisper_config = credential + .whisper_config + .as_ref() + .ok_or("Whisper 本地配置缺失")?; + + // 获取模型文件路径 + let model_path = Self::get_whisper_model_path(&whisper_config.model)?; + + // 将 PCM 字节转换为 i16 采样 + let audio = Self::build_audio_data(audio_data, sample_rate)?; + + // 检查录音时长 + if !audio.is_valid() { + return Err("录音时间过短(需要至少 0.5 秒)".to_string()); + } + + // 转换模型大小枚举 + let model = Self::convert_model_size(&whisper_config.model); + + // 创建 Whisper 识别器 + let transcriber = + voice_core::WhisperTranscriber::new(model_path, model, &credential.language) + .map_err(|e| format!("Whisper 模型加载失败: {e}"))?; + + // 执行识别 + let result = transcriber + .transcribe(&audio) + .map_err(|e| format!("Whisper 识别失败: {e}"))?; + + Ok(result.text) + } + + /// 本地 Whisper 识别(未启用 local-whisper feature 时的 stub) + #[cfg(not(feature = "local-whisper"))] + async fn transcribe_whisper_local( + _credential: &AsrCredentialEntry, + _audio_data: &[u8], + _sample_rate: u32, + ) -> Result { + Err("本地 Whisper 功能未启用。请使用云端 ASR 服务(OpenAI、百度、讯飞)".to_string()) + } + + /// 获取 Whisper 模型文件路径 + #[cfg(feature = "local-whisper")] + fn get_whisper_model_path(model_size: &WhisperModelSize) -> Result { + // 模型文件名 + let filename = match model_size { + WhisperModelSize::Tiny => "ggml-tiny.bin", + WhisperModelSize::Base => "ggml-base.bin", + WhisperModelSize::Small => "ggml-small.bin", + WhisperModelSize::Medium => "ggml-medium.bin", + }; + + // 模型存储目录:~/Library/Application Support/proxycast/models/whisper/ + let models_dir = dirs::data_dir() + .ok_or("无法获取数据目录")? + .join("proxycast") + .join("models") + .join("whisper"); + + let model_path = models_dir.join(filename); + + // 检查模型文件是否存在 + if !model_path.exists() { + return Err(format!( + "Whisper 模型文件不存在: {}\n请下载模型文件到: {}", + filename, + models_dir.display() + )); + } + + Ok(model_path) + } + + /// 转换模型大小枚举 + #[cfg(feature = "local-whisper")] + fn convert_model_size(size: &WhisperModelSize) -> voice_core::types::WhisperModel { + match size { + WhisperModelSize::Tiny => voice_core::types::WhisperModel::Tiny, + WhisperModelSize::Base => voice_core::types::WhisperModel::Base, + WhisperModelSize::Small => voice_core::types::WhisperModel::Small, + WhisperModelSize::Medium => voice_core::types::WhisperModel::Medium, + } + } + + /// OpenAI Whisper API 识别 + 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)?; + + 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()); + } + + let result = client + .transcribe(&audio) + .await + .map_err(|e| format!("OpenAI Whisper 识别失败: {e}"))?; + + Ok(result.text) + } + + /// 百度语音识别 + async fn transcribe_baidu( + credential: &AsrCredentialEntry, + audio_data: &[u8], + sample_rate: u32, + ) -> Result { + let config = credential.baidu_config.as_ref().ok_or("百度配置缺失")?; + let audio = Self::build_audio_data(audio_data, sample_rate)?; + + let client = BaiduClient::new(config.api_key.clone(), config.secret_key.clone()); + let result = client + .transcribe(&audio) + .await + .map_err(|e| format!("百度识别失败: {e}"))?; + + Ok(result.text) + } + + /// 讯飞语音识别 + /// + /// 使用 WebSocket 流式识别,支持实时语音转文字 + async fn transcribe_xunfei( + credential: &AsrCredentialEntry, + audio_data: &[u8], + sample_rate: u32, + ) -> Result { + let config = credential.xunfei_config.as_ref().ok_or("讯飞配置缺失")?; + let audio = Self::build_audio_data(audio_data, sample_rate)?; + + // 创建讯飞客户端 + // 讯飞语言代码转换:zh -> zh_cn, en -> en_us + let xunfei_language = match credential.language.as_str() { + "zh" => "zh_cn".to_string(), + "en" => "en_us".to_string(), + other => other.to_string(), + }; + + let client = XunfeiClient::new( + config.app_id.clone(), + config.api_key.clone(), + config.api_secret.clone(), + ) + .with_language(xunfei_language); + + let result = client + .transcribe(&audio) + .await + .map_err(|e| format!("讯飞识别失败: {e}"))?; + + Ok(result.text) + } + + /// 将 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()); + } + + Ok(audio) + } +} diff --git a/src-tauri/crates/services/src/voice_command_service.rs b/src-tauri/crates/services/src/voice_command_service.rs new file mode 100644 index 000000000..58c953057 --- /dev/null +++ b/src-tauri/crates/services/src/voice_command_service.rs @@ -0,0 +1,155 @@ +//! 语音命令业务服务 +//! +//! 封装语音转写、润色、输出等可复用业务流程。 + +use serde::{Deserialize, Serialize}; + +use super::voice_asr_service::AsrService; +use super::voice_config_service; +use super::voice_output_service; +use super::voice_processor_service; + +/// 语音识别结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TranscribeResult { + /// 识别文本 + pub text: String, + /// 使用的 ASR 服务 + pub provider: String, +} + +/// 润色文本结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PolishResult { + /// 润色后的文本 + pub text: String, + /// 使用的指令 + pub instruction_name: String, +} + +/// 执行语音识别 +pub async fn transcribe_audio( + audio_data: &[u8], + sample_rate: u32, + credential_id: Option<&str>, +) -> Result { + tracing::info!( + "[语音识别] 开始识别,音频大小: {} 字节,采样率: {}", + audio_data.len(), + sample_rate + ); + + if audio_data.is_empty() { + tracing::error!("[语音识别] 音频数据为空!"); + return Err("音频数据为空,请检查麦克风权限".to_string()); + } + + let non_zero_count = audio_data.iter().filter(|&&byte| byte != 0).count(); + let non_zero_ratio = non_zero_count as f32 / audio_data.len() as f32; + tracing::info!( + "[语音识别] 非零字节比例: {:.2}% ({}/{})", + non_zero_ratio * 100.0, + non_zero_count, + audio_data.len() + ); + + if non_zero_ratio < 0.01 { + tracing::warn!("[语音识别] 音频数据几乎全为静音,可能是麦克风权限问题或未正确录音"); + } + + let credential = if let Some(id) = credential_id { + tracing::info!("[语音识别] 使用指定凭证: {}", id); + AsrService::get_credential(id)?.ok_or_else(|| format!("凭证不存在: {id}"))? + } else { + tracing::info!("[语音识别] 获取默认凭证..."); + match AsrService::get_default_credential() { + Ok(Some(credential)) => { + tracing::info!( + "[语音识别] 找到默认凭证: id={}, provider={:?}", + credential.id, + credential.provider + ); + credential + } + Ok(None) => { + if let Ok(credentials) = voice_config_service::list_asr_credentials() { + tracing::error!( + "[语音识别] 未找到默认凭证,当前 ASR 凭证数量: {}", + credentials.len() + ); + for (index, credential) in credentials.iter().enumerate() { + tracing::error!( + "[语音识别] 凭证 {}: id={}, is_default={}, disabled={}", + index, + credential.id, + credential.is_default, + credential.disabled + ); + } + } + return Err("未配置语音识别服务。请在设置 → 凭证池 → ASR 中添加讯飞、百度或 OpenAI Whisper 凭证。".to_string()); + } + Err(error) => { + tracing::error!("[语音识别] 获取默认凭证失败: {}", error); + return Err(format!("获取凭证失败: {error}")); + } + } + }; + + let provider_name = voice_config_service::asr_provider_name(credential.provider); + tracing::info!("[语音识别] 使用服务: {}", provider_name); + + let text = AsrService::transcribe(&credential, audio_data, sample_rate).await?; + tracing::info!("[语音识别] 识别完成,文本长度: {} 字符", text.len()); + + Ok(TranscribeResult { + text, + provider: provider_name.to_string(), + }) +} + +/// 润色文本 +pub async fn polish_voice_text( + text: &str, + instruction_id: Option<&str>, +) -> Result { + let voice_config = voice_config_service::load_voice_config()?; + let instruction_id = instruction_id + .map(|value| value.to_string()) + .unwrap_or_else(|| voice_config.processor.default_instruction_id.clone()); + + let instruction = voice_config + .instructions + .iter() + .find(|item| item.id == instruction_id) + .ok_or_else(|| format!("指令不存在: {instruction_id}"))?; + + if instruction_id == "raw" { + return Ok(PolishResult { + text: text.to_string(), + instruction_name: instruction.name.clone(), + }); + } + + let polished = voice_processor_service::polish_text( + text, + instruction, + voice_config.processor.polish_provider.as_deref(), + voice_config.processor.polish_model.as_deref(), + ) + .await?; + + Ok(PolishResult { + text: polished, + instruction_name: instruction.name.clone(), + }) +} + +/// 输出文本到系统 +pub fn output_voice_text(text: &str, mode: Option<&str>) -> Result<(), String> { + let output_mode = voice_config_service::resolve_output_mode(mode)?; + voice_output_service::output_text(text, output_mode)?; + + tracing::info!("[语音输出] 文本已输出: {} 字符", text.chars().count()); + Ok(()) +} diff --git a/src-tauri/crates/services/src/voice_config_service.rs b/src-tauri/crates/services/src/voice_config_service.rs new file mode 100644 index 000000000..4161f03ed --- /dev/null +++ b/src-tauri/crates/services/src/voice_config_service.rs @@ -0,0 +1,132 @@ +//! 语音输入配置服务 +//! +//! 管理语音输入配置、ASR 凭证与润色指令。 +//! 不依赖 Tauri,可被主 crate 以桥接方式复用。 + +use proxycast_core::config::{ + load_config, save_config, AsrCredentialEntry, AsrProviderType, VoiceInputConfig, + VoiceInstruction, VoiceOutputMode, +}; + +/// 加载语音输入配置 +pub fn load_voice_config() -> Result { + let config = load_config().map_err(|e| e.to_string())?; + Ok(config.experimental.voice_input) +} + +/// 保存语音输入配置 +pub fn save_voice_config(voice_config: VoiceInputConfig) -> Result<(), String> { + let mut config = load_config().map_err(|e| e.to_string())?; + config.experimental.voice_input = voice_config; + save_config(&config).map_err(|e| e.to_string())?; + Ok(()) +} + +/// 获取默认 ASR 凭证 +pub fn get_default_asr_credential() -> Result, String> { + let config = load_config().map_err(|e| e.to_string())?; + Ok(config + .credential_pool + .asr + .into_iter() + .find(|credential| credential.is_default && !credential.disabled)) +} + +/// 获取指定 ID 的 ASR 凭证 +pub fn get_asr_credential(id: &str) -> Result, String> { + let config = load_config().map_err(|e| e.to_string())?; + Ok(config + .credential_pool + .asr + .into_iter() + .find(|credential| credential.id == id)) +} + +/// 列出所有 ASR 凭证 +pub fn list_asr_credentials() -> Result, String> { + let config = load_config().map_err(|e| e.to_string())?; + Ok(config.credential_pool.asr) +} + +/// 获取首个启用的指定 Provider 凭证 +pub fn get_enabled_asr_credential_by_provider( + provider: AsrProviderType, +) -> Result, String> { + let config = load_config().map_err(|e| e.to_string())?; + Ok(config + .credential_pool + .asr + .into_iter() + .find(|credential| credential.provider == provider && !credential.disabled)) +} + +/// 获取指令列表 +pub fn get_instructions() -> Result, String> { + let config = load_config().map_err(|e| e.to_string())?; + Ok(config.experimental.voice_input.instructions) +} + +/// 获取指定 ID 的指令 +pub fn get_instruction(id: &str) -> Result, String> { + let instructions = get_instructions()?; + Ok(instructions + .into_iter() + .find(|instruction| instruction.id == id)) +} + +/// 保存或更新语音指令 +pub fn save_voice_instruction(instruction: VoiceInstruction) -> Result<(), String> { + let mut voice_config = load_voice_config()?; + + if let Some(index) = voice_config + .instructions + .iter() + .position(|item| item.id == instruction.id) + { + voice_config.instructions[index] = instruction; + } else { + voice_config.instructions.push(instruction); + } + + save_voice_config(voice_config) +} + +/// 删除语音指令(预设指令不可删除) +pub fn delete_voice_instruction(id: &str) -> Result<(), String> { + let mut voice_config = load_voice_config()?; + + if let Some(instruction) = voice_config.instructions.iter().find(|item| item.id == id) { + if instruction.is_preset { + return Err("无法删除预设指令".to_string()); + } + } + + voice_config.instructions.retain(|item| item.id != id); + save_voice_config(voice_config) +} + +/// 解析输出模式 +/// +/// 当 `mode` 为 `None` 时,返回配置中的默认输出模式。 +pub fn resolve_output_mode(mode: Option<&str>) -> Result { + match mode { + Some("type") => Ok(VoiceOutputMode::Type), + Some("clipboard") => Ok(VoiceOutputMode::Clipboard), + Some("both") => Ok(VoiceOutputMode::Both), + None => { + let voice_config = load_voice_config()?; + Ok(voice_config.output.mode) + } + Some(other) => Err(format!("未知的输出模式: {other}")), + } +} + +/// 获取 ASR Provider 展示名 +pub fn asr_provider_name(provider: AsrProviderType) -> &'static str { + match provider { + AsrProviderType::WhisperLocal => "本地 Whisper", + AsrProviderType::OpenAI => "OpenAI Whisper", + AsrProviderType::Baidu => "百度语音", + AsrProviderType::Xunfei => "讯飞语音", + } +} diff --git a/src-tauri/crates/services/src/voice_output_service.rs b/src-tauri/crates/services/src/voice_output_service.rs new file mode 100644 index 000000000..b9fa86573 --- /dev/null +++ b/src-tauri/crates/services/src/voice_output_service.rs @@ -0,0 +1,22 @@ +//! 语音文本输出服务 +//! +//! 提供模拟键盘输入和剪贴板输出能力。 + +use proxycast_core::config::VoiceOutputMode; +use voice_core::{OutputHandler, OutputMode}; + +/// 输出文字到系统 +/// +/// 根据配置的输出模式,将文字输出到当前焦点应用。 +pub fn output_text(text: &str, mode: VoiceOutputMode) -> Result<(), String> { + 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/crates/services/src/voice_processor_service.rs b/src-tauri/crates/services/src/voice_processor_service.rs new file mode 100644 index 000000000..15df26b47 --- /dev/null +++ b/src-tauri/crates/services/src/voice_processor_service.rs @@ -0,0 +1,49 @@ +//! 语音文本处理服务 +//! +//! 提供语音识别文本的 Prompt 套用与 LLM 润色能力。 + +use proxycast_core::config::VoiceInstruction; + +/// 处理文本(应用指令模板) +pub fn process_text(text: &str, instruction: &VoiceInstruction) -> String { + voice_core::text_polish::apply_prompt_template(text, &instruction.prompt) +} + +/// 使用 LLM 润色文本 +/// +/// 通过本地 API 服务器调用 LLM 进行文本润色。 +pub async fn polish_text( + text: &str, + instruction: &VoiceInstruction, + _provider: Option<&str>, + model: Option<&str>, +) -> Result { + if instruction.id == "raw" { + return Ok(text.to_string()); + } + + let prompt = process_text(text, instruction); + call_local_llm(&prompt, model, &instruction.id).await +} + +/// 调用本地 API 服务器进行 LLM 推理 +async fn call_local_llm( + prompt: &str, + model: Option<&str>, + instruction_id: &str, +) -> Result { + use proxycast_core::config::load_config; + + let config = load_config().map_err(|e| e.to_string())?; + let base_url = format!("http://{}:{}", config.server.host, config.server.port); + let api_key = &config.server.api_key; + + voice_core::text_polish::polish_with_local_api( + &base_url, + api_key, + prompt, + model, + instruction_id, + ) + .await +} diff --git a/src-tauri/crates/services/src/voice_recording_service.rs b/src-tauri/crates/services/src/voice_recording_service.rs new file mode 100644 index 000000000..a44175a2f --- /dev/null +++ b/src-tauri/crates/services/src/voice_recording_service.rs @@ -0,0 +1,35 @@ +//! 录音服务桥接层 +//! +//! 录音核心逻辑已迁移到 `voice-core` 的 `threaded_recorder` 模块。 +//! 本模块保留 Tauri State 包装和向后兼容导出路径。 + +use parking_lot::Mutex; +use std::sync::Arc; + +pub use voice_core::{AudioDeviceInfo, RecordingCommand, RecordingResponse, RecordingService}; + +/// 获取所有可用的麦克风设备 +pub fn list_audio_devices() -> Result, String> { + voice_core::list_audio_devices().map_err(|e| e.to_string()) +} + +/// 全局录音服务状态(Tauri State 包装) +pub struct RecordingServiceState(pub Arc>); + +impl RecordingServiceState { + /// 创建新的录音服务状态 + pub fn new() -> Self { + Self(Arc::new(Mutex::new(RecordingService::new()))) + } +} + +impl Default for RecordingServiceState { + fn default() -> Self { + Self::new() + } +} + +/// 创建录音服务状态 +pub fn create_recording_service_state() -> RecordingServiceState { + RecordingServiceState::new() +} diff --git a/src-tauri/crates/skills/Cargo.toml b/src-tauri/crates/skills/Cargo.toml index fd0a1c7fe..878fc6da6 100644 --- a/src-tauri/crates/skills/Cargo.toml +++ b/src-tauri/crates/skills/Cargo.toml @@ -6,6 +6,11 @@ authors.workspace = true repository.workspace = true [dependencies] +proxycast-core.workspace = true +proxycast-providers.workspace = true +proxycast-services.workspace = true +proxycast-server-utils.workspace = true + serde.workspace = true serde_json.workspace = true async-trait.workspace = true diff --git a/src-tauri/crates/skills/src/lib.rs b/src-tauri/crates/skills/src/lib.rs index a0d3366b6..641b9dab9 100644 --- a/src-tauri/crates/skills/src/lib.rs +++ b/src-tauri/crates/skills/src/lib.rs @@ -1,10 +1,11 @@ //! ProxyCast Skills Crate //! //! 包含 Skills 系统的 trait 定义和纯逻辑部分。 -//! Tauri 相关的实现(TauriExecutionCallback、ProxyCastLlmProvider)留在主 crate。 +//! Tauri 相关实现(TauriExecutionCallback)保留在主 crate。 mod execution_callback; mod llm_provider; +mod proxycast_llm_provider; mod skill_loader; pub use execution_callback::{ @@ -12,6 +13,7 @@ pub use execution_callback::{ StepStartPayload, }; pub use llm_provider::{LlmProvider, SkillError}; +pub use proxycast_llm_provider::ProxyCastLlmProvider; 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, diff --git a/src-tauri/crates/skills/src/proxycast_llm_provider.rs b/src-tauri/crates/skills/src/proxycast_llm_provider.rs new file mode 100644 index 000000000..f4cfa15b8 --- /dev/null +++ b/src-tauri/crates/skills/src/proxycast_llm_provider.rs @@ -0,0 +1,543 @@ +//! ProxyCast LLM Provider 实现 +//! +//! 使用 ProviderPoolService 选择凭证并调用 LLM API。 +//! trait 定义(LlmProvider, SkillError)已迁移到 proxycast-skills crate。 + +use std::sync::Arc; + +use async_trait::async_trait; + +use proxycast_core::database::DbConnection; +use proxycast_core::models::anthropic::AnthropicMessagesRequest; +#[cfg(test)] +use proxycast_core::models::provider_pool_model::PoolProviderType; +use proxycast_core::models::provider_pool_model::{CredentialData, ProviderCredential}; +use proxycast_providers::providers::claude_custom::ClaudeCustomProvider; +use proxycast_providers::providers::kiro::KiroProvider; +use proxycast_providers::providers::openai_custom::OpenAICustomProvider; +use proxycast_services::api_key_provider_service::ApiKeyProviderService; +use proxycast_services::provider_pool_service::ProviderPoolService; + +use crate::{LlmProvider, SkillError}; + +/// ProxyCast LLM Provider +/// +/// 使用 ProviderPoolService 选择凭证并调用 LLM API。 +/// 实现 aster-rust 定义的 LlmProvider trait。 +pub struct ProxyCastLlmProvider { + /// 凭证池服务 + pool_service: Arc, + /// API Key Provider 服务(用于智能降级) + api_key_service: Arc, + /// 数据库连接 + db: DbConnection, + /// 偏好的 Provider 类型(可选) + preferred_provider: Option, +} + +impl ProxyCastLlmProvider { + /// 创建新的 ProxyCastLlmProvider 实例 + /// + /// # Arguments + /// * `pool_service` - 凭证池服务 + /// * `api_key_service` - API Key 服务 + /// * `db` - 数据库连接 + pub fn new( + pool_service: Arc, + api_key_service: Arc, + db: DbConnection, + ) -> Self { + Self { + pool_service, + api_key_service, + db, + preferred_provider: None, + } + } + + /// 创建带有偏好 Provider 的实例 + /// + /// # Arguments + /// * `pool_service` - 凭证池服务 + /// * `api_key_service` - API Key 服务 + /// * `db` - 数据库连接 + /// * `preferred_provider` - 偏好的 Provider 类型 + pub fn with_preferred_provider( + pool_service: Arc, + api_key_service: Arc, + db: DbConnection, + preferred_provider: String, + ) -> Self { + Self { + pool_service, + api_key_service, + db, + preferred_provider: Some(preferred_provider), + } + } + + /// 设置偏好的 Provider 类型 + pub fn set_preferred_provider(&mut self, provider: Option) { + self.preferred_provider = provider; + } + + /// 获取偏好的 Provider 类型 + pub fn preferred_provider(&self) -> Option<&str> { + self.preferred_provider.as_deref() + } + + /// 将 Skill 的 provider 字段映射到 PoolProviderType + /// + /// # Arguments + /// * `provider` - Provider 名称字符串 + /// + /// # Returns + /// 对应的 PoolProviderType,未知类型返回 None + #[cfg(test)] + fn map_skill_provider_to_pool_type(provider: &str) -> Option { + match provider.to_lowercase().as_str() { + "openai" | "gpt" => Some(PoolProviderType::OpenAI), + "anthropic" | "claude" => Some(PoolProviderType::Claude), + "gemini" | "google" => Some(PoolProviderType::Gemini), + "kiro" | "codewhisperer" => Some(PoolProviderType::Kiro), + "vertex" => Some(PoolProviderType::Vertex), + "codex" => Some(PoolProviderType::Codex), + _ => None, + } + } + + /// 根据凭证调用 LLM API + /// + /// # Arguments + /// * `credential` - 选中的凭证 + /// * `system_prompt` - 系统提示词 + /// * `user_message` - 用户消息 + /// * `model` - 模型名称 + /// + /// # Returns + /// LLM 响应文本或错误 + async fn call_llm_with_credential( + &self, + credential: &ProviderCredential, + system_prompt: &str, + user_message: &str, + model: &str, + ) -> Result { + match &credential.credential { + CredentialData::KiroOAuth { creds_file_path } => { + self.call_kiro_api(creds_file_path, system_prompt, user_message, model) + .await + } + CredentialData::ClaudeKey { api_key, base_url } => { + self.call_claude_api( + api_key, + base_url.as_deref(), + system_prompt, + user_message, + model, + ) + .await + } + CredentialData::OpenAIKey { api_key, base_url } => { + self.call_openai_api( + api_key, + base_url.as_deref(), + system_prompt, + user_message, + model, + ) + .await + } + CredentialData::AnthropicKey { api_key, base_url } => { + // Anthropic API Key 使用 Claude API + self.call_claude_api( + api_key, + base_url.as_deref(), + system_prompt, + user_message, + model, + ) + .await + } + _ => Err(SkillError::ProviderError(format!( + "不支持的凭证类型: {:?}", + credential.provider_type + ))), + } + } + + /// 调用 Kiro API + async fn call_kiro_api( + &self, + creds_file_path: &str, + system_prompt: &str, + user_message: &str, + model: &str, + ) -> Result { + use proxycast_core::models::anthropic::AnthropicMessage; + use proxycast_providers::converter::anthropic_to_openai::convert_anthropic_to_openai; + use proxycast_providers::providers::traits::CredentialProvider; + use proxycast_server_utils::parse_cw_response; + + let mut kiro = KiroProvider::new(); + kiro.load_credentials_from_path(creds_file_path) + .await + .map_err(|e| SkillError::ProviderError(format!("加载 Kiro 凭证失败: {}", e)))?; + + // 确保 Token 有效 + if !kiro.is_token_valid() || kiro.is_token_expiring_soon() { + kiro.refresh_token() + .await + .map_err(|e| SkillError::ProviderError(format!("刷新 Token 失败: {}", e)))?; + } + + // 构建 Anthropic 请求 + let request = AnthropicMessagesRequest { + model: model.to_string(), + max_tokens: Some(4096), + system: Some(serde_json::Value::String(system_prompt.to_string())), + messages: vec![AnthropicMessage { + role: "user".to_string(), + content: serde_json::Value::String(user_message.to_string()), + }], + stream: false, + temperature: None, + tools: None, + tool_choice: None, + }; + + // 转换为 OpenAI 格式并调用 + let openai_request = convert_anthropic_to_openai(&request); + let resp = kiro + .call_api(&openai_request) + .await + .map_err(|e| SkillError::ProviderError(format!("Kiro API 调用失败: {}", e)))?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(SkillError::ProviderError(format!( + "Kiro API 返回错误: status={}, body={}", + status, body + ))); + } + + let bytes = resp + .bytes() + .await + .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; + let body = String::from_utf8_lossy(&bytes).to_string(); + let parsed = parse_cw_response(&body); + + Ok(parsed.content) + } + + /// 调用 Claude API + async fn call_claude_api( + &self, + api_key: &str, + base_url: Option<&str>, + system_prompt: &str, + user_message: &str, + model: &str, + ) -> Result { + use proxycast_core::models::anthropic::AnthropicMessage; + + let claude = + ClaudeCustomProvider::with_config(api_key.to_string(), base_url.map(|s| s.to_string())); + + // 构建 Anthropic 请求 + let request = AnthropicMessagesRequest { + model: model.to_string(), + max_tokens: Some(4096), + system: Some(serde_json::Value::String(system_prompt.to_string())), + messages: vec![AnthropicMessage { + role: "user".to_string(), + content: serde_json::Value::String(user_message.to_string()), + }], + stream: false, + temperature: None, + tools: None, + tool_choice: None, + }; + + let resp = claude + .call_api(&request) + .await + .map_err(|e| SkillError::ProviderError(format!("Claude API 调用失败: {}", e)))?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(SkillError::ProviderError(format!( + "Claude API 返回错误: status={}, body={}", + status, body + ))); + } + + let body = resp + .text() + .await + .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; + + // 解析 Anthropic 响应 + let json: serde_json::Value = serde_json::from_str(&body) + .map_err(|e| SkillError::ProviderError(format!("解析响应失败: {}", e)))?; + + // 提取文本内容 + let content = json["content"] + .as_array() + .and_then(|arr| arr.first()) + .and_then(|block| block["text"].as_str()) + .unwrap_or(""); + + Ok(content.to_string()) + } + + /// 调用 OpenAI API + async fn call_openai_api( + &self, + api_key: &str, + base_url: Option<&str>, + system_prompt: &str, + user_message: &str, + model: &str, + ) -> Result { + use proxycast_core::models::openai::{ChatCompletionRequest, ChatMessage, MessageContent}; + + let openai = + OpenAICustomProvider::with_config(api_key.to_string(), base_url.map(|s| s.to_string())); + + // 构建 OpenAI 请求 + let request = ChatCompletionRequest { + model: model.to_string(), + messages: vec![ + ChatMessage { + role: "system".to_string(), + content: Some(MessageContent::Text(system_prompt.to_string())), + tool_calls: None, + tool_call_id: None, + reasoning_content: None, + }, + ChatMessage { + role: "user".to_string(), + content: Some(MessageContent::Text(user_message.to_string())), + tool_calls: None, + tool_call_id: None, + reasoning_content: None, + }, + ], + max_tokens: Some(4096), + stream: false, + temperature: None, + top_p: None, + tools: None, + tool_choice: None, + reasoning_effort: None, + }; + + let resp = openai + .call_api(&request) + .await + .map_err(|e| SkillError::ProviderError(format!("OpenAI API 调用失败: {}", e)))?; + + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(SkillError::ProviderError(format!( + "OpenAI API 返回错误: status={}, body={}", + status, body + ))); + } + + let body = resp + .text() + .await + .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; + + // 解析 OpenAI 响应 + let json: serde_json::Value = serde_json::from_str(&body) + .map_err(|e| SkillError::ProviderError(format!("解析响应失败: {}", e)))?; + + // 提取文本内容 + let content = json["choices"] + .as_array() + .and_then(|arr| arr.first()) + .and_then(|choice| choice["message"]["content"].as_str()) + .unwrap_or(""); + + Ok(content.to_string()) + } +} + +#[async_trait] +impl LlmProvider for ProxyCastLlmProvider { + /// 调用 LLM 进行对话 + /// + /// # 实现说明 + /// 1. 使用 ProviderPoolService.select_credential_with_fallback() 选择凭证 + /// 2. 如果指定了 preferred_provider,优先选择该类型的凭证 + /// 3. 如果指定了 model,传递给底层 provider + /// 4. 如果没有可用凭证,返回 ProviderError + /// + /// # Requirements + /// - 1.2: 使用 ProviderPoolService 选择可用凭证 + /// - 1.3: 优先选择指定 provider 类型的凭证 + /// - 1.4: 将 model 参数传递给底层 provider + /// - 1.5: 没有可用凭证时返回 ProviderError + async fn chat( + &self, + system_prompt: &str, + user_message: &str, + model: Option<&str>, + ) -> Result { + // 确定要使用的 provider 类型 + let provider_type = self.preferred_provider.as_deref().unwrap_or("claude"); // 默认使用 Claude + + // 确定要使用的模型 + let model_name = model.unwrap_or("claude-sonnet-4-5-20250514"); + + tracing::info!( + "[ProxyCastLlmProvider] chat 调用: provider_type={}, model={}", + provider_type, + model_name + ); + + // 使用 ProviderPoolService 选择凭证(Requirements 1.2, 1.3) + let credential = self + .pool_service + .select_credential_with_fallback( + &self.db, + &self.api_key_service, + provider_type, + Some(model_name), + None, // provider_id_hint + None, // client_type + ) + .await + .map_err(|e| SkillError::ProviderError(format!("选择凭证失败: {}", e)))? + .ok_or_else(|| { + // Requirements 1.5: 没有可用凭证时返回 ProviderError + SkillError::ProviderError(format!( + "没有可用的凭证: provider_type={}, model={}", + provider_type, model_name + )) + })?; + + tracing::info!( + "[ProxyCastLlmProvider] 选中凭证: uuid={}, type={:?}", + &credential.uuid[..8], + credential.provider_type + ); + + // 调用 LLM API(Requirements 1.4: 传递 model 参数) + let result = self + .call_llm_with_credential(&credential, system_prompt, user_message, model_name) + .await; + + // 记录使用情况 + match &result { + Ok(_) => { + let _ = self.pool_service.record_usage(&self.db, &credential.uuid); + let _ = + self.pool_service + .mark_healthy(&self.db, &credential.uuid, Some(model_name)); + } + Err(e) => { + let _ = self.pool_service.mark_unhealthy( + &self.db, + &credential.uuid, + Some(&e.to_string()), + ); + } + } + + result + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_map_skill_provider_openai() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("openai"), + Some(PoolProviderType::OpenAI) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("gpt"), + Some(PoolProviderType::OpenAI) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("OPENAI"), + Some(PoolProviderType::OpenAI) + ); + } + + #[test] + fn test_map_skill_provider_claude() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("claude"), + Some(PoolProviderType::Claude) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("anthropic"), + Some(PoolProviderType::Claude) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("CLAUDE"), + Some(PoolProviderType::Claude) + ); + } + + #[test] + fn test_map_skill_provider_gemini() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("gemini"), + Some(PoolProviderType::Gemini) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("google"), + Some(PoolProviderType::Gemini) + ); + } + + #[test] + fn test_map_skill_provider_kiro() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("kiro"), + Some(PoolProviderType::Kiro) + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("codewhisperer"), + Some(PoolProviderType::Kiro) + ); + } + + #[test] + fn test_map_skill_provider_unknown() { + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type("unknown_provider"), + None + ); + assert_eq!( + ProxyCastLlmProvider::map_skill_provider_to_pool_type(""), + None + ); + } + + #[test] + fn test_skill_error_display() { + let provider_err = SkillError::ProviderError("没有可用凭证".to_string()); + assert!(provider_err.to_string().contains("Provider error")); + assert!(provider_err.to_string().contains("没有可用凭证")); + + let exec_err = SkillError::ExecutionError("执行失败".to_string()); + assert!(exec_err.to_string().contains("Execution error")); + + let config_err = SkillError::ConfigError("配置错误".to_string()); + assert!(config_err.to_string().contains("Config error")); + } +} diff --git a/src-tauri/crates/voice-core/README.md b/src-tauri/crates/voice-core/README.md index 390250b4f..f98c1cd1b 100644 --- a/src-tauri/crates/voice-core/README.md +++ b/src-tauri/crates/voice-core/README.md @@ -19,6 +19,7 @@ src/ ├── device.rs # 音频设备枚举 ├── recorder.rs # 音频录制 ├── threaded_recorder.rs # 线程化录音服务(可跨线程控制) +├── text_polish.rs # 文本润色与本地 LLM 调用 ├── transcriber.rs # Whisper 本地识别 ├── output.rs # 文字输出 └── asr_client/ # 云端 ASR diff --git a/src-tauri/crates/voice-core/src/lib.rs b/src-tauri/crates/voice-core/src/lib.rs index 8d69b3ff8..15eb1f8dc 100644 --- a/src-tauri/crates/voice-core/src/lib.rs +++ b/src-tauri/crates/voice-core/src/lib.rs @@ -8,6 +8,7 @@ pub mod device; pub mod error; pub mod output; pub mod recorder; +pub mod text_polish; pub mod threaded_recorder; #[cfg(feature = "local-whisper")] pub mod transcriber; diff --git a/src-tauri/crates/voice-core/src/text_polish.rs b/src-tauri/crates/voice-core/src/text_polish.rs new file mode 100644 index 000000000..85dbae3f5 --- /dev/null +++ b/src-tauri/crates/voice-core/src/text_polish.rs @@ -0,0 +1,133 @@ +//! 语音文本润色 +//! +//! 提供与 Tauri 无关的文本模板处理和本地 LLM 调用能力。 + +/// 语音润色的 System Prompt +pub const VOICE_POLISH_SYSTEM_PROMPT: &str = r#"你是一个语音转文字的后处理助手。 + +## 背景说明 +用户通过麦克风录音,然后使用语音识别 API(如讯飞、百度等)将语音转换为文字。由于录音环境、口音、语速等因素,识别结果可能存在以下问题: + +1. **重复字词**:语音识别的流式返回机制可能导致字词重复,如"你你好好"实际是"你好" +2. **语气词和填充词**:如"嗯"、"啊"、"那个"、"就是"、"然后"等口语习惯 +3. **同音字错误**:识别引擎可能选错同音字,如"准备"识别成"准被" +4. **环境噪音干扰**:背景噪音可能被误识别为无意义的字词 +5. **断句错误**:缺少标点或标点位置不当 +6. **多余内容**:录音开始或结束时的杂音可能被识别为无关文字 + +## 你的任务 +根据上下文语义,智能还原用户真正想表达的内容: +- 去除明显的重复字词 +- 去除无意义的语气词和填充词 +- 根据语义修正可能的同音字错误 +- 过滤掉噪音产生的无意义内容 +- 添加合适的标点符号 +- 保持用户的原意,不要添加或臆测内容 + +## 输出要求 +只输出处理后的文本,不要添加任何解释、说明或前缀。"#; + +/// 应用 Prompt 模板(将 `{{text}}` 替换为实际文本) +pub fn apply_prompt_template(text: &str, template: &str) -> String { + template.replace("{{text}}", text) +} + +/// 调用本地 API 服务器进行 LLM 润色 +pub async fn polish_with_local_api( + base_url: &str, + api_key: &str, + prompt: &str, + model: Option<&str>, + instruction_id: &str, +) -> Result { + // 使用配置的模型,如果没有配置则使用 deepseek-chat + let model_name = model.filter(|m| !m.is_empty()).unwrap_or("deepseek-chat"); + + tracing::info!( + "[语音润色] 使用模型: {}, 指令: {}", + model_name, + instruction_id + ); + + #[derive(serde::Serialize)] + struct Message { + role: String, + content: String, + } + + #[derive(serde::Serialize)] + struct ChatRequest { + model: String, + messages: Vec, + max_tokens: u32, + temperature: f32, + } + + let messages = if instruction_id == "default" { + vec![ + Message { + role: "system".to_string(), + content: VOICE_POLISH_SYSTEM_PROMPT.to_string(), + }, + Message { + role: "user".to_string(), + content: prompt.to_string(), + }, + ] + } else { + vec![Message { + role: "user".to_string(), + content: prompt.to_string(), + }] + }; + + let request = ChatRequest { + model: model_name.to_string(), + messages, + max_tokens: 2048, + temperature: 0.3, + }; + + let endpoint = format!("{}/v1/chat/completions", base_url.trim_end_matches('/')); + + let response = reqwest::Client::new() + .post(endpoint) + .header("Authorization", format!("Bearer {api_key}")) + .header("Content-Type", "application/json") + .json(&request) + .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!("LLM API 错误: {status} - {body}")); + } + + #[derive(serde::Deserialize)] + struct Choice { + message: ResponseMessage, + } + + #[derive(serde::Deserialize)] + struct ResponseMessage { + content: Option, + } + + #[derive(serde::Deserialize)] + struct ChatResponse { + choices: Vec, + } + + let result: ChatResponse = response + .json() + .await + .map_err(|e| format!("解析响应失败: {e}"))?; + + result + .choices + .first() + .and_then(|c| c.message.content.clone()) + .ok_or_else(|| "LLM 返回空内容".to_string()) +} diff --git a/src-tauri/crates/voice-core/src/threaded_recorder.rs b/src-tauri/crates/voice-core/src/threaded_recorder.rs index 2ebc5389e..837879a5d 100644 --- a/src-tauri/crates/voice-core/src/threaded_recorder.rs +++ b/src-tauri/crates/voice-core/src/threaded_recorder.rs @@ -19,13 +19,13 @@ //! - Tauri 命令通过 channel 发送控制指令 //! - 录音线程通过 channel 返回结果 +use crate::types::AudioData; 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)] diff --git a/src-tauri/src/agent/README.md b/src-tauri/src/agent/README.md index 6eb4b0498..2249eff89 100644 --- a/src-tauri/src/agent/README.md +++ b/src-tauri/src/agent/README.md @@ -19,12 +19,12 @@ AI Agent 集成模块,基于 aster-rust 框架实现。 |------|------| | `mod.rs` | 模块入口,导出公共类型 | | `types.rs` | Agent 相关类型定义 | -| `aster_state.rs` | Aster Agent 状态管理(Provider 配置、取消令牌、Skills 加载) | -| `aster_agent.rs` | Aster Agent 包装器(会话管理) | +| `aster_state.rs` | Aster Agent 状态管理(主状态桥接;会话配置/项目上下文/身份配置/Skills 加载辅助逻辑委托 `crates/agent/src/aster_state_support.rs`) | +| `aster_agent.rs` | Aster Agent 包装器(流式桥接,会话存储逻辑委托 `crates/agent/src/session_store.rs`) | | `event_converter.rs` | Aster 事件到 Tauri 事件转换 | -| `credential_bridge.rs` | 凭证池桥接(连接 ProxyCast 凭证池与 Aster Provider,智能拆分 base_url) | +| `credential_bridge.rs` | 重导出层(纯逻辑已迁移到 `crates/agent/src/credential_bridge.rs`) | | `mcp_bridge.rs` | MCP 服务桥接 | -| `subagent_scheduler.rs` | 子 Agent 调度器 | +| `subagent_scheduler.rs` | Tauri 事件桥接层(纯逻辑已迁移到 `crates/agent/src/subagent_scheduler.rs`) | ## Skills 集成 @@ -34,7 +34,7 @@ Agent 初始化时自动加载 `~/.proxycast/skills/` 目录下的 Skills: ```rust // init_agent_with_db() 内部调用 -Self::load_proxycast_skills(); +proxycast_agent::reload_proxycast_skills(); ``` ### AI 自动调用 @@ -104,7 +104,7 @@ let stream = agent.reply(user_message, session_config, Some(cancel_token)).await ## 凭证池桥接 -`credential_bridge.rs` 模块将 ProxyCast 凭证池与 Aster Provider 系统连接: +`credential_bridge.rs` 在主 crate 中仅作为兼容导出,核心逻辑位于 `crates/agent/src/credential_bridge.rs`: - 自动从凭证池选择可用凭证 - 支持 OAuth 和 API Key 两种凭证类型 diff --git a/src-tauri/src/agent/aster_agent.rs b/src-tauri/src/agent/aster_agent.rs index a9ddc4774..1976a632d 100644 --- a/src-tauri/src/agent/aster_agent.rs +++ b/src-tauri/src/agent/aster_agent.rs @@ -1,16 +1,17 @@ //! Aster Agent 包装器 //! -//! 提供简化的接口来使用 Aster Agent -//! 处理消息发送、事件流转换和会话管理 +//! 提供简化的接口来使用 Aster Agent。 +//! 处理消息发送、事件流转换,并桥接会话存储服务。 use crate::agent::aster_state::{AsterAgentState, SessionConfigBuilder}; -use crate::database::dao::agent::AgentDao; use crate::database::DbConnection; use aster::conversation::message::Message; -use chrono::Utc; use futures::StreamExt; +use proxycast_agent::{convert_agent_event, TauriAgentEvent}; use tauri::{AppHandle, Emitter}; +pub use proxycast_agent::session_store::{SessionDetail, SessionInfo}; + /// Aster Agent 包装器 /// /// 提供与 Tauri 集成的简化接口 @@ -37,72 +38,56 @@ impl AsterAgentWrapper { session_id: String, event_name: String, ) -> Result<(), String> { - // 1. 初始化检查(使用带数据库的版本) if !state.is_initialized().await { state.init_agent_with_db(db).await?; } - // 2. 创建取消令牌 let cancel_token = state.create_cancel_token(&session_id).await; - // 3. 构建消息和配置 let user_message = Message::user().with_text(&message); let session_config = SessionConfigBuilder::new(&session_id).build(); - // 4. 获取 Agent 引用(关键步骤) let agent_arc = state.get_agent_arc(); let guard = agent_arc.read().await; let agent = guard.as_ref().ok_or("Agent not initialized")?; - // 5. 调用 Agent::reply let stream_result = agent .reply(user_message, session_config, Some(cancel_token.clone())) .await; - // 6. 处理流式响应 match stream_result { Ok(mut stream) => { while let Some(event_result) = stream.next().await { match event_result { Ok(agent_event) => { - // 转换并发送事件到前端 - let tauri_events = - crate::agent::event_converter::convert_agent_event(agent_event); + let tauri_events = convert_agent_event(agent_event); for tauri_event in tauri_events { - if let Err(e) = app.emit(&event_name, &tauri_event) { - tracing::error!("[AsterAgentWrapper] 发送事件失败: {}", e); + if let Err(error) = app.emit(&event_name, &tauri_event) { + tracing::error!("[AsterAgentWrapper] 发送事件失败: {}", error); } } } - Err(e) => { - // 发送错误事件 - let error_event = - crate::agent::event_converter::TauriAgentEvent::Error { - message: format!("Stream error: {e}"), - }; + Err(error) => { + let error_event = TauriAgentEvent::Error { + message: format!("Stream error: {error}"), + }; let _ = app.emit(&event_name, &error_event); } } } - // 发送完成事件 - let done_event = - crate::agent::event_converter::TauriAgentEvent::FinalDone { usage: None }; + let done_event = TauriAgentEvent::FinalDone { usage: None }; let _ = app.emit(&event_name, &done_event); } - Err(e) => { - // 发送错误事件并返回错误 - let error_event = crate::agent::event_converter::TauriAgentEvent::Error { - message: format!("Agent error: {e}"), + Err(error) => { + let error_event = TauriAgentEvent::Error { + message: format!("Agent error: {error}"), }; let _ = app.emit(&event_name, &error_event); - return Err(format!("Agent error: {e}")); + return Err(format!("Agent error: {error}")); } } - // guard 在作用域结束时自动释放 - - // 7. 清理取消令牌 state.remove_cancel_token(&session_id).await; Ok(()) @@ -113,133 +98,19 @@ impl AsterAgentWrapper { state.cancel_session(session_id).await } - /// 创建新会话 - 使用 ProxyCast 数据库 + /// 创建新会话 pub fn create_session_sync(db: &DbConnection, name: Option) -> Result { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - let session_name = name.unwrap_or_else(|| "新对话".to_string()); - let session_id = uuid::Uuid::new_v4().to_string(); - let now = Utc::now().to_rfc3339(); - - let session = crate::agent::types::AgentSession { - id: session_id.clone(), - model: "agent:default".to_string(), - messages: Vec::new(), - system_prompt: None, - title: Some(session_name), - created_at: now.clone(), - updated_at: now, - }; - - AgentDao::create_session(&conn, &session).map_err(|e| format!("创建会话失败: {e}"))?; - - Ok(session_id) + proxycast_agent::session_store::create_session_sync(db, name) } - /// 列出所有会话 - 使用 ProxyCast 数据库 + /// 列出所有会话 pub fn list_sessions_sync(db: &DbConnection) -> Result, String> { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let sessions = - AgentDao::list_sessions(&conn).map_err(|e| format!("获取会话列表失败: {e}"))?; - - Ok(sessions - .into_iter() - .map(|s| { - let messages_count = AgentDao::get_message_count(&conn, &s.id).unwrap_or(0); - SessionInfo { - id: s.id, - name: s.title.unwrap_or_else(|| "未命名".to_string()), - created_at: chrono::DateTime::parse_from_rfc3339(&s.created_at) - .map(|dt| dt.timestamp()) - .unwrap_or(0), - updated_at: chrono::DateTime::parse_from_rfc3339(&s.updated_at) - .map(|dt| dt.timestamp()) - .unwrap_or(0), - messages_count, - } - }) - .collect()) + proxycast_agent::session_store::list_sessions_sync(db) } - /// 获取会话详情 - 使用 ProxyCast 数据库 + /// 获取会话详情 pub fn get_session_sync(db: &DbConnection, session_id: &str) -> Result { - let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; - - let session = AgentDao::get_session(&conn, session_id) - .map_err(|e| format!("获取会话失败: {e}"))? - .ok_or_else(|| format!("会话不存在: {session_id}"))?; - - let messages = - AgentDao::get_messages(&conn, session_id).map_err(|e| format!("获取消息失败: {e}"))?; - - Ok(SessionDetail { - id: session.id, - name: session.title.unwrap_or_else(|| "未命名".to_string()), - created_at: chrono::DateTime::parse_from_rfc3339(&session.created_at) - .map(|dt| dt.timestamp()) - .unwrap_or(0), - updated_at: chrono::DateTime::parse_from_rfc3339(&session.updated_at) - .map(|dt| dt.timestamp()) - .unwrap_or(0), - messages: messages - .into_iter() - .map(|m| convert_agent_message(&m)) - .collect(), - }) - } -} - -/// 会话信息(简化版) -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct SessionInfo { - pub id: String, - pub name: String, - pub created_at: i64, - pub updated_at: i64, - pub messages_count: usize, -} - -/// 会话详情(包含消息) -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] -pub struct SessionDetail { - pub id: String, - pub name: String, - pub created_at: i64, - pub updated_at: i64, - pub messages: Vec, -} - -/// 将 AgentMessage 转换为 TauriMessage -fn convert_agent_message( - msg: &crate::agent::types::AgentMessage, -) -> crate::agent::event_converter::TauriMessage { - use crate::agent::event_converter::{TauriMessage, TauriMessageContent}; - use crate::agent::types::MessageContent; - - let content = match &msg.content { - MessageContent::Text(text) => vec![TauriMessageContent::Text { text: text.clone() }], - MessageContent::Parts(parts) => parts - .iter() - .filter_map(|p| { - if let crate::agent::types::ContentPart::Text { text } = p { - Some(TauriMessageContent::Text { text: text.clone() }) - } else { - None - } - }) - .collect(), - }; - - // 解析时间戳 - let timestamp = chrono::DateTime::parse_from_rfc3339(&msg.timestamp) - .map(|dt| dt.timestamp()) - .unwrap_or(0); - - TauriMessage { - id: None, - role: msg.role.clone(), - content, - timestamp, + proxycast_agent::session_store::get_session_sync(db, session_id) } } diff --git a/src-tauri/src/agent/aster_state.rs b/src-tauri/src/agent/aster_state.rs index 5c2e4c72b..e4bef1c18 100644 --- a/src-tauri/src/agent/aster_state.rs +++ b/src-tauri/src/agent/aster_state.rs @@ -1,753 +1,7 @@ -//! Aster Agent 状态管理 +//! Aster Agent 状态管理(桥接层) //! -//! 管理 Aster Agent 实例和相关状态 -//! 提供 Tauri 应用与 Aster 框架的桥接 -//! 支持从 ProxyCast 凭证池自动选择凭证 -//! -//! ## 重要:SessionStore 注入 -//! -//! 为了让 Aster Agent 的消息存储到 ProxyCast 数据库,必须在创建 Agent 时 -//! 注入 `ProxyCastSessionStore`。使用 `init_agent_with_db()` 方法而不是 `init_agent()`。 -//! -//! ## Agent 身份配置 -//! -//! 通过 Aster 框架的 `AgentIdentity` API 设置 ProxyCast 专属的 Agent 身份, -//! 包括名称、语言偏好、产品描述等。这是架构层面的正确做法, -//! 而不是简单地追加提示词。 -//! -//! ## Skills 集成 -//! -//! Agent 初始化时会自动加载 `~/.proxycast/skills/` 目录下的 Skills 到 -//! aster-rust 的 global_registry,使 AI 能够自动发现和调用这些 Skills。 -//! -//! 参考文档:`docs/prd/chat-architecture-redesign.md` +//! 纯逻辑已迁移到 `proxycast-agent` crate, +//! 本模块保留兼容导出。 -use aster::agents::{Agent, AgentIdentity, SessionConfig}; -use aster::model::ModelConfig; -use aster::skills::{global_registry, load_skills_from_directory, SkillSource}; -use std::sync::Arc; -use tokio::sync::RwLock; -use tokio_util::sync::CancellationToken; - -use crate::agent::credential_bridge::{ - create_aster_provider, AsterProviderConfig, CredentialBridge, -}; -use crate::database::DbConnection; -use crate::services::aster_session_store::ProxyCastSessionStore; -use crate::services::project_context_builder::ProjectContextBuilder; - -/// Provider 配置信息 -#[derive(Debug, Clone)] -pub struct ProviderConfig { - /// Provider 名称 (openai, anthropic, google, ollama 等) - pub provider_name: String, - /// 模型名称 - pub model_name: String, - /// API Key (可选,某些 provider 从环境变量读取) - pub api_key: Option, - /// Base URL (可选,用于自定义端点) - pub base_url: Option, - /// 凭证 UUID(来自凭证池,用于记录使用和健康状态) - pub credential_uuid: Option, -} - -/// Aster Agent 全局状态 -/// -/// 在 Tauri 应用中作为 managed state 使用 -pub struct AsterAgentState { - /// Aster Agent 实例 - agent: Arc>>, - /// 当前活跃的取消令牌(用于中止正在进行的对话) - cancel_tokens: Arc>>, - /// 当前 Provider 配置 - current_provider_config: Arc>>, - /// 凭证桥接器 - credential_bridge: CredentialBridge, -} - -impl Default for AsterAgentState { - fn default() -> Self { - Self::new() - } -} - -impl AsterAgentState { - /// 创建新的 Aster Agent 状态 - pub fn new() -> Self { - Self { - agent: Arc::new(RwLock::new(None)), - cancel_tokens: Arc::new(RwLock::new(std::collections::HashMap::new())), - current_provider_config: Arc::new(RwLock::new(None)), - credential_bridge: CredentialBridge::new(), - } - } - - /// 初始化 Agent(带数据库连接) - /// - /// 创建 Agent 并注入 ProxyCastSessionStore,确保消息存储到 ProxyCast 数据库。 - /// 同时设置 ProxyCast 专属的 Agent 身份(名称、语言、描述)。 - /// 自动加载 `~/.proxycast/skills/` 目录下的 Skills 到 aster-rust 的 global_registry。 - /// - /// **推荐使用此方法**而不是 `init_agent()`。 - /// - /// # 参数 - /// - `db`: 数据库连接,用于创建 SessionStore - pub async fn init_agent_with_db(&self, db: &DbConnection) -> Result<(), String> { - let mut agent_guard = self.agent.write().await; - if agent_guard.is_none() { - // 创建 SessionStore - let session_store = Arc::new(ProxyCastSessionStore::new(db.clone())); - tracing::info!("[AsterAgent] 创建 ProxyCastSessionStore 成功"); - - // 创建 Agent 并注入 SessionStore - let agent = Agent::new().with_session_store(session_store); - - // 验证 session_store 是否被正确设置 - let has_store = agent.session_store().is_some(); - tracing::info!( - "[AsterAgent] Agent 创建完成,session_store 已设置: {}", - has_store - ); - - // 使用异步方法设置 ProxyCast 专属身份 - let identity = Self::create_proxycast_identity(); - agent.set_identity(identity).await; - - // 加载 ProxyCast Skills 到 aster-rust 的 global_registry - Self::load_proxycast_skills(); - - *agent_guard = Some(agent); - tracing::info!( - "[AsterAgent] Agent 初始化成功,已注入 ProxyCastSessionStore、ProxyCast 身份和 Skills" - ); - } else { - tracing::debug!("[AsterAgent] Agent 已初始化,跳过"); - } - Ok(()) - } - - /// 加载 ProxyCast Skills 到 aster-rust 的 global_registry - /// - /// 从 `~/.proxycast/skills/` 目录加载 Skills,使 AI 能够自动发现和调用。 - fn load_proxycast_skills() { - let home = match dirs::home_dir() { - Some(h) => h, - None => { - tracing::warn!("[AsterAgent] 无法获取 home 目录,跳过 Skills 加载"); - return; - } - }; - - let skills_dir = home.join(".proxycast").join("skills"); - if !skills_dir.exists() { - tracing::info!( - "[AsterAgent] ProxyCast Skills 目录不存在: {:?},跳过加载", - skills_dir - ); - return; - } - - // 从 ProxyCast skills 目录加载 Skills - let skills = load_skills_from_directory(&skills_dir, SkillSource::User); - let skill_count = skills.len(); - - if skill_count == 0 { - tracing::info!("[AsterAgent] ProxyCast Skills 目录为空,无 Skills 可加载"); - return; - } - - // 注册到 global_registry - let registry = global_registry(); - if let Ok(mut registry_guard) = registry.write() { - for skill in skills { - let skill_name = skill.skill_name.clone(); - registry_guard.register(skill); - tracing::debug!("[AsterAgent] 已注册 Skill: {}", skill_name); - } - tracing::info!( - "[AsterAgent] 成功加载 {} 个 ProxyCast Skills 到 global_registry", - skill_count - ); - } else { - tracing::error!("[AsterAgent] 无法获取 global_registry 写锁,Skills 加载失败"); - } - } - - /// 重新加载 ProxyCast Skills - /// - /// 当用户安装或卸载 Skills 后调用此方法刷新 registry。 - pub fn reload_proxycast_skills() { - Self::load_proxycast_skills(); - } - - /// 创建 ProxyCast 专属的 Agent 身份配置 - fn create_proxycast_identity() -> AgentIdentity { - AgentIdentity::new("ProxyCast 助手") - .with_language("Chinese") - .with_description( - "ProxyCast 是一个 AI 代理服务应用,帮助用户管理和使用各种 AI 模型的凭证。", - ) - .with_custom_prompt(PROXYCAST_IDENTITY_PROMPT.to_string()) - } - - /// 初始化 Agent(无数据库版本) - /// - /// **警告**:此方法创建的 Agent 不会将消息存储到 ProxyCast 数据库, - /// 消息会存储到 Aster 默认的 `~/.aster/sessions.db`。 - /// - /// 建议使用 `init_agent_with_db()` 代替。 - #[deprecated( - since = "0.1.0", - note = "请使用 init_agent_with_db() 以确保消息存储到 ProxyCast 数据库" - )] - pub async fn init_agent(&self) -> Result<(), String> { - let mut agent_guard = self.agent.write().await; - if agent_guard.is_none() { - let agent = Agent::new(); - *agent_guard = Some(agent); - tracing::warn!( - "[AsterAgent] Agent 初始化(无 SessionStore),消息将存储到 Aster 默认数据库" - ); - } - Ok(()) - } - - /// 配置 Provider - /// - /// 根据配置创建并设置 Provider - /// - /// # 参数 - /// - `config`: Provider 配置 - /// - `session_id`: 会话 ID - /// - `db`: 数据库连接(用于初始化 Agent) - pub async fn configure_provider( - &self, - config: ProviderConfig, - session_id: &str, - db: &DbConnection, - ) -> Result<(), String> { - // 确保 Agent 已初始化(使用带数据库的版本) - self.init_agent_with_db(db).await?; - - // 设置环境变量(Aster 的 provider 从环境变量读取配置) - self.set_provider_env_vars(&config); - - // 创建 ModelConfig - let model_config = ModelConfig::new(&config.model_name) - .map_err(|e| format!("创建 ModelConfig 失败: {e}"))?; - - // 创建 Provider - let provider = aster::providers::create(&config.provider_name, model_config) - .await - .map_err(|e| format!("创建 Provider 失败: {e}"))?; - - // 更新 Agent 的 Provider - let agent_guard = self.agent.read().await; - if let Some(agent) = agent_guard.as_ref() { - agent - .update_provider(provider, session_id) - .await - .map_err(|e| format!("更新 Provider 失败: {e}"))?; - } - - // 保存当前配置 - let mut config_guard = self.current_provider_config.write().await; - *config_guard = Some(config.clone()); - - tracing::info!( - "[AsterAgent] Provider 配置成功: {} / {}", - config.provider_name, - config.model_name - ); - - Ok(()) - } - - /// 从凭证池配置 Provider - /// - /// 自动从 ProxyCast 凭证池选择可用凭证并配置 Aster Provider - /// - /// # 参数 - /// - `db`: 数据库连接 - /// - `provider_type`: Provider 类型 (openai, anthropic, kiro 等) - /// - `model`: 模型名称 - /// - `session_id`: 会话 ID - pub async fn configure_provider_from_pool( - &self, - db: &DbConnection, - provider_type: &str, - model: &str, - session_id: &str, - ) -> Result { - // 确保 Agent 已初始化(使用带数据库的版本) - self.init_agent_with_db(db).await?; - - // 从凭证池选择凭证并获取配置 - let aster_config = self - .credential_bridge - .select_and_configure(db, provider_type, model) - .await - .map_err(|e| format!("从凭证池选择凭证失败: {e}"))?; - - // 创建 Provider - let provider = create_aster_provider(&aster_config) - .await - .map_err(|e| format!("创建 Provider 失败: {e}"))?; - - // 更新 Agent 的 Provider - let agent_guard = self.agent.read().await; - if let Some(agent) = agent_guard.as_ref() { - agent - .update_provider(provider, session_id) - .await - .map_err(|e| format!("更新 Provider 失败: {e}"))?; - } - - // 保存当前配置 - let config = ProviderConfig { - provider_name: aster_config.provider_name.clone(), - model_name: aster_config.model_name.clone(), - api_key: aster_config.api_key.clone(), - base_url: aster_config.base_url.clone(), - credential_uuid: Some(aster_config.credential_uuid.clone()), - }; - let mut config_guard = self.current_provider_config.write().await; - *config_guard = Some(config); - - // 记录凭证使用 - if let Err(e) = self - .credential_bridge - .record_usage(db, &aster_config.credential_uuid) - { - tracing::warn!("[AsterAgent] 记录凭证使用失败: {}", e); - } - - tracing::info!( - "[AsterAgent] 从凭证池配置 Provider 成功: {} / {} (凭证: {})", - aster_config.provider_name, - aster_config.model_name, - aster_config.credential_uuid - ); - - Ok(aster_config) - } - - /// 标记当前凭证为健康 - pub fn mark_current_healthy(&self, db: &DbConnection, model: Option<&str>) { - if let Ok(config_guard) = self.current_provider_config.try_read() { - if let Some(config) = config_guard.as_ref() { - if let Some(uuid) = &config.credential_uuid { - if let Err(e) = self.credential_bridge.mark_healthy(db, uuid, model) { - tracing::warn!("[AsterAgent] 标记凭证健康失败: {}", e); - } - } - } - } - } - - /// 标记当前凭证为不健康 - pub fn mark_current_unhealthy(&self, db: &DbConnection, error: Option<&str>) { - if let Ok(config_guard) = self.current_provider_config.try_read() { - if let Some(config) = config_guard.as_ref() { - if let Some(uuid) = &config.credential_uuid { - if let Err(e) = self.credential_bridge.mark_unhealthy(db, uuid, error) { - tracing::warn!("[AsterAgent] 标记凭证不健康失败: {}", e); - } - } - } - } - } - - /// 设置 Provider 相关的环境变量 - fn set_provider_env_vars(&self, config: &ProviderConfig) { - tracing::info!( - "[AsterAgent] set_provider_env_vars: provider_name={}, model_name={}, has_api_key={}, base_url={:?}", - config.provider_name, - config.model_name, - config.api_key.is_some(), - config.base_url - ); - - // 根据 provider 类型设置对应的环境变量 - let env_key = match config.provider_name.as_str() { - "openai" => "OPENAI_API_KEY", - "anthropic" => "ANTHROPIC_API_KEY", - "google" => "GOOGLE_API_KEY", - "deepseek" | "custom_deepseek" => "OPENAI_API_KEY", // DeepSeek 使用 OpenAI 兼容 API - "groq" => "OPENAI_API_KEY", // Groq 使用 OpenAI 兼容 API - "mistral" => "OPENAI_API_KEY", // Mistral 使用 OpenAI 兼容 API - "openrouter" => "OPENROUTER_API_KEY", - "ollama" => return, // Ollama 不需要 API Key - _ => { - tracing::warn!( - "[AsterAgent] 未知的 provider_name: {}, 使用通用 OpenAI 格式", - config.provider_name - ); - // 通用 OpenAI 兼容格式 - if let Some(api_key) = &config.api_key { - std::env::set_var("OPENAI_API_KEY", api_key); - } - if let Some(base_url) = &config.base_url { - std::env::set_var("OPENAI_BASE_URL", base_url); - } - return; - } - }; - - tracing::info!("[AsterAgent] 设置环境变量: {}=***", env_key); - - if let Some(api_key) = &config.api_key { - std::env::set_var(env_key, api_key); - } - - if let Some(base_url) = &config.base_url { - let base_url_key = format!( - "{}_BASE_URL", - config.provider_name.to_uppercase().replace("_", "") - ); - std::env::set_var(base_url_key, base_url); - } - } - - /// 获取当前 Provider 配置 - pub async fn get_provider_config(&self) -> Option { - self.current_provider_config.read().await.clone() - } - - /// 清除当前 Provider 配置 - /// - /// 用于切换凭证后重置状态,下次对话时会重新从凭证池选择凭证 - pub async fn clear_provider_config(&self) { - let mut config_guard = self.current_provider_config.write().await; - *config_guard = None; - tracing::info!("[AsterAgent] Provider 配置已清除"); - } - - /// 检查 Provider 是否已配置 - pub async fn is_provider_configured(&self) -> bool { - self.current_provider_config.read().await.is_some() - } - - /// 获取 Agent 的只读引用并执行同步操作 - pub async fn with_agent(&self, f: F) -> Result - where - F: FnOnce(&Agent) -> R, - { - let guard = self.agent.read().await; - match guard.as_ref() { - Some(agent) => Ok(f(agent)), - None => Err("Agent not initialized".to_string()), - } - } - - /// 获取 Agent 的可变引用并执行同步操作 - pub async fn with_agent_mut(&self, f: F) -> Result - where - F: FnOnce(&mut Agent) -> R, - { - let mut guard = self.agent.write().await; - match guard.as_mut() { - Some(agent) => Ok(f(agent)), - None => Err("Agent not initialized".to_string()), - } - } - - /// 获取 Agent 的 Arc 引用 - /// - /// 用于需要长期持有 Agent 引用的场景 - pub fn get_agent_arc(&self) -> Arc>> { - self.agent.clone() - } - - /// 创建新的取消令牌 - pub async fn create_cancel_token(&self, session_id: &str) -> CancellationToken { - let token = CancellationToken::new(); - let mut tokens = self.cancel_tokens.write().await; - tokens.insert(session_id.to_string(), token.clone()); - token - } - - /// 取消指定会话的操作 - pub async fn cancel_session(&self, session_id: &str) -> bool { - let tokens = self.cancel_tokens.read().await; - if let Some(token) = tokens.get(session_id) { - token.cancel(); - true - } else { - false - } - } - - /// 移除取消令牌 - pub async fn remove_cancel_token(&self, session_id: &str) { - let mut tokens = self.cancel_tokens.write().await; - tokens.remove(session_id); - } - - // ------------------------------------------------------------------------ - // 项目上下文支持 - // ------------------------------------------------------------------------ - - /// 构建带项目上下文的 System Prompt - /// - /// 加载项目的人设、素材、模板配置,构建完整的 AI 提示词。 - /// - /// # 参数 - /// - `db`: 数据库连接 - /// - `project_id`: 项目 ID - /// - /// # 返回 - /// - 成功返回构建好的 System Prompt - /// - 失败返回错误信息 - pub fn build_project_system_prompt( - db: &DbConnection, - project_id: &str, - ) -> Result { - let conn = db - .lock() - .map_err(|e| format!("获取数据库连接失败: {}", e))?; - let context = ProjectContextBuilder::build_context(&conn, project_id) - .map_err(|e| format!("构建项目上下文失败: {}", e))?; - Ok(ProjectContextBuilder::build_system_prompt(&context)) - } - - /// 创建带项目上下文的会话配置 - /// - /// 自动加载项目配置并构建 SessionConfig。 - /// - /// # 参数 - /// - `db`: 数据库连接 - /// - `session_id`: 会话 ID - /// - `project_id`: 项目 ID(可选,如果为 None 则不注入项目上下文) - /// - /// # 返回 - /// - 构建好的 SessionConfig - pub fn create_session_config_with_project( - db: &DbConnection, - session_id: &str, - project_id: Option<&str>, - ) -> SessionConfig { - let system_prompt = - project_id.and_then(|pid| Self::build_project_system_prompt(db, pid).ok()); - - SessionConfigBuilder::new(session_id) - .system_prompt(system_prompt.unwrap_or_default()) - .build() - } - - /// 检查 Agent 是否已初始化 - pub async fn is_initialized(&self) -> bool { - self.agent.read().await.is_some() - } -} - -/// 会话配置构建器 -/// -/// 用于构建 Aster SessionConfig -pub struct SessionConfigBuilder { - id: String, - max_turns: Option, - system_prompt: Option, -} - -impl SessionConfigBuilder { - pub fn new(id: impl Into) -> Self { - Self { - id: id.into(), - max_turns: None, - system_prompt: None, - } - } - - pub fn max_turns(mut self, turns: u32) -> Self { - self.max_turns = Some(turns); - self - } - - pub fn system_prompt(mut self, prompt: impl Into) -> Self { - self.system_prompt = Some(prompt.into()); - self - } - - pub fn build(self) -> SessionConfig { - SessionConfig { - id: self.id, - schedule_id: None, - max_turns: self.max_turns, - retry_config: None, - system_prompt: self.system_prompt, - } - } -} - -/// 消息构建辅助函数 -pub mod message_helpers { - use aster::conversation::message::Message; - - /// 创建用户文本消息 - pub fn user_text(text: impl Into) -> Message { - Message::user().with_text(text) - } - - /// 创建助手文本消息 - pub fn assistant_text(text: impl Into) -> Message { - Message::assistant().with_text(text) - } -} - -// ============================================================================= -// ProxyCast Agent 身份提示词 -// ============================================================================= - -/// ProxyCast 专属的 Agent 身份提示词 -/// -/// 这是完整的身份定义,会替换 Aster 框架默认的 "aster by Block" 身份。 -/// 框架的能力描述(Extensions、Response Guidelines)会自动追加。 -const PROXYCAST_IDENTITY_PROMPT: &str = r#"你是 ProxyCast 助手,一个专业、友好的 AI 技术伙伴。 - -## 关于 ProxyCast - -ProxyCast 是一个 AI 代理服务应用,帮助用户: -- 管理多个 AI 模型提供商的凭证(OpenAI、Claude、Gemini、Kiro 等) -- 通过统一的 API 接口访问不同的 AI 模型 -- 实现凭证池的负载均衡和健康检查 - -## 语言规范 - -1. **始终使用中文回复**:除非用户明确要求使用其他语言 -2. **代码注释使用中文**:生成代码时,注释应使用中文 -3. **技术术语保持原文**:API、JSON、HTTP、Token 等专业术语保持英文 - -## 交互风格 - -- 简洁专业,直接给出解决方案 -- 友好但不啰嗦,像经验丰富的技术伙伴 -- 遇到问题时,先分析原因再提供方案 -"#; - -#[cfg(test)] -mod tests { - use super::*; - use std::fs; - use tempfile::TempDir; - - #[tokio::test] - async fn test_aster_state_init() { - let state = AsterAgentState::new(); - assert!(!state.is_initialized().await); - - #[allow(deprecated)] - state.init_agent().await.unwrap(); - assert!(state.is_initialized().await); - } - - #[tokio::test] - async fn test_cancel_token() { - let state = AsterAgentState::new(); - let session_id = "test-session"; - - let token = state.create_cancel_token(session_id).await; - assert!(!token.is_cancelled()); - - assert!(state.cancel_session(session_id).await); - assert!(token.is_cancelled()); - - state.remove_cancel_token(session_id).await; - assert!(!state.cancel_session(session_id).await); - } - - // ========================================================================= - // Skills 集成测试 - // ========================================================================= - - /// 测试辅助函数:创建测试用的 Skill 目录 - fn create_test_skill(skills_dir: &std::path::Path, skill_name: &str, description: &str) { - let skill_path = skills_dir.join(skill_name); - fs::create_dir_all(&skill_path).unwrap(); - let skill_md = format!( - r#"--- -name: {} -description: {} ---- - -# {} - -这是一个测试 Skill。 -"#, - skill_name, description, skill_name - ); - fs::write(skill_path.join("SKILL.md"), skill_md).unwrap(); - } - - /// 测试:load_skills_from_directory 能正确加载 Skills - #[test] - fn test_load_skills_from_directory() { - let temp_dir = TempDir::new().unwrap(); - let skills_dir = temp_dir.path(); - - // 创建测试 Skills - create_test_skill(skills_dir, "test-skill-1", "第一个测试技能"); - create_test_skill(skills_dir, "test-skill-2", "第二个测试技能"); - - // 加载 Skills - let skills = load_skills_from_directory(skills_dir, SkillSource::User); - - // 验证 - assert_eq!(skills.len(), 2); - let names: Vec<_> = skills.iter().map(|s| s.display_name.as_str()).collect(); - assert!(names.contains(&"test-skill-1")); - assert!(names.contains(&"test-skill-2")); - } - - /// 测试:空目录返回空列表 - #[test] - fn test_load_skills_empty_directory() { - let temp_dir = TempDir::new().unwrap(); - let skills = load_skills_from_directory(temp_dir.path(), SkillSource::User); - assert!(skills.is_empty()); - } - - /// 测试:不存在的目录返回空列表 - #[test] - fn test_load_skills_nonexistent_directory() { - let nonexistent = std::path::Path::new("/nonexistent/path/to/skills"); - let skills = load_skills_from_directory(nonexistent, SkillSource::User); - assert!(skills.is_empty()); - } - - /// 测试:global_registry 能正确注册和查找 Skills - #[test] - fn test_global_registry_register_and_find() { - let temp_dir = TempDir::new().unwrap(); - let skills_dir = temp_dir.path(); - - // 创建测试 Skill - create_test_skill(skills_dir, "registry-test-skill", "注册表测试技能"); - - // 加载并注册到 global_registry - let skills = load_skills_from_directory(skills_dir, SkillSource::User); - let registry = global_registry(); - - if let Ok(mut registry_guard) = registry.write() { - for skill in skills { - registry_guard.register(skill); - } - } - - // 验证能找到注册的 Skill - if let Ok(registry_guard) = registry.read() { - let found = registry_guard.find("registry-test-skill"); - assert!(found.is_some()); - assert_eq!(found.unwrap().display_name, "registry-test-skill"); - } - } - - /// 测试:reload_proxycast_skills 不会 panic(即使目录不存在) - #[test] - fn test_reload_proxycast_skills_no_panic() { - // 这个测试确保 reload_proxycast_skills 在各种情况下都不会 panic - // 即使 ~/.proxycast/skills/ 目录不存在 - AsterAgentState::reload_proxycast_skills(); - // 如果没有 panic,测试通过 - } -} +pub use proxycast_agent::aster_state::{AsterAgentState, ProviderConfig}; +pub use proxycast_agent::aster_state_support::{message_helpers, SessionConfigBuilder}; diff --git a/src-tauri/src/agent/credential_bridge.rs b/src-tauri/src/agent/credential_bridge.rs index f04aa392a..904dbdb34 100644 --- a/src-tauri/src/agent/credential_bridge.rs +++ b/src-tauri/src/agent/credential_bridge.rs @@ -1,561 +1,6 @@ -//! 凭证池桥接模块 +//! 凭证池桥接模块(重导出层) //! -//! 将 ProxyCast 凭证池与 Aster Provider 系统连接 -//! 支持从凭证池自动选择凭证并配置 Aster Provider -//! -//! ## 功能 -//! - 从凭证池选择可用凭证 -//! - 将凭证转换为 Aster Provider 配置 -//! - 支持 OAuth 和 API Key 两种凭证类型 -//! - 自动刷新过期的 OAuth Token -//! - 智能拆分 base_url 为 host + path,避免路径重复(如智谱 /v4/v1 问题) +//! 纯逻辑已迁移到 `proxycast-agent` crate, +//! 本模块仅保留兼容导出。 -use crate::database::DbConnection; -use crate::models::provider_pool_model::{CredentialData, PoolProviderType, ProviderCredential}; -use crate::services::api_key_provider_service::ApiKeyProviderService; -use crate::services::provider_pool_service::ProviderPoolService; -use aster::model::ModelConfig; -use aster::providers::base::Provider; -use std::sync::Arc; - -/// 凭证桥接错误 -#[derive(Debug, Clone)] -pub enum CredentialBridgeError { - /// 没有可用凭证 - NoCredentials(String), - /// 凭证类型不支持 - UnsupportedCredentialType(String), - /// Provider 创建失败 - ProviderCreationFailed(String), - /// Token 刷新失败 - TokenRefreshFailed(String), - /// 数据库错误 - DatabaseError(String), -} - -impl std::fmt::Display for CredentialBridgeError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - Self::NoCredentials(msg) => write!(f, "没有可用凭证: {msg}"), - Self::UnsupportedCredentialType(msg) => write!(f, "不支持的凭证类型: {msg}"), - Self::ProviderCreationFailed(msg) => write!(f, "Provider 创建失败: {msg}"), - Self::TokenRefreshFailed(msg) => write!(f, "Token 刷新失败: {msg}"), - Self::DatabaseError(msg) => write!(f, "数据库错误: {msg}"), - } - } -} - -impl std::error::Error for CredentialBridgeError {} - -/// Aster Provider 配置 -#[derive(Debug, Clone)] -pub struct AsterProviderConfig { - /// Provider 名称 (openai, anthropic, google 等) - pub provider_name: String, - /// 模型名称 - pub model_name: String, - /// API Key - pub api_key: Option, - /// Base URL - pub base_url: Option, - /// 凭证 UUID(用于记录使用和健康状态) - pub credential_uuid: String, -} - -/// 凭证池桥接器 -/// -/// 负责从 ProxyCast 凭证池选择凭证并转换为 Aster Provider 配置 -pub struct CredentialBridge { - pool_service: ProviderPoolService, - api_key_service: ApiKeyProviderService, -} - -impl Default for CredentialBridge { - fn default() -> Self { - Self::new() - } -} - -impl CredentialBridge { - pub fn new() -> Self { - Self { - pool_service: ProviderPoolService::new(), - api_key_service: ApiKeyProviderService::new(), - } - } - - /// 从凭证池选择凭证并创建 Aster Provider 配置 - /// - /// # 参数 - /// - `db`: 数据库连接 - /// - `provider_type`: Provider 类型 (openai, anthropic, kiro, deepseek 等) - /// - `model`: 模型名称 - /// - /// # 返回 - /// 成功时返回 AsterProviderConfig,失败时返回错误 - pub async fn select_and_configure( - &self, - db: &DbConnection, - provider_type: &str, - model: &str, - ) -> Result { - // 1. 从凭证池选择凭证 - // 将 provider_type 同时作为 provider_id_hint 传递,支持 60+ API Key Provider - // 例如 "deepseek", "moonshot", "qwen" 等 - let credential = self - .pool_service - .select_credential_with_fallback( - db, - &self.api_key_service, - provider_type, - Some(model), - Some(provider_type), // 传递 provider_id_hint 支持智能降级 - None, - ) - .await - .map_err(CredentialBridgeError::DatabaseError)? - .ok_or_else(|| { - CredentialBridgeError::NoCredentials(format!( - "没有找到 {provider_type} 类型的可用凭证" - )) - })?; - - // 2. 转换为 Aster Provider 配置,传递 provider_type 以便正确识别 Provider - self.credential_to_config(&credential, model, provider_type, db) - .await - } - - /// 将 ProxyCast 凭证转换为 Aster Provider 配置 - async fn credential_to_config( - &self, - credential: &ProviderCredential, - model: &str, - provider_type_hint: &str, - db: &DbConnection, - ) -> Result { - tracing::info!( - "[CredentialBridge] credential_to_config: provider_type_hint={}, credential_type={:?}", - provider_type_hint, - credential.provider_type - ); - - let (provider_name, api_key, base_url) = match &credential.credential { - // OpenAI API Key - 根据 provider_type_hint 确定实际的 Provider - CredentialData::OpenAIKey { api_key, base_url } => { - // 使用 provider_type_hint 来确定 aster provider 名称 - let provider = map_provider_type_to_aster(provider_type_hint); - tracing::info!( - "[CredentialBridge] OpenAIKey: provider_type_hint={} -> aster_provider={}", - provider_type_hint, - provider - ); - ( - provider.to_string(), - Some(api_key.clone()), - base_url.clone(), - ) - } - - // Claude/Anthropic API Key - CredentialData::ClaudeKey { api_key, base_url } - | CredentialData::AnthropicKey { api_key, base_url } => ( - "anthropic".to_string(), - Some(api_key.clone()), - base_url.clone(), - ), - - // Kiro OAuth - 需要获取 access_token - CredentialData::KiroOAuth { creds_file_path } => { - let token = self - .get_kiro_token(creds_file_path, db, &credential.uuid) - .await?; - // Kiro 使用 CodeWhisperer API,映射到 bedrock provider - ("bedrock".to_string(), Some(token), None) - } - - // Gemini OAuth - CredentialData::GeminiOAuth { - creds_file_path, .. - } => { - let token = self.get_oauth_token(creds_file_path).await?; - ("google".to_string(), Some(token), None) - } - - // Gemini API Key - CredentialData::GeminiApiKey { - api_key, base_url, .. - } => ( - "google".to_string(), - Some(api_key.clone()), - base_url.clone(), - ), - - // Vertex AI - CredentialData::VertexKey { - api_key, base_url, .. - } => ( - "gcpvertexai".to_string(), - Some(api_key.clone()), - base_url.clone(), - ), - - // Codex OAuth - CredentialData::CodexOAuth { - creds_file_path, - api_base_url, - } => { - let token = self.get_codex_token(creds_file_path).await?; - ("codex".to_string(), Some(token), api_base_url.clone()) - } - - // Claude OAuth - CredentialData::ClaudeOAuth { creds_file_path } => { - let token = self.get_oauth_token(creds_file_path).await?; - ("anthropic".to_string(), Some(token), None) - } - - // Antigravity OAuth - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => { - let token = self.get_oauth_token(creds_file_path).await?; - ("google".to_string(), Some(token), None) - } - }; - - Ok(AsterProviderConfig { - provider_name, - model_name: model.to_string(), - api_key, - base_url, - credential_uuid: credential.uuid.clone(), - }) - } - - /// 获取 Kiro OAuth Token - async fn get_kiro_token( - &self, - creds_path: &str, - _db: &DbConnection, - _uuid: &str, - ) -> Result { - use crate::providers::kiro::KiroProvider; - - let mut provider = KiroProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("加载 Kiro 凭证失败: {e}")) - })?; - - // 检查 token 是否过期,如果过期则刷新 - if provider.is_token_expired() { - tracing::info!("[CredentialBridge] Kiro token 已过期,尝试刷新"); - self.pool_service - .refresh_kiro_token(creds_path) - .await - .map_err(CredentialBridgeError::TokenRefreshFailed)?; - - // 重新加载凭证 - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("重新加载凭证失败: {e}")) - })?; - } - - provider.credentials.access_token.ok_or_else(|| { - CredentialBridgeError::TokenRefreshFailed("缺少 access_token".to_string()) - }) - } - - /// 获取通用 OAuth Token - async fn get_oauth_token(&self, creds_path: &str) -> Result { - let content = std::fs::read_to_string(creds_path).map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("读取凭证文件失败: {e}")) - })?; - - let creds: serde_json::Value = serde_json::from_str(&content) - .map_err(|e| CredentialBridgeError::TokenRefreshFailed(format!("解析凭证失败: {e}")))?; - - creds["access_token"] - .as_str() - .map(String::from) - .ok_or_else(|| { - CredentialBridgeError::TokenRefreshFailed("凭证中缺少 access_token".to_string()) - }) - } - - /// 获取 Codex OAuth Token - async fn get_codex_token(&self, creds_path: &str) -> Result { - use crate::providers::codex::CodexProvider; - - let mut provider = CodexProvider::new(); - provider - .load_credentials_from_path(creds_path) - .await - .map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("加载 Codex 凭证失败: {e}")) - })?; - - provider.ensure_valid_token().await.map_err(|e| { - CredentialBridgeError::TokenRefreshFailed(format!("获取 Codex token 失败: {e}")) - }) - } - - /// 记录凭证使用 - pub fn record_usage(&self, db: &DbConnection, uuid: &str) -> Result<(), CredentialBridgeError> { - self.pool_service - .record_usage(db, uuid) - .map_err(CredentialBridgeError::DatabaseError) - } - - /// 标记凭证为健康 - pub fn mark_healthy( - &self, - db: &DbConnection, - uuid: &str, - model: Option<&str>, - ) -> Result<(), CredentialBridgeError> { - self.pool_service - .mark_healthy(db, uuid, model) - .map_err(CredentialBridgeError::DatabaseError) - } - - /// 标记凭证为不健康 - pub fn mark_unhealthy( - &self, - db: &DbConnection, - uuid: &str, - error: Option<&str>, - ) -> Result<(), CredentialBridgeError> { - self.pool_service - .mark_unhealthy(db, uuid, error) - .map_err(CredentialBridgeError::DatabaseError) - } -} - -/// 从 AsterProviderConfig 创建 Aster Provider -/// -/// 设置环境变量并调用 aster::providers::create -pub async fn create_aster_provider( - config: &AsterProviderConfig, -) -> Result, CredentialBridgeError> { - // 设置环境变量 - set_provider_env_vars(config); - - // 创建 ModelConfig - let model_config = ModelConfig::new(&config.model_name).map_err(|e| { - CredentialBridgeError::ProviderCreationFailed(format!("创建 ModelConfig 失败: {e}")) - })?; - - // 创建 Provider - aster::providers::create(&config.provider_name, model_config) - .await - .map_err(|e| { - CredentialBridgeError::ProviderCreationFailed(format!("创建 Provider 失败: {e}")) - }) -} - -/// 设置 Provider 环境变量 -/// 从 URL 中拆分 host(scheme+authority)和 path 部分 -/// -/// 例如: -/// - `https://api.openai.com` -> (`https://api.openai.com`, ``) -/// - `https://open.bigmodel.cn/api/paas/v4` -> (`https://open.bigmodel.cn`, `api/paas/v4`) -/// - `https://localhost:8080/v1` -> (`https://localhost:8080`, `v1`) -fn split_url_host_and_path(url: &str) -> (String, String) { - // 找到 scheme 之后的 authority 部分 - let after_scheme = if let Some(pos) = url.find("://") { - pos + 3 - } else { - return (url.to_string(), String::new()); - }; - - // 找到 authority 之后的第一个 /(即路径开始) - let path_start = url[after_scheme..].find('/').map(|p| p + after_scheme); - - match path_start { - Some(pos) => { - let host = url[..pos].to_string(); - let path = url[pos..].trim_matches('/').to_string(); - (host, path) - } - None => (url.to_string(), String::new()), - } -} - -fn set_provider_env_vars(config: &AsterProviderConfig) { - tracing::info!( - "[CredentialBridge] set_provider_env_vars: provider_name={}, has_api_key={}, base_url={:?}", - config.provider_name, - config.api_key.is_some(), - config.base_url - ); - - let env_key = match config.provider_name.as_str() { - "openai" => "OPENAI_API_KEY", - "anthropic" => "ANTHROPIC_API_KEY", - "google" => "GOOGLE_API_KEY", - "bedrock" => "AWS_ACCESS_KEY_ID", // Bedrock 使用 AWS 凭证 - "gcpvertexai" => "GOOGLE_API_KEY", - "codex" => "OPENAI_API_KEY", // Codex 兼容 OpenAI - "deepseek" | "custom_deepseek" => "OPENAI_API_KEY", // DeepSeek 使用 OpenAI 兼容 API - "groq" => "OPENAI_API_KEY", // Groq 使用 OpenAI 兼容 API - "mistral" => "OPENAI_API_KEY", // Mistral 使用 OpenAI 兼容 API - "openrouter" => "OPENROUTER_API_KEY", - _ => "OPENAI_API_KEY", // 默认使用 OpenAI 格式 - }; - - tracing::info!("[CredentialBridge] 设置环境变量: {}=***", env_key); - - if let Some(api_key) = &config.api_key { - std::env::set_var(env_key, api_key); - } - - // 设置 base_url - // Aster 的 OpenAI Provider 使用 OPENAI_HOST(仅 scheme+host+port)和 - // OPENAI_BASE_PATH(路径部分 + /chat/completions)环境变量 - if let Some(base_url) = &config.base_url { - match config.provider_name.as_str() { - "openai" => { - // 解析 base_url,将路径部分拆分到 OPENAI_BASE_PATH - // 例如 https://open.bigmodel.cn/api/paas/v4 - // -> OPENAI_HOST = https://open.bigmodel.cn - // -> OPENAI_BASE_PATH = api/paas/v4/chat/completions - let (host_part, path_part) = split_url_host_and_path(base_url); - if path_part.is_empty() { - // 无路径部分(如 https://api.openai.com),直接设置 - std::env::set_var("OPENAI_HOST", base_url); - // 清除可能残留的 OPENAI_BASE_PATH,使用 Aster 默认值 - std::env::remove_var("OPENAI_BASE_PATH"); - tracing::info!("[CredentialBridge] 设置 OPENAI_HOST={}", base_url); - } else { - // base_url 包含路径,需要拆分 - let base_path = format!("{}/chat/completions", path_part); - std::env::set_var("OPENAI_HOST", &host_part); - std::env::set_var("OPENAI_BASE_PATH", &base_path); - tracing::info!( - "[CredentialBridge] 设置 OPENAI_HOST={}, OPENAI_BASE_PATH={}", - host_part, - base_path - ); - } - } - "anthropic" => { - std::env::set_var("ANTHROPIC_BASE_URL", base_url); - } - _ => { - // 其他 Provider 使用通用格式 - let base_url_key = format!( - "{}_BASE_URL", - config.provider_name.to_uppercase().replace('-', "_") - ); - std::env::set_var(&base_url_key, base_url); - } - } - } -} - -/// Provider 类型映射 -/// -/// 将 ProxyCast PoolProviderType 映射到 Aster Provider 名称 -pub fn map_pool_type_to_aster(pool_type: &PoolProviderType) -> &'static str { - match pool_type { - PoolProviderType::Kiro => "bedrock", - PoolProviderType::Gemini => "google", - PoolProviderType::Antigravity => "google", - PoolProviderType::OpenAI => "openai", - PoolProviderType::Claude => "anthropic", - PoolProviderType::Anthropic => "anthropic", - PoolProviderType::AnthropicCompatible => "anthropic", - PoolProviderType::Vertex => "gcpvertexai", - PoolProviderType::GeminiApiKey => "google", - PoolProviderType::Codex => "codex", - PoolProviderType::ClaudeOAuth => "anthropic", - PoolProviderType::AzureOpenai => "azure", - PoolProviderType::AwsBedrock => "bedrock", - PoolProviderType::Ollama => "ollama", - } -} - -/// 将 provider_type 字符串映射到 Aster Provider 名称 -/// -/// 支持 60+ API Key Provider,包括 deepseek, moonshot, qwen 等 -fn map_provider_type_to_aster(provider_type: &str) -> &'static str { - match provider_type { - // 标准 Provider - "openai" => "openai", - "anthropic" | "claude" => "anthropic", - "google" | "gemini" => "google", - "bedrock" | "kiro" => "bedrock", - "gcpvertexai" | "vertex" => "gcpvertexai", - "codex" => "codex", - "azure" | "azure-openai" => "azure", - "ollama" => "ollama", - - // DeepSeek - 使用 openai 兼容 provider(Aster 会通过 alias 映射) - "deepseek" | "custom_deepseek" => "openai", - - // 其他 OpenAI 兼容 Provider - 使用 openai provider - // 这些 Provider 都使用 OpenAI 兼容 API,通过 base_url 区分 - "groq" => "openai", - "mistral" => "openai", - "openrouter" => "openrouter", - - // 默认使用 openai(OpenAI 兼容格式) - _ => "openai", - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_map_pool_type_to_aster() { - assert_eq!(map_pool_type_to_aster(&PoolProviderType::OpenAI), "openai"); - assert_eq!( - map_pool_type_to_aster(&PoolProviderType::Claude), - "anthropic" - ); - assert_eq!(map_pool_type_to_aster(&PoolProviderType::Gemini), "google"); - assert_eq!(map_pool_type_to_aster(&PoolProviderType::Kiro), "bedrock"); - } - - #[test] - fn test_credential_bridge_error_display() { - let err = CredentialBridgeError::NoCredentials("test".to_string()); - assert!(err.to_string().contains("没有可用凭证")); - } - - #[test] - fn test_split_url_host_and_path() { - // 无路径 - let (host, path) = split_url_host_and_path("https://api.openai.com"); - assert_eq!(host, "https://api.openai.com"); - assert_eq!(path, ""); - - // 带路径(智谱) - let (host, path) = split_url_host_and_path("https://open.bigmodel.cn/api/paas/v4"); - assert_eq!(host, "https://open.bigmodel.cn"); - assert_eq!(path, "api/paas/v4"); - - // 带端口 - let (host, path) = split_url_host_and_path("https://localhost:8080/v1"); - assert_eq!(host, "https://localhost:8080"); - assert_eq!(path, "v1"); - - // 尾部斜杠 - let (host, path) = split_url_host_and_path("https://api.deepseek.com/v1/"); - assert_eq!(host, "https://api.deepseek.com"); - assert_eq!(path, "v1"); - - // 仅根路径 - let (host, path) = split_url_host_and_path("https://api.openai.com/"); - assert_eq!(host, "https://api.openai.com"); - assert_eq!(path, ""); - } -} +pub use proxycast_agent::credential_bridge::*; diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index ab27e8be5..48e3d3aea 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -1,7 +1,7 @@ //! AI Agent 集成模块 //! //! 纯逻辑部分已迁移到 proxycast-agent crate, -//! 本模块保留深耦合部分(依赖 database, services, AppHandle)。 +//! 本模块保留深耦合部分(Aster 状态与 Tauri 桥接)。 pub mod aster_agent; pub mod aster_state; @@ -15,6 +15,7 @@ pub use proxycast_agent::prompt; // types 已迁移到 proxycast-core pub use proxycast_core::agent::types; +pub use proxycast_core::agent::types::*; pub use aster_agent::{AsterAgentWrapper, SessionDetail, SessionInfo}; pub use aster_state::AsterAgentState; @@ -25,4 +26,3 @@ pub use proxycast_agent::{convert_agent_event, convert_to_tauri_message, TauriAg pub use subagent_scheduler::{ ProxyCastScheduler, ProxyCastSubAgentExecutor, SubAgentProgressEvent, }; -pub use types::*; diff --git a/src-tauri/src/agent/subagent_scheduler.rs b/src-tauri/src/agent/subagent_scheduler.rs index 127c23f5f..1e76bca65 100644 --- a/src-tauri/src/agent/subagent_scheduler.rs +++ b/src-tauri/src/agent/subagent_scheduler.rs @@ -1,190 +1,26 @@ -//! SubAgent 调度器集成 +//! SubAgent 调度器集成(Tauri 桥接层) //! -//! 将 aster-rust 的 SubAgent 调度器与 ProxyCast 凭证池集成 -//! -//! ## 功能 -//! - 自动从凭证池选择健康凭证 -//! - 支持凭证 fallback 策略 -//! - 集成 Tauri 事件系统进行进度通知 +//! 纯逻辑已迁移到 `proxycast-agent` crate, +//! 本模块负责 Tauri 事件桥接。 -use std::collections::HashMap; use std::sync::Arc; -use std::time::Duration; use aster::agents::context::AgentContext; use aster::agents::subagent_scheduler::{ - SchedulerConfig, SchedulerError, SchedulerExecutionResult, SchedulerResult, SubAgentExecutor, - SubAgentResult, SubAgentScheduler, SubAgentTask, TokenUsage as SchedulerTokenUsage, + SchedulerConfig, SchedulerExecutionResult, SchedulerResult, SubAgentTask, }; -use aster::conversation::message::Message; -use chrono::Utc; use tauri::{AppHandle, Emitter}; -use tokio::sync::RwLock; -use tracing::{debug, info, warn}; -use crate::agent::credential_bridge::{ - create_aster_provider, AsterProviderConfig, CredentialBridge, -}; use crate::database::DbConnection; -/// ProxyCast SubAgent 执行器 -/// -/// 实现 aster-rust 的 SubAgentExecutor trait, -/// 集成 ProxyCast 凭证池进行 LLM 调用 -pub struct ProxyCastSubAgentExecutor { - /// 凭证桥接器 - credential_bridge: CredentialBridge, - /// 数据库连接 - db: DbConnection, - /// 默认模型 - default_model: String, - /// 默认 Provider 类型 - default_provider: String, - /// Tauri AppHandle(用于事件通知) - app_handle: Option, -} +pub use proxycast_agent::subagent_scheduler::{ + ProxyCastSubAgentExecutor, SchedulerEventEmitter, SubAgentProgressEvent, +}; -impl ProxyCastSubAgentExecutor { - /// 创建新的执行器 - pub fn new(db: DbConnection) -> Self { - Self { - credential_bridge: CredentialBridge::new(), - db, - default_model: "claude-sonnet-4-20250514".to_string(), - default_provider: "anthropic".to_string(), - app_handle: None, - } - } - - /// 设置 Tauri AppHandle - pub fn with_app_handle(mut self, handle: AppHandle) -> Self { - self.app_handle = Some(handle); - self - } - - /// 设置默认模型 - pub fn with_default_model(mut self, model: impl Into) -> Self { - self.default_model = model.into(); - self - } - - /// 设置默认 Provider - pub fn with_default_provider(mut self, provider: impl Into) -> Self { - self.default_provider = provider.into(); - self - } - - /// 从凭证池选择凭证 - async fn select_credential(&self, task: &SubAgentTask) -> SchedulerResult { - // 根据任务类型和模型选择 provider - let model = task.model.as_deref().unwrap_or(&self.default_model); - let provider_type = &self.default_provider; - - // 使用 CredentialBridge 选择凭证 - let config = self - .credential_bridge - .select_and_configure(&self.db, provider_type, model) - .await - .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; - - Ok(config) - } - - /// 发送 Tauri 事件 - #[allow(dead_code)] - fn emit_event(&self, event_name: &str, payload: impl serde::Serialize + Clone) { - if let Some(handle) = &self.app_handle { - if let Err(e) = handle.emit(event_name, payload) { - warn!("发送 Tauri 事件失败: {}", e); - } - } - } -} - -#[async_trait::async_trait] -impl SubAgentExecutor for ProxyCastSubAgentExecutor { - async fn execute_task( - &self, - task: &SubAgentTask, - context: &AgentContext, - ) -> SchedulerResult { - let start_time = Utc::now(); - info!("执行 SubAgent 任务: {}", task.id); - - // 选择凭证 - let provider_config = self.select_credential(task).await?; - debug!("使用凭证: {}", provider_config.credential_uuid); - - // 创建 provider - let provider = create_aster_provider(&provider_config) - .await - .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; - - // 构建提示 - let system_prompt = context.system_prompt.clone().unwrap_or_default(); - let user_message = Message::user().with_text(&task.prompt); - - // 调用 LLM(使用 complete 方法) - let (response_msg, usage) = provider - .complete(&system_prompt, &[user_message], &[]) - .await - .map_err(|e| SchedulerError::ProviderError(e.to_string()))?; - - let response = response_msg.as_concat_text(); - - let end_time = Utc::now(); - let duration = (end_time - start_time).to_std().unwrap_or(Duration::ZERO); - - // 生成摘要 - let summary = if task.return_summary { - Some(self.generate_summary(&response, task)) - } else { - None - }; - - // 转换 token 使用 - let token_usage = Some(SchedulerTokenUsage { - input_tokens: usage.usage.input_tokens.unwrap_or(0) as usize, - output_tokens: usage.usage.output_tokens.unwrap_or(0) as usize, - total_tokens: usage.usage.total_tokens.unwrap_or(0) as usize, - }); - - Ok(SubAgentResult { - task_id: task.id.clone(), - success: true, - output: Some(response), - summary, - error: None, - duration, - retries: 0, - started_at: start_time, - completed_at: end_time, - token_usage, - metadata: HashMap::new(), - }) - } -} - -impl ProxyCastSubAgentExecutor { - /// 生成摘要 - fn generate_summary(&self, output: &str, task: &SubAgentTask) -> String { - // 简单摘要:取前 500 字符 - let max_len = 500; - if output.chars().count() <= max_len { - format!("任务 {} 完成:\n{}", task.id, output) - } else { - let truncated: String = output.chars().take(max_len - 3).collect(); - format!("任务 {} 完成:\n{}...", task.id, truncated) - } - } -} - -/// ProxyCast SubAgent 调度器包装器 +/// ProxyCast SubAgent 调度器(Tauri 桥接) pub struct ProxyCastScheduler { - /// 内部调度器 - scheduler: Arc>>>, - /// 数据库连接 - db: DbConnection, + /// 内部纯逻辑调度器 + inner: proxycast_agent::subagent_scheduler::ProxyCastScheduler, /// Tauri AppHandle app_handle: Option, } @@ -193,8 +29,7 @@ impl ProxyCastScheduler { /// 创建新的调度器 pub fn new(db: DbConnection) -> Self { Self { - scheduler: Arc::new(RwLock::new(None)), - db, + inner: proxycast_agent::subagent_scheduler::ProxyCastScheduler::new(db), app_handle: None, } } @@ -207,26 +42,17 @@ impl ProxyCastScheduler { /// 初始化调度器 pub async fn init(&self, config: Option) { - let executor = ProxyCastSubAgentExecutor::new(self.db.clone()); - let executor = if let Some(handle) = &self.app_handle { - executor.with_app_handle(handle.clone()) - } else { - executor - }; - - let config = config.unwrap_or_default(); - - // 创建调度器并设置事件回调 - let app_handle = self.app_handle.clone(); - let scheduler = - SubAgentScheduler::new(config, executor).with_event_callback(move |event| { - if let Some(handle) = &app_handle { - let _ = handle.emit("subagent-scheduler-event", &event); + let event_emitter = self.app_handle.clone().map(|handle| { + Arc::new(move |event: &serde_json::Value| { + if let Err(err) = handle.emit("subagent-scheduler-event", event) { + tracing::warn!("发送 Tauri 事件失败: {}", err); } - }); + }) as SchedulerEventEmitter + }); - *self.scheduler.write().await = Some(scheduler); - info!("ProxyCast SubAgent 调度器初始化完成"); + self.inner + .init_with_event_emitter(config, event_emitter) + .await; } /// 执行任务 @@ -235,49 +61,11 @@ impl ProxyCastScheduler { tasks: Vec, parent_context: Option<&AgentContext>, ) -> SchedulerResult { - let scheduler = self.scheduler.read().await; - let scheduler = scheduler - .as_ref() - .ok_or_else(|| SchedulerError::ContextError("调度器未初始化".to_string()))?; - - scheduler.execute(tasks, parent_context).await + self.inner.execute(tasks, parent_context).await } /// 取消执行 pub async fn cancel(&self) { - if let Some(scheduler) = self.scheduler.read().await.as_ref() { - scheduler.cancel().await; - } - } -} - -/// Tauri 事件:SubAgent 进度 -#[derive(Debug, Clone, serde::Serialize)] -#[serde(rename_all = "camelCase")] -pub struct SubAgentProgressEvent { - /// 总任务数 - pub total: usize, - /// 已完成数 - pub completed: usize, - /// 失败数 - pub failed: usize, - /// 运行中数 - pub running: usize, - /// 进度百分比 - pub percentage: f64, - /// 当前任务 - pub current_tasks: Vec, -} - -impl From for SubAgentProgressEvent { - fn from(p: aster::agents::subagent_scheduler::SchedulerProgress) -> Self { - Self { - total: p.total, - completed: p.completed, - failed: p.failed, - running: p.running, - percentage: p.percentage, - current_tasks: p.current_tasks, - } + self.inner.cancel().await; } } diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index fb91968f5..4196b67dc 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -21,89 +21,27 @@ use crate::commands::skill_cmd::SkillServiceState; use crate::commands::terminal_cmd::TerminalManagerState; use crate::commands::tool_hooks::ToolHooksServiceState; use crate::commands::webview_cmd::{WebviewManagerState, WebviewManagerWrapper}; -use crate::config::{self, Config, ConfigManager, GlobalConfigManager, GlobalConfigManagerState}; +use crate::config::{GlobalConfigManager, GlobalConfigManagerState}; use crate::database::{self, DbConnection}; use crate::logger; use crate::mcp::McpManagerState; use crate::plugin; -use crate::server; -use crate::services::api_key_provider_service::ApiKeyProviderService; -use crate::services::aster_session_store::ProxyCastSessionStore; -use crate::services::context_memory_service::{ContextMemoryConfig, ContextMemoryService}; -use crate::services::provider_pool_service::ProviderPoolService; -use crate::services::skill_service::SkillService; -use crate::services::token_cache_service::TokenCacheService; -use crate::services::tool_hooks_service::ToolHooksService; -use crate::services::update_check_service::UpdateCheckServiceState; use crate::telemetry; use crate::voice::recording_service::{create_recording_service_state, RecordingServiceState}; +use proxycast_core::config::{Config, ConfigManager}; +use proxycast_server as server; +use proxycast_services::api_key_provider_service::ApiKeyProviderService; +use proxycast_services::aster_session_store::ProxyCastSessionStore; +use proxycast_services::context_memory_service::{ContextMemoryConfig, ContextMemoryService}; +use proxycast_services::provider_pool_service::ProviderPoolService; +use proxycast_services::skill_service::SkillService; +use proxycast_services::token_cache_service::TokenCacheService; +use proxycast_services::tool_hooks_service::ToolHooksService; +use proxycast_services::update_check_service::UpdateCheckServiceState; use super::types::{AppState, LogState, TokenCacheServiceState}; -use super::utils::{generate_api_key, is_valid_bind_host}; -/// 配置验证错误 -#[derive(Debug)] -pub enum ConfigError { - LoadFailed(String), - SaveFailed(String), - InvalidHost, - DefaultApiKeyWithNonLocalBind, - TlsNotSupported, - RemoteManagementNotSupported, -} - -impl std::fmt::Display for ConfigError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - ConfigError::LoadFailed(e) => write!(f, "配置加载失败: {e}"), - ConfigError::SaveFailed(e) => write!(f, "配置保存失败: {e}"), - ConfigError::InvalidHost => { - write!( - f, - "无效的监听地址。允许的地址:127.0.0.1、localhost、::1、0.0.0.0、::" - ) - } - ConfigError::DefaultApiKeyWithNonLocalBind => write!( - f, - "监听所有网络接口 (0.0.0.0 或 ::) 时,必须设置非默认的 API Key" - ), - ConfigError::TlsNotSupported => write!(f, "当前版本尚未支持 TLS"), - ConfigError::RemoteManagementNotSupported => { - write!(f, "远程管理需要 TLS 支持,当前版本未启用") - } - } - } -} - -/// 加载并验证配置 -pub fn load_and_validate_config() -> Result { - let mut config = config::load_config().map_err(|e| ConfigError::LoadFailed(e.to_string()))?; - - // 验证主机地址 - if !is_valid_bind_host(&config.server.host) { - return Err(ConfigError::InvalidHost); - } - - // 如果使用默认 API key,自动生成新密钥 - if config.server.api_key == config::DEFAULT_API_KEY { - let new_key = generate_api_key(); - config.server.api_key = new_key; - config::save_config(&config).map_err(|e| ConfigError::SaveFailed(e.to_string()))?; - tracing::info!("检测到默认 API key,已自动生成并保存新密钥"); - } - - // 检查 TLS 配置 - if config.server.tls.enable { - return Err(ConfigError::TlsNotSupported); - } - - // 检查远程管理配置 - if config.remote_management.allow_remote { - return Err(ConfigError::RemoteManagementNotSupported); - } - - Ok(config) -} +pub use proxycast_core::app_bootstrap::{load_and_validate_config, ConfigError}; /// 应用状态集合 pub struct AppStates { @@ -167,7 +105,7 @@ pub fn init_states(config: &Config) -> Result { let token_cache_service = TokenCacheService::new(); let token_cache_service_state = TokenCacheServiceState(Arc::new(token_cache_service)); - let machine_id_service = crate::services::machine_id_service::MachineIdService::new() + let machine_id_service = proxycast_services::machine_id_service::MachineIdService::new() .map_err(|e| format!("MachineIdService 初始化失败: {e}"))?; let machine_id_service_state: MachineIdState = Arc::new(RwLock::new(machine_id_service)); diff --git a/src-tauri/src/app/commands/config.rs b/src-tauri/src/app/commands/config.rs index 6a69f6978..5484a99cf 100644 --- a/src-tauri/src/app/commands/config.rs +++ b/src-tauri/src/app/commands/config.rs @@ -203,7 +203,7 @@ pub async fn update_provider_env_vars( api_host: String, api_key: Option, ) -> Result<(), String> { - use crate::services::live_sync::write_env_to_shell_config; + use proxycast_services::live_sync::write_env_to_shell_config; use serde_json::{json, Value}; use std::fs; diff --git a/src-tauri/src/app/commands/server.rs b/src-tauri/src/app/commands/server.rs index 95ea6f19b..ff3bd3856 100644 --- a/src-tauri/src/app/commands/server.rs +++ b/src-tauri/src/app/commands/server.rs @@ -7,7 +7,7 @@ use crate::app::TokenCacheServiceState; use crate::commands::provider_pool_cmd::ProviderPoolServiceState; use crate::commands::telemetry_cmd::TelemetryState; use crate::database; -use crate::server; +use proxycast_server as server; /// 启动服务器 #[tauri::command] diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index ad84a92af..2f1681ec2 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -318,7 +318,7 @@ pub fn run() { tauri::async_runtime::spawn(async move { // 创建 ModelRegistryService - let mut service = crate::services::model_registry_service::ModelRegistryService::new(db_clone); + let mut service = proxycast_services::model_registry_service::ModelRegistryService::new(db_clone); // 设置资源目录路径 service.set_resource_dir(resource_dir); @@ -344,7 +344,7 @@ pub fn run() { // 初始化终端会话管理器 { let app_handle = app.handle().clone(); - let terminal_manager = crate::terminal::TerminalSessionManager::new(crate::terminal::TauriEmitter(app_handle.clone())); + let terminal_manager = proxycast_terminal::TerminalSessionManager::new(crate::terminal::TauriEmitter(app_handle.clone())); if let Some(state) = app_handle.try_state::() { let mut guard = state.inner().0.blocking_write(); *guard = Some(terminal_manager); diff --git a/src-tauri/src/app/setup.rs b/src-tauri/src/app/setup.rs index cc32b7ed9..90595180e 100644 --- a/src-tauri/src/app/setup.rs +++ b/src-tauri/src/app/setup.rs @@ -8,11 +8,11 @@ use tauri::{App, Manager}; // use crate::agent::tools::{set_term_scrollback_tool_app_handle, set_terminal_tool_app_handle}; use crate::agent::AsterAgentState; use crate::database; -use crate::services::aster_session_store::ProxyCastSessionStore; -use crate::services::provider_pool_service::ProviderPoolService; -use crate::services::token_cache_service::TokenCacheService; use crate::telemetry; use crate::tray::{TrayIconStatus, TrayManager, TrayStateSnapshot}; +use proxycast_services::aster_session_store::ProxyCastSessionStore; +use proxycast_services::provider_pool_service::ProviderPoolService; +use proxycast_services::token_cache_service::TokenCacheService; use super::types::{AppState, LogState, TrayManagerState}; diff --git a/src-tauri/src/app/state.rs b/src-tauri/src/app/state.rs index 2104e2364..d0f530abd 100644 --- a/src-tauri/src/app/state.rs +++ b/src-tauri/src/app/state.rs @@ -15,20 +15,21 @@ use crate::commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoo use crate::commands::resilience_cmd::ResilienceConfigState; use crate::commands::skill_cmd::SkillServiceState; use crate::commands::tool_hooks::ToolHooksServiceState; -use crate::config::{Config, ConfigManager, GlobalConfigManager, GlobalConfigManagerState}; +use crate::config::{GlobalConfigManager, GlobalConfigManagerState}; use crate::database; use crate::plugin; -use crate::services::api_key_provider_service::ApiKeyProviderService; -use crate::services::context_memory_service::{ContextMemoryConfig, ContextMemoryService}; -use crate::services::provider_pool_service::ProviderPoolService; -use crate::services::skill_service::SkillService; -use crate::services::token_cache_service::TokenCacheService; -use crate::services::tool_hooks_service::ToolHooksService; use crate::telemetry; +use proxycast_core::config::{Config, ConfigManager}; +use proxycast_services::api_key_provider_service::ApiKeyProviderService; +use proxycast_services::context_memory_service::{ContextMemoryConfig, ContextMemoryService}; +use proxycast_services::provider_pool_service::ProviderPoolService; +use proxycast_services::skill_service::SkillService; +use proxycast_services::token_cache_service::TokenCacheService; +use proxycast_services::tool_hooks_service::ToolHooksService; use super::types::{AppState, LogState, TokenCacheServiceState}; use crate::logger; -use crate::server; +use proxycast_server as server; /// 初始化核心应用状态 pub fn init_core_state(config: Config) -> (AppState, LogState) { @@ -85,7 +86,7 @@ pub fn init_service_states() -> ServiceStates { let token_cache_service_state = TokenCacheServiceState(Arc::new(token_cache_service)); // Initialize MachineIdService - let machine_id_service = crate::services::machine_id_service::MachineIdService::new() + let machine_id_service = proxycast_services::machine_id_service::MachineIdService::new() .expect("Failed to initialize MachineIdService"); let machine_id_service_state: MachineIdState = Arc::new(RwLock::new(machine_id_service)); diff --git a/src-tauri/src/app/types.rs b/src-tauri/src/app/types.rs index fcd63b649..5ed6624e1 100644 --- a/src-tauri/src/app/types.rs +++ b/src-tauri/src/app/types.rs @@ -7,9 +7,9 @@ use tauri::{Emitter, Runtime}; use tokio::sync::RwLock; use crate::logger; -use crate::server; -use crate::services::token_cache_service::TokenCacheService; use crate::tray::TrayManager; +use proxycast_server as server; +use proxycast_services::token_cache_service::TokenCacheService; use proxycast_core::event_emit::EventEmit; diff --git a/src-tauri/src/app/utils.rs b/src-tauri/src/app/utils.rs index 931e8903b..97c215454 100644 --- a/src-tauri/src/app/utils.rs +++ b/src-tauri/src/app/utils.rs @@ -1,130 +1,8 @@ -//! 辅助函数 +//! 辅助函数(桥接层) //! -//! 包含通用工具函数。 +//! 纯逻辑已迁移到 `proxycast-core` crate, +//! 本模块保留兼容导出。 -use crate::config; - -/// 生成安全的 API Key -pub fn generate_api_key() -> String { - config::generate_secure_api_key() -} - -/// 检查是否为回环地址 -pub fn is_loopback_host(host: &str) -> bool { - if host == "localhost" { - return true; - } - match host.parse::() { - Ok(addr) => addr.is_loopback(), - Err(_) => false, - } -} - -/// 检查是否为有效的绑定地址 -/// 允许回环地址、0.0.0.0 和私有网络地址 -pub fn is_valid_bind_host(host: &str) -> bool { - if is_loopback_host(host) { - return true; - } - // 允许 0.0.0.0 和 :: (监听所有接口) - if host == "0.0.0.0" || host == "::" { - return true; - } - - // 允许私有网络地址 - if let Ok(addr) = host.parse::() { - if let std::net::IpAddr::V4(ipv4) = addr { - // 10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16 - let octets = ipv4.octets(); - return octets[0] == 10 - || (octets[0] == 172 && (octets[1] >= 16 && octets[1] <= 31)) - || (octets[0] == 192 && octets[1] == 168); - } - } - - false -} - -/// 检查是否为非本地绑定地址(需要强 API Key) -pub fn is_non_local_bind(host: &str) -> bool { - if host == "0.0.0.0" || host == "::" { - return true; - } - - // 私有网络地址也算非本地绑定 - if let Ok(addr) = host.parse::() { - if let std::net::IpAddr::V4(ipv4) = addr { - let octets = ipv4.octets(); - return octets[0] == 10 - || (octets[0] == 172 && (octets[1] >= 16 && octets[1] <= 31)) - || (octets[0] == 192 && octets[1] == 168); - } - } - - false -} - -/// 掩码敏感 Token -pub fn mask_token(token: &str) -> String { - let chars: Vec = token.chars().collect(); - if chars.len() <= 12 { - "****".to_string() - } else { - let prefix: String = chars[..6].iter().collect(); - let suffix: String = chars[chars.len() - 4..].iter().collect(); - format!("{prefix}****{suffix}") - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_is_loopback_host() { - assert!(is_loopback_host("localhost")); - assert!(is_loopback_host("127.0.0.1")); - assert!(is_loopback_host("::1")); - assert!(!is_loopback_host("0.0.0.0")); - assert!(!is_loopback_host("192.168.1.1")); - } - - #[test] - fn test_is_valid_bind_host() { - // 回环地址 - assert!(is_valid_bind_host("localhost")); - assert!(is_valid_bind_host("127.0.0.1")); - assert!(is_valid_bind_host("::1")); - // 监听所有接口 - assert!(is_valid_bind_host("0.0.0.0")); - assert!(is_valid_bind_host("::")); - // 私有网络地址(局域网)- 应该允许 - assert!(is_valid_bind_host("192.168.1.1")); - assert!(is_valid_bind_host("10.0.0.1")); - assert!(is_valid_bind_host("172.16.0.1")); - assert!(is_valid_bind_host("172.31.255.255")); - // 公网地址不允许 - assert!(!is_valid_bind_host("8.8.8.8")); - assert!(!is_valid_bind_host("1.1.1.1")); - } - - #[test] - fn test_is_non_local_bind() { - // 监听所有接口 - assert!(is_non_local_bind("0.0.0.0")); - assert!(is_non_local_bind("::")); - // 回环地址不是非本地绑定 - assert!(!is_non_local_bind("127.0.0.1")); - assert!(!is_non_local_bind("localhost")); - // 私有网络地址是非本地绑定(需要强 API Key) - assert!(is_non_local_bind("192.168.1.1")); - assert!(is_non_local_bind("10.0.0.1")); - assert!(is_non_local_bind("172.16.0.1")); - } - - #[test] - fn test_mask_token() { - assert_eq!(mask_token("short"), "****"); - assert_eq!(mask_token("abcdefghijklmnop"), "abcdef****mnop"); - } -} +pub use proxycast_core::app_utils::{ + generate_api_key, is_loopback_host, is_non_local_bind, is_valid_bind_host, mask_token, +}; diff --git a/src-tauri/src/commands/api_key_provider_cmd.rs b/src-tauri/src/commands/api_key_provider_cmd.rs index 287a78e68..bc14e4fb1 100644 --- a/src-tauri/src/commands/api_key_provider_cmd.rs +++ b/src-tauri/src/commands/api_key_provider_cmd.rs @@ -10,7 +10,7 @@ use crate::database::dao::api_key_provider::{ }; use crate::database::system_providers::get_system_providers; use crate::database::DbConnection; -use crate::services::api_key_provider_service::{ +use proxycast_services::api_key_provider_service::{ ApiKeyProviderService, ChatTestResult, ConnectionTestResult, ImportResult, }; use serde::{Deserialize, Serialize}; diff --git a/src-tauri/src/commands/aster_agent_cmd.rs b/src-tauri/src/commands/aster_agent_cmd.rs index bdcf57c95..bc6ec70dd 100644 --- a/src-tauri/src/commands/aster_agent_cmd.rs +++ b/src-tauri/src/commands/aster_agent_cmd.rs @@ -5,17 +5,17 @@ //! 支持从 ProxyCast 凭证池自动选择凭证 use crate::agent::aster_state::{ProviderConfig, SessionConfigBuilder}; -use crate::agent::event_converter::convert_agent_event; use crate::agent::{ AsterAgentState, AsterAgentWrapper, SessionDetail, SessionInfo, TauriAgentEvent, }; use crate::database::dao::agent::AgentDao; use crate::database::DbConnection; use crate::mcp::{McpManagerState, McpServerConfig}; -use crate::services::mcp_service::McpService; use aster::agents::extension::{Envs, ExtensionConfig}; use aster::conversation::message::Message; use futures::StreamExt; +use proxycast_agent::event_converter::convert_agent_event; +use proxycast_services::mcp_service::McpService; use serde::{Deserialize, Serialize}; use tauri::{AppHandle, Emitter, State}; @@ -651,7 +651,7 @@ async fn ensure_proxycast_mcp_servers_running( return (0, 0); } - let candidates: Vec<&crate::models::McpServer> = + let candidates: Vec<&crate::models::mcp_model::McpServer> = servers.iter().filter(|s| s.enabled_proxycast).collect(); if candidates.is_empty() { diff --git a/src-tauri/src/commands/config_cmd.rs b/src-tauri/src/commands/config_cmd.rs index 90beb5165..44eb775dd 100644 --- a/src-tauri/src/commands/config_cmd.rs +++ b/src-tauri/src/commands/config_cmd.rs @@ -2,7 +2,7 @@ use crate::config::{ Config, ConfigManager, ExportBundle, ExportOptions as ExportServiceOptions, ExportService, ImportOptions as ImportServiceOptions, ImportService, ValidationResult, }; -use crate::models::AppType; +use crate::models::app_type::AppType; use serde::{Deserialize, Serialize}; use std::path::PathBuf; use tauri::{AppHandle, Manager}; diff --git a/src-tauri/src/commands/connection_cmd.rs b/src-tauri/src/commands/connection_cmd.rs index 51d6aa78f..937fd73d2 100644 --- a/src-tauri/src/commands/connection_cmd.rs +++ b/src-tauri/src/commands/connection_cmd.rs @@ -11,7 +11,7 @@ //! - `connection_get_raw_config` - 获取原始配置内容 //! - `connection_save_raw_config` - 保存原始配置内容 -use crate::terminal::connections::{ +use proxycast_terminal::connections::{ ConnectionConfig, ConnectionConfigManager, ConnectionConfigType, ConnectionListEntry, }; use serde::{Deserialize, Serialize}; diff --git a/src-tauri/src/commands/context_memory.rs b/src-tauri/src/commands/context_memory.rs index e985a20a6..b6b0ada0c 100644 --- a/src-tauri/src/commands/context_memory.rs +++ b/src-tauri/src/commands/context_memory.rs @@ -1,6 +1,6 @@ //! 上下文记忆管理相关的 Tauri 命令 -use crate::services::context_memory_service::{ +use proxycast_services::context_memory_service::{ ContextMemoryService, MemoryEntry, MemoryFileType, MemoryStats, }; use serde::{Deserialize, Serialize}; diff --git a/src-tauri/src/commands/general_chat_cmd.rs b/src-tauri/src/commands/general_chat_cmd.rs index 46eca18c9..b64f99372 100644 --- a/src-tauri/src/commands/general_chat_cmd.rs +++ b/src-tauri/src/commands/general_chat_cmd.rs @@ -14,7 +14,7 @@ use crate::database::dao::general_chat::GeneralChatDao; use crate::database::DbConnection; -use crate::services::general_chat::{ +use proxycast_services::general_chat::{ ChatMessage, ChatSession, ContentBlock, MessageRole, SessionDetail, }; use serde::{Deserialize, Serialize}; diff --git a/src-tauri/src/commands/kiro_local.rs b/src-tauri/src/commands/kiro_local.rs index 57a939b23..bbfc0ed8c 100644 --- a/src-tauri/src/commands/kiro_local.rs +++ b/src-tauri/src/commands/kiro_local.rs @@ -6,7 +6,7 @@ use crate::commands::provider_pool_cmd::ProviderPoolServiceState; use crate::database::DbConnection; use crate::models::kiro_fingerprint::{KiroFingerprintStore, SwitchToLocalResult}; use crate::models::provider_pool_model::CredentialData; -use crate::services::machine_id_service::MachineIdService; +use proxycast_services::machine_id_service::MachineIdService; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; use std::fs; diff --git a/src-tauri/src/commands/machine_id_cmd.rs b/src-tauri/src/commands/machine_id_cmd.rs index 5c2c38210..cbc3d9e74 100644 --- a/src-tauri/src/commands/machine_id_cmd.rs +++ b/src-tauri/src/commands/machine_id_cmd.rs @@ -1,5 +1,5 @@ use crate::models::machine_id::*; -use crate::services::machine_id_service::MachineIdService; +use proxycast_services::machine_id_service::MachineIdService; use std::sync::Arc; use tauri::State; use tokio::sync::RwLock; diff --git a/src-tauri/src/commands/material_cmd.rs b/src-tauri/src/commands/material_cmd.rs index f0bf135d6..12cacdbf8 100644 --- a/src-tauri/src/commands/material_cmd.rs +++ b/src-tauri/src/commands/material_cmd.rs @@ -18,7 +18,7 @@ use crate::database::DbConnection; use crate::models::project_model::{ Material, MaterialFilter, MaterialUpdate, UploadMaterialRequest, }; -use crate::services::material_service::MaterialService; +use proxycast_services::material_service::MaterialService; // ============================================================================ // Tauri 命令 diff --git a/src-tauri/src/commands/mcp_cmd.rs b/src-tauri/src/commands/mcp_cmd.rs index e60a82c04..7d901cc1a 100644 --- a/src-tauri/src/commands/mcp_cmd.rs +++ b/src-tauri/src/commands/mcp_cmd.rs @@ -39,8 +39,8 @@ use crate::mcp::{ McpManagerState, McpPromptDefinition, McpPromptResult, McpResourceContent, McpResourceDefinition, McpServerConfig, McpServerInfo, McpToolDefinition, McpToolResult, }; -use crate::models::McpServer; -use crate::services::mcp_service::McpService; +use crate::models::mcp_model::McpServer; +use proxycast_services::mcp_service::McpService; use tauri::State; use tracing::{debug, error, info}; diff --git a/src-tauri/src/commands/model_cmd.rs b/src-tauri/src/commands/model_cmd.rs index 031ef54c2..a68d51d3a 100644 --- a/src-tauri/src/commands/model_cmd.rs +++ b/src-tauri/src/commands/model_cmd.rs @@ -2,7 +2,7 @@ use crate::database::dao::provider_pool::ProviderPoolDao; use crate::database::DbConnection; -use crate::services::model_service::ModelService; +use proxycast_services::model_service::ModelService; use std::collections::HashMap; use tauri::State; diff --git a/src-tauri/src/commands/model_registry_cmd.rs b/src-tauri/src/commands/model_registry_cmd.rs index b110b52e8..df02364df 100644 --- a/src-tauri/src/commands/model_registry_cmd.rs +++ b/src-tauri/src/commands/model_registry_cmd.rs @@ -5,7 +5,7 @@ use crate::models::model_registry::{ EnhancedModelMetadata, ModelSyncState, ModelTier, ProviderAliasConfig, UserModelPreference, }; -use crate::services::model_registry_service::{FetchModelsResult, ModelRegistryService}; +use proxycast_services::model_registry_service::{FetchModelsResult, ModelRegistryService}; use std::collections::BTreeSet; use std::sync::Arc; use tauri::State; diff --git a/src-tauri/src/commands/network_cmd.rs b/src-tauri/src/commands/network_cmd.rs index 7f56ee1c2..f2fc4e3aa 100644 --- a/src-tauri/src/commands/network_cmd.rs +++ b/src-tauri/src/commands/network_cmd.rs @@ -3,9 +3,7 @@ //! 核心逻辑已迁移到 proxycast-core::network,本文件保留 Tauri 命令包装。 // 重新导出核心类型 -pub use proxycast_core::network::{ - get_accessible_host, get_accessible_url, get_local_url, NetworkInfo, -}; +pub use proxycast_core::network::{get_accessible_url, NetworkInfo}; /// 获取本地网络信息(Tauri 命令包装) #[tauri::command] diff --git a/src-tauri/src/commands/orchestrator_cmd.rs b/src-tauri/src/commands/orchestrator_cmd.rs index 47bc12985..f9abcd640 100644 --- a/src-tauri/src/commands/orchestrator_cmd.rs +++ b/src-tauri/src/commands/orchestrator_cmd.rs @@ -4,7 +4,7 @@ use crate::database::dao::provider_pool::ProviderPoolDao; use crate::database::DbConnection; -use crate::orchestrator::{ +use proxycast_core::orchestrator::{ get_global_orchestrator, init_global_orchestrator, AvailableModel, CredentialInfo, OrchestratorConfig, PoolStats, ProviderType, SelectionContext, SelectionResult, ServiceTier, StrategyInfo, TaskHint, diff --git a/src-tauri/src/commands/persona_cmd.rs b/src-tauri/src/commands/persona_cmd.rs index 89abbf807..1b688c5c5 100644 --- a/src-tauri/src/commands/persona_cmd.rs +++ b/src-tauri/src/commands/persona_cmd.rs @@ -24,7 +24,7 @@ use crate::models::project_model::{ BrandPersona, BrandPersonaExtension, BrandPersonaTemplate, CreateBrandExtensionRequest, CreatePersonaRequest, Persona, PersonaTemplate, PersonaUpdate, UpdateBrandExtensionRequest, }; -use crate::services::persona_service::PersonaService; +use proxycast_services::persona_service::PersonaService; // ============================================================================ // Tauri 命令 diff --git a/src-tauri/src/commands/plugin_cmd.rs b/src-tauri/src/commands/plugin_cmd.rs index a10f59945..371573b5f 100644 --- a/src-tauri/src/commands/plugin_cmd.rs +++ b/src-tauri/src/commands/plugin_cmd.rs @@ -11,7 +11,7 @@ #![allow(dead_code)] -use crate::plugin::{PluginConfig, PluginInfo, PluginManager, PluginManifest, PluginType}; +use proxycast_core::plugin::{PluginConfig, PluginInfo, PluginManager, PluginManifest, PluginType}; use serde::{Deserialize, Serialize}; use std::path::Path; use std::sync::Arc; @@ -360,7 +360,7 @@ pub async fn get_plugins_with_ui( // 插件 UI 相关命令 // ============================================================================ -use crate::plugin::{UIMessage, UserAction}; +use proxycast_core::plugin::{UIMessage, UserAction}; /// 获取插件 UI 定义 /// 返回插件的初始 UI 消息列表 diff --git a/src-tauri/src/commands/plugin_install_cmd.rs b/src-tauri/src/commands/plugin_install_cmd.rs index c99ad4097..76f6e2a1c 100644 --- a/src-tauri/src/commands/plugin_install_cmd.rs +++ b/src-tauri/src/commands/plugin_install_cmd.rs @@ -8,7 +8,7 @@ //! //! _需求: 1.1, 2.1, 2.2, 2.4, 3.1, 3.2, 3.3, 4.2, 6.1_ -use crate::plugin::installer::{ +use proxycast_core::plugin::installer::{ InstallProgress, InstalledPlugin, PluginInstaller, ProgressCallback, }; use serde::{Deserialize, Serialize}; @@ -44,7 +44,7 @@ impl TauriProgressCallback { } } -impl crate::plugin::installer::ProgressCallback for TauriProgressCallback { +impl proxycast_core::plugin::installer::ProgressCallback for TauriProgressCallback { fn on_progress(&self, progress: InstallProgress) { // 发送进度事件到前端 let _ = self.app_handle.emit(INSTALL_PROGRESS_EVENT, &progress); diff --git a/src-tauri/src/commands/prompt_cmd.rs b/src-tauri/src/commands/prompt_cmd.rs index 83eaf8a27..8c8eaf29a 100644 --- a/src-tauri/src/commands/prompt_cmd.rs +++ b/src-tauri/src/commands/prompt_cmd.rs @@ -1,6 +1,6 @@ use crate::database::DbConnection; -use crate::models::Prompt; -use crate::services::prompt_service::PromptService; +use crate::models::prompt_model::Prompt; +use proxycast_services::prompt_service::PromptService; use std::collections::HashMap; use tauri::State; diff --git a/src-tauri/src/commands/provider_pool_cmd.rs b/src-tauri/src/commands/provider_pool_cmd.rs index 9c9e8eebe..bf9d46b86 100644 --- a/src-tauri/src/commands/provider_pool_cmd.rs +++ b/src-tauri/src/commands/provider_pool_cmd.rs @@ -2,15 +2,15 @@ #![allow(dead_code)] -use crate::credential::CredentialSyncService; use crate::database::dao::provider_pool::ProviderPoolDao; use crate::database::DbConnection; use crate::models::provider_pool_model::{ AddCredentialRequest, CredentialData, CredentialDisplay, HealthCheckResult, OAuthStatus, PoolProviderType, ProviderCredential, ProviderPoolOverview, UpdateCredentialRequest, }; -use crate::services::provider_pool_service::ProviderPoolService; use chrono::Utc; +use proxycast_credential::CredentialSyncService; +use proxycast_services::provider_pool_service::ProviderPoolService; use std::fs; use std::path::{Path, PathBuf}; use std::sync::Arc; @@ -3852,7 +3852,7 @@ pub async fn get_credential_health( db: State<'_, DbConnection>, pool_service: State<'_, ProviderPoolServiceState>, uuid: String, -) -> Result, String> { +) -> Result, String> { pool_service.0.get_credential_health(&db, &uuid) } @@ -3862,6 +3862,6 @@ pub async fn get_credential_health( pub async fn get_all_credential_health( db: State<'_, DbConnection>, pool_service: State<'_, ProviderPoolServiceState>, -) -> Result, String> { +) -> Result, String> { pool_service.0.get_all_credential_health(&db) } diff --git a/src-tauri/src/commands/screenshot_cmd.rs b/src-tauri/src/commands/screenshot_cmd.rs index cf3df58ee..80b01d295 100644 --- a/src-tauri/src/commands/screenshot_cmd.rs +++ b/src-tauri/src/commands/screenshot_cmd.rs @@ -10,10 +10,7 @@ use crate::config::{ExperimentalFeatures, GlobalConfigManagerState}; use crate::screenshot::{capture, shortcut}; -use base64::{engine::general_purpose::STANDARD, Engine}; -use std::path::Path; use tauri::{AppHandle, Emitter, Manager, State}; -use tokio::fs; use tracing::{debug, error, info}; /// 获取实验室功能配置 @@ -112,18 +109,15 @@ pub async fn save_experimental_config( /// /// 启动交互式截图,返回截图文件路径 /// -/// # 参数 -/// - `app`: Tauri 应用句柄 -/// /// # 返回 /// 成功返回截图文件路径,用户取消返回空字符串,失败返回错误信息 /// /// _需求: 3.1_ #[tauri::command] -pub async fn start_screenshot(app: AppHandle) -> Result { +pub async fn start_screenshot() -> Result { info!("启动截图命令"); - match capture::start_capture(&app).await { + match capture::start_capture().await { Ok(path) => { info!("截图成功: {:?}", path); Ok(path.to_string_lossy().to_string()) @@ -211,9 +205,6 @@ pub async fn update_screenshot_shortcut( /// /// 关闭当前打开的截图对话悬浮窗口 /// -/// # 参数 -/// - `app`: Tauri 应用句柄 -/// /// # 返回 /// 成功返回 Ok(()), 失败返回错误信息 #[tauri::command] @@ -254,30 +245,7 @@ pub fn open_input_with_text(app: AppHandle, text: String) -> Result<(), String> /// _需求: 5.1_ #[tauri::command] pub async fn read_image_as_base64(path: String) -> Result { - debug!("读取图片为 Base64: {}", path); - - let path = Path::new(&path); - - // 检查文件是否存在 - if !path.exists() { - return Err(format!("文件不存在: {}", path.display())); - } - - // 读取文件内容 - let bytes = fs::read(path) - .await - .map_err(|e| format!("读取文件失败: {e}"))?; - - // 检查文件是否为空 - if bytes.is_empty() { - return Err("文件为空".to_string()); - } - - // 编码为 Base64 - let base64 = STANDARD.encode(&bytes); - - debug!("图片读取成功,大小: {} 字节", bytes.len()); - Ok(base64) + proxycast_services::screenshot_image_service::read_image_as_base64(&path).await } /// 截图对话消息结构 diff --git a/src-tauri/src/commands/skill_cmd.rs b/src-tauri/src/commands/skill_cmd.rs index ac9b8074a..6f11c91df 100644 --- a/src-tauri/src/commands/skill_cmd.rs +++ b/src-tauri/src/commands/skill_cmd.rs @@ -1,9 +1,10 @@ use crate::agent::aster_state::AsterAgentState; use crate::database::dao::skills::SkillDao; use crate::database::DbConnection; -use crate::models::{AppType, Skill, SkillRepo, SkillState}; -use crate::services::skill_service::SkillService; +use crate::models::app_type::AppType; +use crate::models::skill_model::{Skill, SkillRepo, SkillState}; use chrono::Utc; +use proxycast_services::skill_service::SkillService; use std::path::Path; use std::sync::Arc; use tauri::State; diff --git a/src-tauri/src/commands/skill_exec_cmd.rs b/src-tauri/src/commands/skill_exec_cmd.rs index c6aeda244..35a20e4eb 100644 --- a/src-tauri/src/commands/skill_exec_cmd.rs +++ b/src-tauri/src/commands/skill_exec_cmd.rs @@ -23,15 +23,15 @@ use uuid::Uuid; use aster::conversation::message::Message; use crate::agent::aster_state::SessionConfigBuilder; -use crate::agent::event_converter::convert_agent_event; use crate::agent::{AsterAgentState, TauriAgentEvent}; use crate::database::DbConnection; -use crate::skills::{ +use crate::skills::TauriExecutionCallback; +use proxycast_agent::event_converter::convert_agent_event; +use proxycast_skills::{ find_skill_by_name, get_proxycast_skills_dir, load_skills_from_directory, ExecutionCallback, - TauriExecutionCallback, }; #[cfg(test)] -use crate::skills::{ +use proxycast_skills::{ load_skill_from_file, parse_allowed_tools, parse_boolean, parse_skill_frontmatter, }; diff --git a/src-tauri/src/commands/switch_cmd.rs b/src-tauri/src/commands/switch_cmd.rs index 7dffc8a14..383eea06a 100644 --- a/src-tauri/src/commands/switch_cmd.rs +++ b/src-tauri/src/commands/switch_cmd.rs @@ -1,7 +1,8 @@ use crate::database::DbConnection; -use crate::models::{AppType, Provider}; -use crate::services::live_sync::{check_config_sync, sync_from_external, SyncCheckResult}; -use crate::services::switch::SwitchService; +use crate::models::app_type::AppType; +use crate::models::provider_model::Provider; +use proxycast_services::live_sync::{check_config_sync, sync_from_external, SyncCheckResult}; +use proxycast_services::switch::SwitchService; use serde_json::Value; use tauri::State; diff --git a/src-tauri/src/commands/template_cmd.rs b/src-tauri/src/commands/template_cmd.rs index 781cc845e..2e1430695 100644 --- a/src-tauri/src/commands/template_cmd.rs +++ b/src-tauri/src/commands/template_cmd.rs @@ -15,7 +15,7 @@ use tauri::State; use crate::database::DbConnection; use crate::models::project_model::{CreateTemplateRequest, Template, TemplateUpdate}; -use crate::services::template_service::TemplateService; +use proxycast_services::template_service::TemplateService; // ============================================================================ // Tauri 命令 diff --git a/src-tauri/src/commands/terminal_cmd.rs b/src-tauri/src/commands/terminal_cmd.rs index 1fda8932d..0bfb79213 100644 --- a/src-tauri/src/commands/terminal_cmd.rs +++ b/src-tauri/src/commands/terminal_cmd.rs @@ -18,7 +18,7 @@ use serde::Serialize; use tauri::State; use tokio::sync::RwLock; -use crate::terminal::{SessionMetadata, TerminalSessionManager}; +use proxycast_terminal::{SessionMetadata, TerminalSessionManager}; /// 终端会话管理器状态包装 pub struct TerminalManagerState(pub Arc>>); diff --git a/src-tauri/src/commands/tool_hooks.rs b/src-tauri/src/commands/tool_hooks.rs index 44619043f..951e99b09 100644 --- a/src-tauri/src/commands/tool_hooks.rs +++ b/src-tauri/src/commands/tool_hooks.rs @@ -1,6 +1,6 @@ //! 工具钩子管理相关的 Tauri 命令 -use crate::services::tool_hooks_service::{ +use proxycast_services::tool_hooks_service::{ HookContext, HookExecutionStats, HookRule, HookTrigger, ToolHooksService, }; use serde::{Deserialize, Serialize}; diff --git a/src-tauri/src/commands/unified_chat_cmd.rs b/src-tauri/src/commands/unified_chat_cmd.rs index 8e0270f84..1d50442ef 100644 --- a/src-tauri/src/commands/unified_chat_cmd.rs +++ b/src-tauri/src/commands/unified_chat_cmd.rs @@ -14,12 +14,12 @@ //! - `docs/prd/chat-architecture-redesign.md` use crate::agent::aster_state::SessionConfigBuilder; -use crate::agent::event_converter::convert_agent_event; use crate::agent::{AsterAgentState, TauriAgentEvent}; use crate::database::dao::chat::{ChatDao, ChatMessage, ChatMode, ChatSession}; use crate::database::DbConnection; use aster::conversation::message::Message; use futures::StreamExt; +use proxycast_agent::event_converter::convert_agent_event; use serde::{Deserialize, Serialize}; use tauri::{AppHandle, Emitter, State}; diff --git a/src-tauri/src/commands/update_cmd.rs b/src-tauri/src/commands/update_cmd.rs index e6307b725..2807301bd 100644 --- a/src-tauri/src/commands/update_cmd.rs +++ b/src-tauri/src/commands/update_cmd.rs @@ -8,10 +8,10 @@ use crate::app::AppState; use crate::config; -use crate::services::update_check_service::{ +use crate::services::update_window; +use proxycast_services::update_check_service::{ UpdateCheckService, UpdateCheckServiceState, UpdateInfo, }; -use crate::services::update_window; use serde::{Deserialize, Serialize}; use std::sync::Arc; use tauri::{AppHandle, Manager, State}; diff --git a/src-tauri/src/commands/usage_cmd.rs b/src-tauri/src/commands/usage_cmd.rs index f448fb081..54f5a67bf 100644 --- a/src-tauri/src/commands/usage_cmd.rs +++ b/src-tauri/src/commands/usage_cmd.rs @@ -5,8 +5,8 @@ use crate::database::dao::provider_pool::ProviderPoolDao; use crate::database::DbConnection; use crate::models::provider_pool_model::{CredentialData, PoolProviderType}; -use crate::services::usage_service::{self, UsageInfo}; use crate::TokenCacheServiceState; +use proxycast_services::usage_service::{self, UsageInfo}; use tauri::State; /// 默认 Kiro 版本号 diff --git a/src-tauri/src/commands/websocket_cmd.rs b/src-tauri/src/commands/websocket_cmd.rs index 3705ac37b..2fd5d1a0e 100644 --- a/src-tauri/src/commands/websocket_cmd.rs +++ b/src-tauri/src/commands/websocket_cmd.rs @@ -1,6 +1,6 @@ //! WebSocket 相关的 Tauri 命令 -use crate::websocket::{WsConnection, WsStatsSnapshot}; +use proxycast_websocket::{WsConnection, WsStatsSnapshot}; use serde::{Deserialize, Serialize}; use std::sync::Arc; use tokio::sync::RwLock; diff --git a/src-tauri/src/commands/workspace_cmd.rs b/src-tauri/src/commands/workspace_cmd.rs index 3eba72e53..296105088 100644 --- a/src-tauri/src/commands/workspace_cmd.rs +++ b/src-tauri/src/commands/workspace_cmd.rs @@ -13,10 +13,10 @@ use crate::database::DbConnection; use crate::models::project_model::ProjectContext; -use crate::services::project_context_builder::ProjectContextBuilder; use crate::workspace::{ Workspace, WorkspaceManager, WorkspaceSettings, WorkspaceType, WorkspaceUpdate, }; +use proxycast_services::project_context_builder::ProjectContextBuilder; use serde::{Deserialize, Serialize}; use std::path::PathBuf; use std::sync::Arc; diff --git a/src-tauri/src/config/mod.rs b/src-tauri/src/config/mod.rs index 55a7085bb..2b0075108 100644 --- a/src-tauri/src/config/mod.rs +++ b/src-tauri/src/config/mod.rs @@ -3,19 +3,16 @@ //! 核心配置类型、YAML 支持、热重载和导入导出功能已迁移到 proxycast-core crate。 //! 本模块保留 observer(依赖 Tauri)和集成测试。 -#![allow(unused_imports)] - -// 从 core crate 重新导出所有配置类型 -pub use proxycast_core::config::*; - // observer 模块保留在主 crate(依赖 Tauri) pub mod observer; +// 兼容导出:配置核心能力已迁移到 proxycast-core crate +pub use proxycast_core::config::*; + // 重新导出观察者模块的核心类型 -pub use observer::{ - ConfigChangeEvent, ConfigChangeSource, ConfigObserver, ConfigSubject, GlobalConfigManager, - GlobalConfigManagerState, -}; +pub use observer::ConfigChangeSource; +pub use proxycast_config::observer::manager::GlobalConfigManager; +pub use proxycast_config::GlobalConfigManagerState; #[cfg(test)] mod tests; diff --git a/src-tauri/src/config/observer/mod.rs b/src-tauri/src/config/observer/mod.rs index 63ba04412..dbb410330 100644 --- a/src-tauri/src/config/observer/mod.rs +++ b/src-tauri/src/config/observer/mod.rs @@ -1,30 +1,15 @@ //! 配置观察者模块 //! //! 核心逻辑已迁移到 proxycast-config crate。 -//! 本模块保留 Tauri 相关实现和重新导出。 +//! 本模块保留 Tauri 相关实现和必要的重新导出。 mod tauri_emitter; mod tauri_observer; -// 从 proxycast-config crate 重新导出所有类型 -pub use proxycast_config::observer::emitter::{ConfigEventEmit, NoOpEmitter}; +// 从 proxycast-config crate 重新导出被使用的类型 pub use proxycast_config::observer::events::{ - AmpConfigChangeEvent, ConfigChangeEvent, ConfigChangeSource, CredentialPoolChangeEvent, - EndpointProvidersChangeEvent, FullReloadEvent, InjectionChangeEvent, LoggingChangeEvent, - NativeAgentChangeEvent, RetryChangeEvent, RoutingChangeEvent, ServerChangeEvent, + ConfigChangeEvent, ConfigChangeSource, EndpointProvidersChangeEvent, RoutingChangeEvent, }; -pub use proxycast_config::observer::manager::GlobalConfigManager; -pub use proxycast_config::observer::observers::{ - DefaultProviderRefObserver, EndpointObserver, InjectorObserver, LoggingObserver, RouterObserver, -}; -pub use proxycast_config::observer::subject::{ - ConfigSubject, CONFIG_CHANGED_EVENT, CONFIG_RELOAD_EVENT, -}; -pub use proxycast_config::observer::traits::{ - ConfigObserver, FnObserver, SyncConfigObserver, SyncObserverWrapper, -}; -pub use proxycast_config::GlobalConfigManagerState; // Tauri 相关实现 pub use tauri_emitter::TauriConfigEmitter; -pub use tauri_observer::TauriObserver; diff --git a/src-tauri/src/config/observer/tauri_observer.rs b/src-tauri/src/config/observer/tauri_observer.rs index 0ced00e61..761cef213 100644 --- a/src-tauri/src/config/observer/tauri_observer.rs +++ b/src-tauri/src/config/observer/tauri_observer.rs @@ -13,6 +13,7 @@ pub struct TauriObserver { app_handle: AppHandle, } +#[allow(dead_code)] impl TauriObserver { pub fn new(app_handle: AppHandle) -> Self { Self { app_handle } diff --git a/src-tauri/src/config/tests.rs b/src-tauri/src/config/tests.rs index f89e42162..c12f1f689 100644 --- a/src-tauri/src/config/tests.rs +++ b/src-tauri/src/config/tests.rs @@ -2,13 +2,13 @@ //! //! 使用 proptest 进行属性测试 -use crate::config::{ +use proptest::prelude::*; +use proxycast_core::config::{ collapse_tilde, contains_tilde, expand_tilde, Config, ConfigManager, CustomProviderConfig, HotReloadManager, InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig, ReloadResult, RetrySettings, RoutingConfig, ServerConfig, YamlService, }; -use crate::config::{ContentCreatorConfig, NavigationConfig}; -use proptest::prelude::*; +use proxycast_core::config::{ContentCreatorConfig, NavigationConfig}; use std::io::Write; use tempfile::NamedTempFile; @@ -38,7 +38,7 @@ fn arb_server_config() -> impl Strategy { host, port, api_key, - tls: crate::config::TlsConfig::default(), + tls: proxycast_core::config::TlsConfig::default(), }) } @@ -179,17 +179,17 @@ fn arb_config() -> impl Strategy { logging, injection: InjectionSettings::default(), auth_dir: "~/.proxycast/auth".to_string(), - credential_pool: crate::config::CredentialPoolConfig::default(), - remote_management: crate::config::RemoteManagementConfig::default(), - quota_exceeded: crate::config::QuotaExceededConfig::default(), + credential_pool: proxycast_core::config::CredentialPoolConfig::default(), + remote_management: proxycast_core::config::RemoteManagementConfig::default(), + quota_exceeded: proxycast_core::config::QuotaExceededConfig::default(), proxy_url: None, - ampcode: crate::config::AmpConfig::default(), - endpoint_providers: crate::config::EndpointProvidersConfig::default(), + ampcode: proxycast_core::config::AmpConfig::default(), + endpoint_providers: proxycast_core::config::EndpointProvidersConfig::default(), minimize_to_tray: true, - models: crate::config::ModelsConfig::default(), - agent: crate::config::NativeAgentConfig::default(), + models: proxycast_core::config::ModelsConfig::default(), + agent: proxycast_core::config::NativeAgentConfig::default(), language: "zh".to_string(), - experimental: crate::config::ExperimentalFeatures::default(), + experimental: proxycast_core::config::ExperimentalFeatures::default(), content_creator: ContentCreatorConfig::default(), navigation: NavigationConfig::default(), }) @@ -356,7 +356,7 @@ fn arb_valid_server_config() -> impl Strategy { host, port, api_key, - tls: crate::config::TlsConfig::default(), + tls: proxycast_core::config::TlsConfig::default(), }) } @@ -419,17 +419,17 @@ fn arb_valid_config() -> impl Strategy { logging, injection: InjectionSettings::default(), auth_dir: "~/.proxycast/auth".to_string(), - credential_pool: crate::config::CredentialPoolConfig::default(), - remote_management: crate::config::RemoteManagementConfig::default(), - quota_exceeded: crate::config::QuotaExceededConfig::default(), + credential_pool: proxycast_core::config::CredentialPoolConfig::default(), + remote_management: proxycast_core::config::RemoteManagementConfig::default(), + quota_exceeded: proxycast_core::config::QuotaExceededConfig::default(), proxy_url: None, - ampcode: crate::config::AmpConfig::default(), - endpoint_providers: crate::config::EndpointProvidersConfig::default(), + ampcode: proxycast_core::config::AmpConfig::default(), + endpoint_providers: proxycast_core::config::EndpointProvidersConfig::default(), minimize_to_tray: true, - models: crate::config::ModelsConfig::default(), - agent: crate::config::NativeAgentConfig::default(), + models: proxycast_core::config::ModelsConfig::default(), + agent: proxycast_core::config::NativeAgentConfig::default(), language: "zh".to_string(), - experimental: crate::config::ExperimentalFeatures::default(), + experimental: proxycast_core::config::ExperimentalFeatures::default(), content_creator: ContentCreatorConfig::default(), navigation: NavigationConfig::default(), }) @@ -470,17 +470,17 @@ fn arb_invalid_config() -> impl Strategy { logging, injection: InjectionSettings::default(), auth_dir: "~/.proxycast/auth".to_string(), - credential_pool: crate::config::CredentialPoolConfig::default(), - remote_management: crate::config::RemoteManagementConfig::default(), - quota_exceeded: crate::config::QuotaExceededConfig::default(), + credential_pool: proxycast_core::config::CredentialPoolConfig::default(), + remote_management: proxycast_core::config::RemoteManagementConfig::default(), + quota_exceeded: proxycast_core::config::QuotaExceededConfig::default(), proxy_url: None, - ampcode: crate::config::AmpConfig::default(), - endpoint_providers: crate::config::EndpointProvidersConfig::default(), + ampcode: proxycast_core::config::AmpConfig::default(), + endpoint_providers: proxycast_core::config::EndpointProvidersConfig::default(), minimize_to_tray: true, - models: crate::config::ModelsConfig::default(), - agent: crate::config::NativeAgentConfig::default(), + models: proxycast_core::config::ModelsConfig::default(), + agent: proxycast_core::config::NativeAgentConfig::default(), language: "zh".to_string(), - experimental: crate::config::ExperimentalFeatures::default(), + experimental: proxycast_core::config::ExperimentalFeatures::default(), content_creator: ContentCreatorConfig::default(), navigation: NavigationConfig::default(), }; @@ -1178,7 +1178,7 @@ providers: // Property 4: Export Scope Filtering // ============================================================================ -use crate::config::{ +use proxycast_core::config::{ ApiKeyEntry, CredentialEntry, CredentialPoolConfig, ExportOptions, ExportService, }; @@ -1357,7 +1357,7 @@ proptest! { // Property 5: Redaction Completeness // ============================================================================ -use crate::config::REDACTED_PLACEHOLDER; +use proxycast_core::config::REDACTED_PLACEHOLDER; /// 生成包含敏感信息的配置 fn arb_config_with_secrets() -> impl Strategy { @@ -1589,7 +1589,7 @@ proptest! { // Property 6: Import Validation // ============================================================================ -use crate::config::{ExportBundle, ImportService}; +use proxycast_core::config::{ExportBundle, ImportService}; /// 生成有效的导出包 fn arb_valid_export_bundle() -> impl Strategy { @@ -1761,7 +1761,7 @@ proptest! { // Property 7: Import Merge vs Replace // ============================================================================ -use crate::config::ImportOptions; +use proxycast_core::config::ImportOptions; proptest! { #![proptest_config(ProptestConfig::with_cases(100))] @@ -2176,7 +2176,7 @@ fn arb_oauth_credential_entry() -> impl Strategy { } /// 生成随机的 Gemini API Key 条目 -fn arb_gemini_api_key_entry() -> impl Strategy { +fn arb_gemini_api_key_entry() -> impl Strategy { ( "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), "AIzaSy[a-zA-Z0-9_-]{33}".prop_map(|s| s), @@ -2187,7 +2187,7 @@ fn arb_gemini_api_key_entry() -> impl Strategy impl Strategy impl Strategy { +fn arb_vertex_api_key_entry() -> impl Strategy { ( "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), "vk-[a-zA-Z0-9]{20,40}".prop_map(|s| s), @@ -2216,13 +2216,13 @@ fn arb_vertex_api_key_entry() -> impl Strategy(), ) .prop_map(|(id, api_key, base_url, models, proxy_url, disabled)| { - crate::config::VertexApiKeyEntry { + proxycast_core::config::VertexApiKeyEntry { id, api_key, base_url, models: models .into_iter() - .map(|(name, alias)| crate::config::VertexModelAlias { name, alias }) + .map(|(name, alias)| proxycast_core::config::VertexModelAlias { name, alias }) .collect(), proxy_url, disabled, @@ -2346,7 +2346,7 @@ proptest! { // Property 3: EndpointProvidersConfig 序列化往返一致性 // ============================================================================ -use crate::config::EndpointProvidersConfig; +use proxycast_core::config::EndpointProvidersConfig; /// 生成随机的 Provider 名称 fn arb_provider_name() -> impl Strategy { diff --git a/src-tauri/src/content/mod.rs b/src-tauri/src/content/mod.rs deleted file mode 100644 index 21accbb49..000000000 --- a/src-tauri/src/content/mod.rs +++ /dev/null @@ -1,3 +0,0 @@ -//! Content 模块 - 从 proxycast-core 重新导出 - -pub use proxycast_core::content::*; diff --git a/src-tauri/src/credential/mod.rs b/src-tauri/src/credential/mod.rs deleted file mode 100644 index 9b6c28718..000000000 --- a/src-tauri/src/credential/mod.rs +++ /dev/null @@ -1,29 +0,0 @@ -//! 凭证池管理模块 -//! -//! 提供多凭证管理、负载均衡和健康检查功能 -//! -//! ## 模块结构 -//! -//! - `types` - 凭证相关类型定义(来自 proxycast-core) -//! - `pool` - 凭证池管理(来自 proxycast-core) -//! - `health` - 健康检查(来自 proxycast-core) -//! - `risk` - 风控模块(来自 proxycast-core) -//! - `balancer` - 负载均衡策略(来自 proxycast-credential) -//! - `quota` - 配额管理(来自 proxycast-credential) -//! - `sync` - 数据库同步(来自 proxycast-credential) - -// 从 proxycast-core 重新导出核心类型模块 -pub use proxycast_core::credential::{health, pool, risk, types}; - -// 重新导出 core 类型 -pub use proxycast_core::credential::{ - CooldownConfig, Credential, CredentialData, CredentialPool, CredentialStats, CredentialStatus, - HealthCheckConfig, HealthCheckResult, HealthChecker, HealthStatus, PoolError, PoolStatus, - RateLimitEvent, RateLimitStats, RiskController, RiskLevel, -}; - -// 从 proxycast-credential crate 重新导出 -pub use proxycast_credential::*; - -#[cfg(test)] -mod tests; diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs deleted file mode 100644 index b1bcb272d..000000000 --- a/src-tauri/src/database/mod.rs +++ /dev/null @@ -1,5 +0,0 @@ -//! 数据库模块 - 从 proxycast-core 重新导出 -//! -//! 核心数据库类型和 DAO 已迁移到 proxycast-core crate - -pub use proxycast_core::database::*; diff --git a/src-tauri/src/dev_bridge.rs b/src-tauri/src/dev_bridge.rs index 648e540ee..a79a77da1 100644 --- a/src-tauri/src/dev_bridge.rs +++ b/src-tauri/src/dev_bridge.rs @@ -24,7 +24,7 @@ use tokio::sync::RwLock; #[cfg(debug_assertions)] use tower_http::cors::CorsLayer; -use crate::server::AppState; +use proxycast_server::AppState; #[cfg(debug_assertions)] #[derive(Debug, Deserialize)] diff --git a/src-tauri/src/dev_bridge/dispatcher.rs b/src-tauri/src/dev_bridge/dispatcher.rs index 828075a40..1212fe7d3 100644 --- a/src-tauri/src/dev_bridge/dispatcher.rs +++ b/src-tauri/src/dev_bridge/dispatcher.rs @@ -2,59 +2,10 @@ //! //! 将 HTTP 请求路由到现有的 Tauri 命令函数。 -use crate::server::AppState; -use serde::Deserialize; +use proxycast_server::AppState; +use proxycast_server_utils::load_model_registry_provider_ids_from_resources; use serde_json::Value as JsonValue; -#[derive(Debug, Deserialize)] -struct ModelRegistryIndex { - providers: Vec, -} - -fn resolve_models_index_path() -> Option { - let mut candidates: Vec = Vec::new(); - - if let Ok(current_dir) = std::env::current_dir() { - candidates.push(current_dir.join("src-tauri/resources/models/index.json")); - candidates.push(current_dir.join("resources/models/index.json")); - } - - if let Ok(exe_path) = std::env::current_exe() { - if let Some(parent) = exe_path.parent() { - candidates.push(parent.join("resources/models/index.json")); - candidates.push(parent.join("../../src-tauri/resources/models/index.json")); - candidates.push(parent.join("../../../src-tauri/resources/models/index.json")); - candidates.push(parent.join("../Resources/resources/models/index.json")); - candidates.push(parent.join("../../Resources/resources/models/index.json")); - candidates.push(parent.join("../../../Resources/resources/models/index.json")); - } - } - - candidates.into_iter().find(|path| path.exists()) -} - -fn load_model_registry_provider_ids_from_resources() -> Result, String> { - let index_path = - resolve_models_index_path().ok_or_else(|| "未找到 models index.json".to_string())?; - - let index_content = std::fs::read_to_string(&index_path) - .map_err(|e| format!("读取 models index.json 失败 ({index_path:?}): {e}"))?; - - let index: ModelRegistryIndex = serde_json::from_str(&index_content) - .map_err(|e| format!("解析 models index.json 失败: {e}"))?; - - let mut provider_ids: Vec = index - .providers - .into_iter() - .map(|id| id.trim().to_string()) - .filter(|id| !id.is_empty()) - .collect(); - - provider_ids.sort(); - provider_ids.dedup(); - Ok(provider_ids) -} - fn load_model_registry_provider_ids_from_db( state: &AppState, ) -> Result, Box> { @@ -88,16 +39,16 @@ pub async fn handle_command( // ========== P0 - 核心配置 ========== "get_config" => { // 从配置文件读取 - let config_path = crate::config::ConfigManager::default_config_path(); - let manager = crate::config::ConfigManager::load(&config_path)?; + let config_path = proxycast_core::config::ConfigManager::default_config_path(); + let manager = proxycast_core::config::ConfigManager::load(&config_path)?; let config = manager.config(); Ok(serde_json::to_value(config)?) } "save_config" => { // 保存配置到文件 - let config: crate::config::Config = serde_json::from_value(args.unwrap_or_default())?; - crate::config::save_config(&config)?; + let config: proxycast_core::config::Config = serde_json::from_value(args.unwrap_or_default())?; + proxycast_core::config::save_config(&config)?; Ok(serde_json::json!({ "success": true })) } @@ -292,7 +243,7 @@ pub async fn handle_command( if let Some(db) = &state.db { // 简化版本:直接创建会话,不需要 agent_state use crate::database::dao::agent::AgentDao; - use crate::agent::types::AgentSession; + use proxycast_core::agent::types::AgentSession; let session_id = uuid::Uuid::new_v4().to_string(); let model_name = model.clone().unwrap_or_else(|| "claude-sonnet-4-20250514".to_string()); diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 2ebbf37e5..8bcbbc62e 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -14,77 +14,52 @@ // 该警告来自 cocoa/objc 依赖的 msg_send! 宏,是已知的 issue #![allow(unexpected_cfgs)] -// 重新导出子 crate 的类型 -// 注意:主 crate 保留了自己的 data, logger, models 模块,所以只导出 core 的具体类型 -pub use proxycast_core::{LogEntry, LogStore, LogStoreConfig, SharedLogStore}; -// infra crate 的类型通过 proxycast_infra 前缀访问,避免与 core 的 InjectionMode/InjectionRule 冲突 -pub use proxycast_infra::{ - injection, proxy, resilience, telemetry, Failover, FailoverConfig, InjectionConfig, - InjectionMode, InjectionResult, InjectionRule, Injector, LogRotationConfig, LoggerError, - ModelStats, ModelTokenStats, PeriodTokenStats, ProviderStats, ProviderTokenStats, - ProxyClientFactory, ProxyError, ProxyProtocol, RequestLog, RequestLogger, RequestStatus, - Retrier, RetryConfig, StatsAggregator, StatsSummary, TimeRange, TimeoutConfig, - TimeoutController, TokenSource, TokenStatsSummary, TokenTracker, TokenUsageRecord, -}; - // 从 providers crate 重新导出(保持 crate::xxx 路径兼容) -pub use proxycast_providers::converter; pub use proxycast_providers::providers; -pub use proxycast_providers::stream; -pub use proxycast_providers::streaming; -pub use proxycast_providers::translator; // 从 core crate 重新导出(保持 crate::xxx 路径兼容) -pub use proxycast_core::backends; pub use proxycast_core::connect; -pub use proxycast_core::orchestrator; +pub use proxycast_core::content; +pub use proxycast_core::credential; +pub use proxycast_core::database; +pub use proxycast_core::memory; pub use proxycast_core::session_files; +pub use proxycast_core::workspace; -// 核心模块 +// 从 infra crate 重新导出(保持 crate::xxx 路径兼容) +pub use proxycast_infra::{injection, resilience, telemetry}; + +// MCP 模块(从 proxycast-mcp crate 重新导出) +pub use proxycast_mcp as mcp; + +// 核心模块(Tauri 相关业务逻辑) pub mod agent; pub mod app; -pub mod content; -pub mod credential; -pub mod database; -pub mod memory; pub mod plugin; pub mod screenshot; pub mod services; -pub mod session; +pub mod skills; pub mod terminal; pub mod tray; pub mod voice; -pub mod workspace; - -// Skills 集成模块 -pub mod skills; - -// MCP 集成模块 -pub mod mcp; // 内部模块 mod commands; mod config; mod data; #[cfg(debug_assertions)] +#[allow(dead_code)] mod dev_bridge; mod logger; -mod models; -mod server_utils; +use proxycast_core::models; -// 从 core crate 重新导出 errors -pub use proxycast_core::errors; - -// 服务器相关模块 -mod middleware; -mod processor; -mod router; -mod server; -mod websocket; +// 测试模块 +#[cfg(test)] +mod tests; // 重新导出核心类型以保持向后兼容 pub use app::{AppState, LogState, ProviderType, TokenCacheServiceState, TrayManagerState}; -pub use services::provider_pool_service::ProviderPoolService; +pub use proxycast_services::provider_pool_service::ProviderPoolService; -// 重新导出 run 函数 +// 重新导出 run 函数(main.rs 入口) pub use app::run; diff --git a/src-tauri/src/logger.rs b/src-tauri/src/logger.rs index 28d7e4f9f..f7d9a4dae 100644 --- a/src-tauri/src/logger.rs +++ b/src-tauri/src/logger.rs @@ -1,12 +1,6 @@ -//! 日志管理模块 +//! 日志管理模块(桥接层) //! -//! 核心逻辑已迁移到 proxycast-core crate,本文件保留扩展函数。 +//! 纯逻辑已迁移到 `proxycast-core` crate, +//! 本模块保留兼容导出。 pub use proxycast_core::logger::*; - -use crate::config::LoggingConfig; - -/// 使用 LoggingConfig 创建 LogStore -pub fn create_log_store_from_config(logging: &LoggingConfig) -> LogStore { - LogStore::with_custom_config(logging.retention_days, logging.enabled) -} diff --git a/src-tauri/src/mcp/README.md b/src-tauri/src/mcp/README.md deleted file mode 100644 index 33611c9e6..000000000 --- a/src-tauri/src/mcp/README.md +++ /dev/null @@ -1,42 +0,0 @@ -# MCP 模块 - -MCP(Model Context Protocol)集成模块,提供 MCP 协议的客户端实现。 - -## 模块结构 - -| 文件 | 说明 | -|------|------| -| `mod.rs` | 模块导出和文档 | -| `types.rs` | MCP 数据类型定义(配置、工具、提示词、资源、错误) | -| `client.rs` | MCP 客户端实现(rmcp ClientHandler) | -| `manager.rs` | MCP 客户端管理器(连接池、缓存、生命周期) | -| `tool_converter.rs` | 工具格式转换器(OpenAI/Anthropic/Gemini) | - -## 功能概览 - -### 服务器生命周期管理 -- 启动/停止 MCP 服务器进程 -- stdio 传输连接 -- 状态监控和事件通知 - -### 工具管理 -- 工具发现和缓存 -- 工具调用路由 -- 名称冲突解决(服务器前缀) - -### 格式转换 -- MCP → OpenAI function calling -- MCP → Anthropic tool use -- MCP → Gemini function declaration - -## 依赖 - -- `rmcp`: Rust MCP SDK -- `tokio`: 异步运行时 -- `serde`: 序列化/反序列化 -- `thiserror`: 错误类型定义 - -## 相关文档 - -- 设计文档: `.kiro/specs/mcp-integration/design.md` -- 需求文档: `.kiro/specs/mcp-integration/requirements.md` diff --git a/src-tauri/src/mcp/mod.rs b/src-tauri/src/mcp/mod.rs deleted file mode 100644 index 422e623a0..000000000 --- a/src-tauri/src/mcp/mod.rs +++ /dev/null @@ -1,19 +0,0 @@ -//! MCP(Model Context Protocol)模块 -//! -//! 业务逻辑已迁移到 proxycast-mcp crate, -//! 本模块仅作为桥接层 re-export。 - -// 从 proxycast-mcp crate re-export 所有公开类型 -pub use proxycast_mcp::client; -pub use proxycast_mcp::manager; -pub use proxycast_mcp::tool_converter; -pub use proxycast_mcp::types; - -pub use proxycast_mcp::{McpClientManager, ProxyCastMcpClient}; -pub use proxycast_mcp::{ - McpClientWrapper, McpContent, McpError, McpManagerState, McpPromptArgument, - McpPromptDefinition, McpPromptMessage, McpPromptResult, McpResourceContent, - McpResourceDefinition, McpServerCapabilities, McpServerConfig, McpServerErrorPayload, - McpServerInfo, McpServerStartedPayload, McpServerStoppedPayload, McpToolCall, - McpToolDefinition, McpToolResult, McpToolsUpdatedPayload, ToolConverter, -}; diff --git a/src-tauri/src/memory/mod.rs b/src-tauri/src/memory/mod.rs deleted file mode 100644 index 266771c25..000000000 --- a/src-tauri/src/memory/mod.rs +++ /dev/null @@ -1,3 +0,0 @@ -//! Memory 模块 - 从 proxycast-core 重新导出 - -pub use proxycast_core::memory::*; diff --git a/src-tauri/src/middleware/mod.rs b/src-tauri/src/middleware/mod.rs deleted file mode 100644 index 643b4d354..000000000 --- a/src-tauri/src/middleware/mod.rs +++ /dev/null @@ -1,5 +0,0 @@ -//! Middleware 模块 -//! -//! 从 proxycast-core 重新导出 - -pub use proxycast_core::middleware::*; diff --git a/src-tauri/src/models/mod.rs b/src-tauri/src/models/mod.rs deleted file mode 100644 index c031a8689..000000000 --- a/src-tauri/src/models/mod.rs +++ /dev/null @@ -1,39 +0,0 @@ -//! 数据模型模块 -//! -//! 从 proxycast-core crate 重新导出所有模型类型。 -//! 仅保留依赖主 crate 业务模块的类型在本地定义。 - -// 从 core crate 重新导出所有模型 -pub use proxycast_core::models::anthropic; -#[allow(unused_imports)] -pub use proxycast_core::models::app_type; -#[allow(unused_imports)] -pub use proxycast_core::models::codewhisperer; -pub use proxycast_core::models::kiro_fingerprint; -pub use proxycast_core::models::machine_id; -pub use proxycast_core::models::mcp_model; -pub use proxycast_core::models::model_registry; -pub use proxycast_core::models::openai; -#[allow(unused_imports)] -pub use proxycast_core::models::prompt_model; -#[allow(unused_imports)] -pub use proxycast_core::models::provider_model; -pub use proxycast_core::models::provider_pool_model; -pub use proxycast_core::models::route_model; -// project_model 已迁移到 core -pub use proxycast_core::models::project_model; - -// 重新导出常用类型(保持向后兼容) -#[allow(unused_imports)] -pub use proxycast_core::models::anthropic::*; -pub use proxycast_core::models::app_type::AppType; -#[allow(unused_imports)] -pub use proxycast_core::models::codewhisperer::*; -pub use proxycast_core::models::mcp_model::McpServer; -#[allow(unused_imports)] -pub use proxycast_core::models::openai::*; -pub use proxycast_core::models::prompt_model::Prompt; -pub use proxycast_core::models::provider_model::Provider; -#[allow(unused_imports)] -pub use proxycast_core::models::provider_pool_model::*; -pub use proxycast_core::models::skill_model::{Skill, SkillMetadata, SkillRepo, SkillState}; diff --git a/src-tauri/src/plugin/mod.rs b/src-tauri/src/plugin/mod.rs index 5d857dd77..c9f23a79e 100644 --- a/src-tauri/src/plugin/mod.rs +++ b/src-tauri/src/plugin/mod.rs @@ -1,25 +1,10 @@ //! 插件系统模块 //! //! 核心逻辑从 proxycast-core 重新导出, -//! ui_events 依赖 Tauri 保留在主 crate +//! ui_events 依赖 Tauri 保留在主 crate。 -// 从 core 重新导出所有插件类型和模块 -pub use proxycast_core::plugin::binary_downloader; -pub use proxycast_core::plugin::examples; -pub use proxycast_core::plugin::installer; -pub use proxycast_core::plugin::ui_builder; -pub use proxycast_core::plugin::ui_trait; -pub use proxycast_core::plugin::ui_types; - -pub use proxycast_core::plugin::{ - Action, BoundValue, ChildrenDef, ComponentDef, ComponentType, DataEntry, DataModelUpdate, - SurfaceDefinition, SurfaceUpdate, UIMessage, UserAction, -}; -pub use proxycast_core::plugin::{ - BinaryComponentStatus, BinaryDownloader, BinaryManifest, HookResult, NoUI, PlatformBinaries, - Plugin, PluginConfig, PluginContext, PluginError, PluginInfo, PluginLoader, PluginManager, - PluginManifest, PluginState, PluginStatus, PluginType, PluginUI, -}; +// 核心插件能力从 core crate 导出 +pub use proxycast_core::plugin::*; // Tauri 依赖的 UI 事件模块保留在主 crate pub mod ui_events; diff --git a/src-tauri/src/processor/mod.rs b/src-tauri/src/processor/mod.rs deleted file mode 100644 index f59942270..000000000 --- a/src-tauri/src/processor/mod.rs +++ /dev/null @@ -1,9 +0,0 @@ -//! 请求处理器模块(重导出层) -//! -//! 核心逻辑已迁移到 `proxycast-processor` crate。 -//! 本模块保留向后兼容路径和本地测试入口。 - -pub use proxycast_processor::*; - -#[cfg(test)] -mod tests; diff --git a/src-tauri/src/processor/steps/mod.rs b/src-tauri/src/processor/steps/mod.rs deleted file mode 100644 index d9334d1e3..000000000 --- a/src-tauri/src/processor/steps/mod.rs +++ /dev/null @@ -1,4 +0,0 @@ -//! 管道步骤模块(re-export from proxycast-processor) - -#[allow(unused_imports)] -pub use proxycast_processor::steps::*; diff --git a/src-tauri/src/router/mod.rs b/src-tauri/src/router/mod.rs deleted file mode 100644 index a879c7486..000000000 --- a/src-tauri/src/router/mod.rs +++ /dev/null @@ -1,3 +0,0 @@ -//! 路由系统模块 - 从 proxycast-core 重新导出 - -pub use proxycast_core::router::*; diff --git a/src-tauri/src/screenshot/README.md b/src-tauri/src/screenshot/README.md index 78bae5224..aa15a7e5e 100644 --- a/src-tauri/src/screenshot/README.md +++ b/src-tauri/src/screenshot/README.md @@ -7,7 +7,7 @@ ``` screenshot/ ├── mod.rs # 模块入口,导出公共接口和初始化函数 -├── capture.rs # 屏幕截图服务,跨平台截图实现 +├── capture.rs # 截图桥接层(纯逻辑在 crates/services/src/screenshot_capture_service.rs) ├── config.rs # 实验室功能配置管理 ├── shortcut.rs # 全局快捷键注册/注销/更新 ├── window.rs # 悬浮对话窗口管理 @@ -36,8 +36,8 @@ screenshot/ ### 4. 截图服务 (capture.rs) -- `start_capture(app)`: 启动交互式截图 -- 支持 macOS (screencapture)、Windows (PowerShell)、Linux (gnome-screenshot/scrot) +- `start_capture()`: 启动交互式截图 +- 支持 macOS (`screencapture`)、Windows (PowerShell)、Linux (`gnome-screenshot` / `scrot`) ### 5. 窗口管理 (window.rs) diff --git a/src-tauri/src/screenshot/capture.rs b/src-tauri/src/screenshot/capture.rs index 9dcd97f72..ab4ac26a7 100644 --- a/src-tauri/src/screenshot/capture.rs +++ b/src-tauri/src/screenshot/capture.rs @@ -1,248 +1,8 @@ -//! 屏幕截图服务 +//! 屏幕截图服务(桥接层) //! -//! 提供跨平台的屏幕截图功能,支持交互式区域选择 +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块保留兼容导出。 -use std::path::PathBuf; -use tauri::AppHandle; -use tracing::{debug, error, info}; - -#[cfg(target_os = "macos")] -use tracing::warn; - -/// 截图错误类型 -#[derive(Debug, thiserror::Error)] -pub enum CaptureError { - #[error("用户取消截图")] - Cancelled, - #[error("权限不足: {0}")] - PermissionDenied(String), - #[error("系统错误: {0}")] - SystemError(String), - #[error("IO 错误: {0}")] - IoError(#[from] std::io::Error), - #[error("临时文件创建失败: {0}")] - TempFileError(String), -} - -/// 截图结果 -#[derive(Debug)] -pub enum CaptureResult { - /// 截图成功,包含图片路径 - Success(PathBuf), - /// 用户取消截图 - Cancelled, -} - -/// 启动交互式截图 -/// -/// 调用系统截图工具进行交互式区域选择 -/// -/// # 参数 -/// - `app`: Tauri 应用句柄 -/// -/// # 返回 -/// 成功返回图片路径,用户取消返回 Cancelled 错误,其他情况返回相应错误 -/// -/// # 平台支持 -/// - macOS: 使用 `screencapture -i -x` 命令 -/// - Windows: 使用 Windows API 或系统截图工具 -/// - Linux: 使用 `gnome-screenshot` 或 `scrot` -pub async fn start_capture(_app: &AppHandle) -> Result { - info!("启动交互式截图"); - - // 生成临时文件路径 - let temp_dir = std::env::temp_dir(); - let timestamp = chrono::Utc::now().format("%Y%m%d_%H%M%S_%3f"); - let filename = format!("proxycast_screenshot_{timestamp}.png"); - let temp_path = temp_dir.join(&filename); - - debug!("截图临时文件路径: {:?}", temp_path); - - // 根据平台调用不同的截图命令 - #[cfg(target_os = "macos")] - { - capture_macos(&temp_path).await?; - } - - #[cfg(target_os = "windows")] - { - capture_windows(&temp_path).await?; - } - - #[cfg(target_os = "linux")] - { - capture_linux(&temp_path).await?; - } - - // 检查文件是否存在(用户可能取消了截图) - if !temp_path.exists() { - info!("截图文件不存在,用户可能取消了截图"); - return Err(CaptureError::Cancelled); - } - - // 检查文件大小(空文件表示取消) - let metadata = std::fs::metadata(&temp_path)?; - if metadata.len() == 0 { - info!("截图文件为空,用户取消了截图"); - std::fs::remove_file(&temp_path)?; - return Err(CaptureError::Cancelled); - } - - info!("截图成功: {:?}", temp_path); - Ok(temp_path) -} - -/// macOS 截图实现 -#[cfg(target_os = "macos")] -async fn capture_macos(output_path: &PathBuf) -> Result<(), CaptureError> { - use std::process::Command; - - debug!("使用 macOS screencapture 命令"); - - // 先检查屏幕录制权限 - // 通过尝试执行一个快速的全屏截图到 /dev/null 来检测权限 - let permission_check = Command::new("screencapture") - .args(["-x", "-c"]) // -c 截图到剪贴板,快速检测权限 - .output(); - - if let Ok(output) = permission_check { - if !output.status.success() { - let stderr = String::from_utf8_lossy(&output.stderr); - if stderr.contains("cannot") || stderr.contains("permission") { - warn!("屏幕录制权限不足: {}", stderr); - return Err(CaptureError::PermissionDenied( - "请在系统设置 → 隐私与安全性 → 录屏与系统录音 中授权 ProxyCast".to_string(), - )); - } - } - } - - // 使用 screencapture 命令 - // -i: 交互式选择区域 - // -x: 不播放截图声音 - let output = Command::new("screencapture") - .args(["-i", "-x", output_path.to_str().unwrap()]) - .output() - .map_err(|e| CaptureError::SystemError(format!("执行 screencapture 失败: {e}")))?; - - if !output.status.success() { - let stderr = String::from_utf8_lossy(&output.stderr); - // 用户按 ESC 取消时,screencapture 返回非零状态但不创建文件 - if !output_path.exists() { - return Err(CaptureError::Cancelled); - } - error!("screencapture 命令失败: {}", stderr); - return Err(CaptureError::SystemError(format!( - "screencapture 失败: {stderr}" - ))); - } - - Ok(()) -} - -/// Windows 截图实现 -#[cfg(target_os = "windows")] -async fn capture_windows(output_path: &PathBuf) -> Result<(), CaptureError> { - use std::process::Command; - - debug!("使用 Windows 截图工具"); - - // 使用 PowerShell 调用截图功能 - // 注意:这是一个简化实现,实际可能需要使用 Windows API - let script = format!( - r#" - Add-Type -AssemblyName System.Windows.Forms - $screen = [System.Windows.Forms.Screen]::PrimaryScreen - $bitmap = New-Object System.Drawing.Bitmap($screen.Bounds.Width, $screen.Bounds.Height) - $graphics = [System.Drawing.Graphics]::FromImage($bitmap) - $graphics.CopyFromScreen($screen.Bounds.Location, [System.Drawing.Point]::Empty, $screen.Bounds.Size) - $bitmap.Save('{}') - "#, - output_path.to_str().unwrap().replace("\\", "\\\\") - ); - - let output = Command::new("powershell") - .args(["-Command", &script]) - .output() - .map_err(|e| CaptureError::SystemError(format!("执行 PowerShell 失败: {}", e)))?; - - if !output.status.success() { - let stderr = String::from_utf8_lossy(&output.stderr); - error!("PowerShell 截图失败: {}", stderr); - return Err(CaptureError::SystemError(format!("截图失败: {}", stderr))); - } - - Ok(()) -} - -/// Linux 截图实现 -#[cfg(target_os = "linux")] -async fn capture_linux(output_path: &PathBuf) -> Result<(), CaptureError> { - use std::process::Command; - - debug!("使用 Linux 截图工具"); - - // 尝试使用 gnome-screenshot - let gnome_result = Command::new("gnome-screenshot") - .args(["-a", "-f", output_path.to_str().unwrap()]) - .output(); - - if let Ok(output) = gnome_result { - if output.status.success() { - return Ok(()); - } - } - - // 回退到 scrot - let scrot_result = Command::new("scrot") - .args(["-s", output_path.to_str().unwrap()]) - .output() - .map_err(|e| { - CaptureError::SystemError(format!( - "无法找到截图工具 (gnome-screenshot 或 scrot): {}", - e - )) - })?; - - if !scrot_result.status.success() { - let stderr = String::from_utf8_lossy(&scrot_result.stderr); - if !output_path.exists() { - return Err(CaptureError::Cancelled); - } - return Err(CaptureError::SystemError(format!("scrot 失败: {}", stderr))); - } - - Ok(()) -} - -/// 清理临时截图文件 -/// -/// # 参数 -/// - `path`: 要删除的文件路径 -pub fn cleanup_temp_file(path: &PathBuf) { - if path.exists() { - if let Err(e) = std::fs::remove_file(path) { - error!("删除临时截图文件失败: {}", e); - } else { - debug!("已删除临时截图文件: {:?}", path); - } - } -} - -#[cfg(test)] -mod tests { - - #[test] - fn test_temp_path_generation() { - let temp_dir = std::env::temp_dir(); - let timestamp = chrono::Utc::now().format("%Y%m%d_%H%M%S_%3f"); - let filename = format!("proxycast_screenshot_{timestamp}.png"); - let temp_path = temp_dir.join(&filename); - - assert!(temp_path - .to_str() - .unwrap() - .contains("proxycast_screenshot_")); - assert!(temp_path.to_str().unwrap().ends_with(".png")); - } -} +pub use proxycast_services::screenshot_capture_service::{ + cleanup_temp_file, start_capture, CaptureError, CaptureResult, +}; diff --git a/src-tauri/src/screenshot/config.rs b/src-tauri/src/screenshot/config.rs index a9a0950e3..7b8086367 100644 --- a/src-tauri/src/screenshot/config.rs +++ b/src-tauri/src/screenshot/config.rs @@ -2,7 +2,8 @@ //! //! 提供实验室功能配置的加载和检查功能 -use crate::config::{ExperimentalFeatures, GlobalConfigManagerState, ScreenshotChatConfig}; +use crate::config::GlobalConfigManagerState; +use proxycast_core::config::{ExperimentalFeatures, ScreenshotChatConfig}; use tauri::{AppHandle, Manager}; use tracing::debug; diff --git a/src-tauri/src/screenshot/shortcut.rs b/src-tauri/src/screenshot/shortcut.rs index f428079dc..c60a22bc4 100644 --- a/src-tauri/src/screenshot/shortcut.rs +++ b/src-tauri/src/screenshot/shortcut.rs @@ -252,7 +252,7 @@ fn handle_shortcut_triggered(app: &AppHandle) { // 短暂延迟,确保窗口完全最小化 tokio::time::sleep(std::time::Duration::from_millis(100)).await; - match capture::start_capture(&app_clone).await { + match capture::start_capture().await { Ok(image_path) => { info!("截图成功: {:?}", image_path); // 打开悬浮窗口 diff --git a/src-tauri/src/server/mod.rs b/src-tauri/src/server/mod.rs deleted file mode 100644 index 11e3b3b50..000000000 --- a/src-tauri/src/server/mod.rs +++ /dev/null @@ -1,5 +0,0 @@ -//! HTTP API 服务器 -//! -//! 核心逻辑已迁移到 proxycast-server crate,本模块仅做重新导出。 - -pub use proxycast_server::*; diff --git a/src-tauri/src/server_utils.rs b/src-tauri/src/server_utils.rs deleted file mode 100644 index 8323e1c2c..000000000 --- a/src-tauri/src/server_utils.rs +++ /dev/null @@ -1,5 +0,0 @@ -//! 服务器工具函数 -//! -//! 已迁移至 proxycast-server-utils crate,此文件为 re-export 层。 - -pub use proxycast_server_utils::*; diff --git a/src-tauri/src/services/README.md b/src-tauri/src/services/README.md index 52abbbdca..95a27d4ce 100644 --- a/src-tauri/src/services/README.md +++ b/src-tauri/src/services/README.md @@ -21,8 +21,10 @@ - `backup_service.rs` - 备份服务 - `live_sync.rs` - 实时同步服务 - `switch.rs` - 开关服务 -- `sysinfo_service.rs` - 系统信息服务(CPU/内存监控) -- `file_browser_service.rs` - 文件浏览器服务(目录列表、文件预览) +- `sysinfo_service.rs` - Tauri 命令桥接(纯逻辑已迁移到 `crates/services/src/sysinfo_service.rs`) +- `file_browser_service.rs` - Tauri 命令桥接(纯逻辑已迁移到 `crates/services/src/file_browser_service.rs`) +- `update_check_service.rs` - 兼容导出层(纯逻辑已迁移到 `crates/services/src/update_check_service.rs`) +- `update_window.rs` - 更新提醒独立窗口管理 - `general_chat/` - 通用对话服务模块(会话管理、消息存储) - `api_key_provider_service.rs` - API Key Provider 服务 - `kiro_event_service.rs` - Kiro 事件服务 @@ -31,8 +33,20 @@ - `persona_service.rs` - 人设服务(创建、列表、更新、删除、设置默认、模板) - `material_service.rs` - 素材服务(上传、存储、删除、内容读取) - `template_service.rs` - 排版模板服务(创建、列表、更新、删除、设置默认) -- `update_check_service.rs` - 自动更新检查服务(每日检查、系统通知) -- `update_window.rs` - 更新提醒独立窗口管理 + +## 已迁移补充 + +以下语音相关纯逻辑已迁移到 `crates/services/src/`: +- `voice_asr_service.rs` - ASR 识别与云端失败回退 +- `voice_config_service.rs` - 语音配置、ASR 凭证与指令管理 +- `voice_processor_service.rs` - 文本模板处理与 LLM 润色 +- `voice_output_service.rs` - 文本输出模式与系统输出 +- `voice_command_service.rs` - 转写/润色/输出业务流程 +- `voice_recording_service.rs` - 录音状态封装与设备查询 + +以下截图相关纯逻辑已迁移到 `crates/services/src/`: +- `screenshot_capture_service.rs` - 跨平台截图与临时文件清理 +- `screenshot_image_service.rs` - 图片读取与 Base64 编码 ## Aster Agent 集成 diff --git a/src-tauri/src/services/file_browser_service.rs b/src-tauri/src/services/file_browser_service.rs index f0c271875..1c2cde995 100644 --- a/src-tauri/src/services/file_browser_service.rs +++ b/src-tauri/src/services/file_browser_service.rs @@ -1,526 +1,15 @@ -//! 文件浏览器服务 +//! 文件浏览器服务(Tauri 命令桥接层) //! -//! 提供文件系统浏览功能 -//! 支持目录列表、文件预览等操作 -//! -//! # 功能 -//! - 列出目录内容 -//! - 读取文件预览 -//! - 获取文件元信息 -//! - 获取文件权限和 MIME 类型 +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块仅保留 Tauri 命令封装。 -use serde::{Deserialize, Serialize}; -use std::fs::{self, Metadata}; -#[cfg(unix)] -use std::os::unix::fs::PermissionsExt; -use std::path::{Path, PathBuf}; -use std::time::UNIX_EPOCH; -use tracing::{debug, error}; - -/// 文件条目 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FileEntry { - /// 文件名 - pub name: String, - /// 完整路径 - pub path: String, - /// 是否为目录 - #[serde(rename = "isDir")] - pub is_dir: bool, - /// 文件大小(字节) - pub size: u64, - /// 修改时间(Unix 时间戳毫秒) - #[serde(rename = "modifiedAt")] - pub modified_at: u64, - /// 文件类型/扩展名 - #[serde(rename = "fileType")] - pub file_type: Option, - /// 是否隐藏文件 - #[serde(rename = "isHidden")] - pub is_hidden: bool, - /// 文件权限字符串(如 -rw-r--r--) - #[serde(rename = "modeStr")] - pub mode_str: Option, - /// 文件权限数字(8进制) - pub mode: Option, - /// MIME 类型 - #[serde(rename = "mimeType")] - pub mime_type: Option, - /// 是否为符号链接 - #[serde(rename = "isSymlink")] - pub is_symlink: bool, -} - -/// 目录列表结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DirectoryListing { - /// 当前路径 - pub path: String, - /// 父目录路径 - #[serde(rename = "parentPath")] - pub parent_path: Option, - /// 文件列表 - pub entries: Vec, - /// 错误信息 - pub error: Option, -} - -/// 文件预览结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FilePreview { - /// 文件路径 - pub path: String, - /// 文件内容(文本) - pub content: Option, - /// 是否为二进制文件 - #[serde(rename = "isBinary")] - pub is_binary: bool, - /// 文件大小 - pub size: u64, - /// 错误信息 - pub error: Option, -} - -/// 获取文件扩展名 -fn get_file_extension(path: &Path) -> Option { - path.extension() - .and_then(|ext| ext.to_str()) - .map(|s| s.to_lowercase()) -} - -/// 判断是否为隐藏文件 -fn is_hidden_file(name: &str) -> bool { - name.starts_with('.') -} - -/// 将 Unix 文件模式转换为权限字符串(如 -rw-r--r--) -#[cfg(unix)] -fn mode_to_string(mode: u32, is_dir: bool, is_symlink: bool) -> String { - let mut result = String::with_capacity(10); - - // 文件类型 - if is_symlink { - result.push('l'); - } else if is_dir { - result.push('d'); - } else { - result.push('-'); - } - - // 用户权限 - result.push(if mode & 0o400 != 0 { 'r' } else { '-' }); - result.push(if mode & 0o200 != 0 { 'w' } else { '-' }); - result.push(if mode & 0o100 != 0 { 'x' } else { '-' }); - - // 组权限 - result.push(if mode & 0o040 != 0 { 'r' } else { '-' }); - result.push(if mode & 0o020 != 0 { 'w' } else { '-' }); - result.push(if mode & 0o010 != 0 { 'x' } else { '-' }); - - // 其他用户权限 - result.push(if mode & 0o004 != 0 { 'r' } else { '-' }); - result.push(if mode & 0o002 != 0 { 'w' } else { '-' }); - result.push(if mode & 0o001 != 0 { 'x' } else { '-' }); - - result -} - -/// 根据文件扩展名和元数据获取 MIME 类型 -fn get_mime_type(path: &Path, metadata: &Metadata) -> String { - // 特殊类型检测 - if metadata.is_dir() { - return "directory".to_string(); - } - - #[cfg(unix)] - { - use std::os::unix::fs::FileTypeExt; - let ft = metadata.file_type(); - if ft.is_symlink() { - return "symlink".to_string(); - } - if ft.is_block_device() { - return "block-device".to_string(); - } - if ft.is_char_device() { - return "char-device".to_string(); - } - if ft.is_fifo() { - return "pipe".to_string(); - } - if ft.is_socket() { - return "socket".to_string(); - } - } - - // 基于扩展名的 MIME 类型映射 - let ext = path - .extension() - .and_then(|e| e.to_str()) - .map(|s| s.to_lowercase()); - - match ext.as_deref() { - // 文本文件 - Some("txt") => "text/plain", - Some("md" | "markdown") => "text/markdown", - Some("html" | "htm") => "text/html", - Some("css") => "text/css", - Some("xml") => "text/xml", - Some("csv") => "text/csv", - - // 代码文件 - Some("js" | "mjs" | "cjs") => "text/javascript", - Some("ts" | "mts" | "cts") => "text/typescript", - Some("tsx") => "text/tsx", - Some("jsx") => "text/jsx", - Some("json") => "application/json", - Some("yaml" | "yml") => "text/yaml", - Some("toml") => "text/toml", - Some("rs") => "text/x-rust", - Some("py") => "text/x-python", - Some("go") => "text/x-go", - Some("java") => "text/x-java", - Some("c") => "text/x-c", - Some("cpp" | "cc" | "cxx") => "text/x-c++", - Some("h" | "hpp") => "text/x-c-header", - Some("sh" | "bash" | "zsh") => "text/x-shellscript", - Some("sql") => "text/x-sql", - Some("vue") => "text/x-vue", - Some("svelte") => "text/x-svelte", - Some("swift") => "text/x-swift", - Some("kt" | "kts") => "text/x-kotlin", - Some("rb") => "text/x-ruby", - Some("php") => "text/x-php", - Some("lua") => "text/x-lua", - - // 图片 - Some("png") => "image/png", - Some("jpg" | "jpeg") => "image/jpeg", - Some("gif") => "image/gif", - Some("webp") => "image/webp", - Some("svg") => "image/svg+xml", - Some("ico") => "image/x-icon", - Some("bmp") => "image/bmp", - Some("tiff" | "tif") => "image/tiff", - - // 音频 - Some("mp3") => "audio/mpeg", - Some("wav") => "audio/wav", - Some("ogg") => "audio/ogg", - Some("flac") => "audio/flac", - Some("aac") => "audio/aac", - Some("m4a") => "audio/mp4", - - // 视频 - Some("mp4") => "video/mp4", - Some("webm") => "video/webm", - Some("avi") => "video/x-msvideo", - Some("mov") => "video/quicktime", - Some("mkv") => "video/x-matroska", - Some("wmv") => "video/x-ms-wmv", - - // 压缩文件 - Some("zip") => "application/zip", - Some("tar") => "application/x-tar", - Some("gz" | "gzip") => "application/gzip", - Some("bz2") => "application/x-bzip2", - Some("xz") => "application/x-xz", - Some("7z") => "application/x-7z-compressed", - Some("rar") => "application/vnd.rar", - - // 文档 - Some("pdf") => "application/pdf", - Some("doc") => "application/msword", - Some("docx") => "application/vnd.openxmlformats-officedocument.wordprocessingml.document", - Some("xls") => "application/vnd.ms-excel", - Some("xlsx") => "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", - Some("ppt") => "application/vnd.ms-powerpoint", - Some("pptx") => "application/vnd.openxmlformats-officedocument.presentationml.presentation", - - // 可执行文件 - Some("exe") => "application/x-msdownload", - Some("dmg") => "application/x-apple-diskimage", - Some("app") => "application/x-apple-application", - Some("deb") => "application/x-debian-package", - Some("rpm") => "application/x-rpm", - - // 字体 - Some("ttf") => "font/ttf", - Some("otf") => "font/otf", - Some("woff") => "font/woff", - Some("woff2") => "font/woff2", - - // 其他 - Some("wasm") => "application/wasm", - - _ => "application/octet-stream", - } - .to_string() -} - -/// 判断是否为文本文件(基于扩展名) -fn is_text_file(extension: Option<&str>) -> bool { - match extension { - Some(ext) => matches!( - ext, - "txt" - | "md" - | "json" - | "yaml" - | "yml" - | "toml" - | "xml" - | "html" - | "htm" - | "css" - | "js" - | "ts" - | "tsx" - | "jsx" - | "rs" - | "py" - | "go" - | "java" - | "c" - | "cpp" - | "h" - | "hpp" - | "sh" - | "bash" - | "zsh" - | "fish" - | "sql" - | "graphql" - | "vue" - | "svelte" - | "astro" - | "log" - | "env" - | "gitignore" - | "dockerignore" - | "editorconfig" - | "prettierrc" - | "eslintrc" - | "babelrc" - | "conf" - | "cfg" - | "ini" - | "properties" - ), - None => false, - } -} - -/// 列出目录内容 -pub fn list_directory(path: &str) -> DirectoryListing { - let path_buf = if path.is_empty() || path == "~" { - dirs::home_dir().unwrap_or_else(|| PathBuf::from("/")) - } else if path.starts_with('~') { - let home = dirs::home_dir().unwrap_or_else(|| PathBuf::from("/")); - home.join(&path[2..]) - } else { - PathBuf::from(path) - }; - - let canonical_path = match path_buf.canonicalize() { - Ok(p) => p, - Err(e) => { - error!("无法解析路径 {}: {}", path, e); - return DirectoryListing { - path: path.to_string(), - parent_path: None, - entries: vec![], - error: Some(format!("无法解析路径: {e}")), - }; - } - }; - - let parent_path = canonical_path - .parent() - .map(|p| p.to_string_lossy().to_string()); - - let entries = match fs::read_dir(&canonical_path) { - Ok(read_dir) => { - let mut entries: Vec = read_dir - .filter_map(|entry| { - let entry = entry.ok()?; - // 先获取符号链接信息 - let symlink_metadata = entry.metadata().ok(); - let is_symlink = entry - .file_type() - .ok() - .map(|ft| ft.is_symlink()) - .unwrap_or(false); - // 获取真实文件的元数据(解析符号链接) - let metadata = if is_symlink { - fs::metadata(entry.path()).ok().or(symlink_metadata)? - } else { - symlink_metadata? - }; - let name = entry.file_name().to_string_lossy().to_string(); - let path = entry.path(); - - let modified_at = metadata - .modified() - .ok() - .and_then(|t| t.duration_since(UNIX_EPOCH).ok()) - .map(|d| d.as_millis() as u64) - .unwrap_or(0); - - let file_type = if metadata.is_dir() { - Some("folder".to_string()) - } else { - get_file_extension(&path) - }; - - // 获取文件权限(仅 Unix) - #[cfg(unix)] - let (mode, mode_str) = { - let m = metadata.permissions().mode(); - ( - Some(m & 0o777), - Some(mode_to_string(m & 0o777, metadata.is_dir(), is_symlink)), - ) - }; - #[cfg(not(unix))] - let (mode, mode_str): (Option, Option) = (None, None); - - // 获取 MIME 类型 - let mime_type = get_mime_type(&path, &metadata); - - Some(FileEntry { - name: name.clone(), - path: path.to_string_lossy().to_string(), - is_dir: metadata.is_dir(), - size: metadata.len(), - modified_at, - file_type, - is_hidden: is_hidden_file(&name), - mode_str, - mode, - mime_type: Some(mime_type), - is_symlink, - }) - }) - .collect(); - - // 排序:目录在前,然后按名称排序 - entries.sort_by(|a, b| match (a.is_dir, b.is_dir) { - (true, false) => std::cmp::Ordering::Less, - (false, true) => std::cmp::Ordering::Greater, - _ => a.name.to_lowercase().cmp(&b.name.to_lowercase()), - }); - - entries - } - Err(e) => { - error!("无法读取目录 {}: {}", canonical_path.display(), e); - return DirectoryListing { - path: canonical_path.to_string_lossy().to_string(), - parent_path, - entries: vec![], - error: Some(format!("无法读取目录: {e}")), - }; - } - }; - - debug!( - "列出目录 {}: {} 个条目", - canonical_path.display(), - entries.len() - ); - - DirectoryListing { - path: canonical_path.to_string_lossy().to_string(), - parent_path, - entries, - error: None, - } -} - -/// 读取文件预览 -pub fn read_file_preview(path: &str, max_size: Option) -> FilePreview { - let max_size = max_size.unwrap_or(100 * 1024); // 默认 100KB - let path_buf = PathBuf::from(path); - - let metadata = match fs::metadata(&path_buf) { - Ok(m) => m, - Err(e) => { - return FilePreview { - path: path.to_string(), - content: None, - is_binary: false, - size: 0, - error: Some(format!("无法读取文件元信息: {e}")), - }; - } - }; - - if metadata.is_dir() { - return FilePreview { - path: path.to_string(), - content: None, - is_binary: false, - size: 0, - error: Some("不能预览目录".to_string()), - }; - } - - let size = metadata.len(); - let extension = get_file_extension(&path_buf); - let is_text = is_text_file(extension.as_deref()); - - if !is_text { - return FilePreview { - path: path.to_string(), - content: None, - is_binary: true, - size, - error: None, - }; - } - - // 读取文件内容 - let content = match fs::read(&path_buf) { - Ok(bytes) => { - let bytes_to_read = bytes.len().min(max_size); - match String::from_utf8(bytes[..bytes_to_read].to_vec()) { - Ok(s) => Some(s), - Err(_) => { - return FilePreview { - path: path.to_string(), - content: None, - is_binary: true, - size, - error: None, - }; - } - } - } - Err(e) => { - return FilePreview { - path: path.to_string(), - content: None, - is_binary: false, - size, - error: Some(format!("无法读取文件: {e}")), - }; - } - }; - - FilePreview { - path: path.to_string(), - content, - is_binary: false, - size, - error: None, - } -} +pub use proxycast_services::file_browser_service::{list_directory, read_file_preview}; +pub use proxycast_services::file_browser_service::{DirectoryListing, FileEntry, FilePreview}; /// Tauri 命令:列出目录 #[tauri::command] pub async fn list_dir(path: String) -> Result { - Ok(list_directory(&path)) + proxycast_services::file_browser_service::list_dir(path).await } /// Tauri 命令:读取文件预览 @@ -529,212 +18,53 @@ pub async fn read_file_preview_cmd( path: String, max_size: Option, ) -> Result { - Ok(read_file_preview(&path, max_size)) + proxycast_services::file_browser_service::read_file_preview_cmd(path, max_size).await } /// Tauri 命令:获取用户主目录 #[tauri::command] pub async fn get_home_dir() -> Result { - dirs::home_dir() - .map(|p| p.to_string_lossy().to_string()) - .ok_or_else(|| "无法获取主目录".to_string()) + proxycast_services::file_browser_service::get_home_dir().await } /// Tauri 命令:创建新文件 #[tauri::command] pub async fn create_file(path: String) -> Result<(), String> { - let path_buf = PathBuf::from(&path); - - // 检查文件是否已存在 - if path_buf.exists() { - return Err("文件已存在".to_string()); - } - - // 确保父目录存在 - if let Some(parent) = path_buf.parent() { - if !parent.exists() { - fs::create_dir_all(parent).map_err(|e| format!("无法创建父目录: {e}"))?; - } - } - - // 创建空文件 - fs::File::create(&path_buf).map_err(|e| format!("无法创建文件: {e}"))?; - - debug!("创建文件: {}", path); - Ok(()) + proxycast_services::file_browser_service::create_file(path).await } /// Tauri 命令:创建新目录 #[tauri::command] pub async fn create_directory(path: String) -> Result<(), String> { - let path_buf = PathBuf::from(&path); - - // 检查目录是否已存在 - if path_buf.exists() { - return Err("目录已存在".to_string()); - } - - fs::create_dir_all(&path_buf).map_err(|e| format!("无法创建目录: {e}"))?; - - debug!("创建目录: {}", path); - Ok(()) + proxycast_services::file_browser_service::create_directory(path).await } /// Tauri 命令:删除文件或目录 #[tauri::command] pub async fn delete_file(path: String, recursive: bool) -> Result<(), String> { - let path_buf = PathBuf::from(&path); - - if !path_buf.exists() { - return Err("文件或目录不存在".to_string()); - } - - if path_buf.is_dir() { - if recursive { - fs::remove_dir_all(&path_buf).map_err(|e| format!("无法删除目录: {e}"))?; - } else { - fs::remove_dir(&path_buf) - .map_err(|e| format!("无法删除目录(目录非空,需要递归删除): {e}"))?; - } - debug!("删除目录: {}", path); - } else { - fs::remove_file(&path_buf).map_err(|e| format!("无法删除文件: {e}"))?; - debug!("删除文件: {}", path); - } - - Ok(()) + proxycast_services::file_browser_service::delete_file(path, recursive).await } /// Tauri 命令:重命名文件或目录 #[tauri::command] pub async fn rename_file(old_path: String, new_path: String) -> Result<(), String> { - let old_path_buf = PathBuf::from(&old_path); - let new_path_buf = PathBuf::from(&new_path); - - if !old_path_buf.exists() { - return Err("源文件或目录不存在".to_string()); - } - - if new_path_buf.exists() { - return Err("目标文件或目录已存在".to_string()); - } - - fs::rename(&old_path_buf, &new_path_buf).map_err(|e| format!("无法重命名: {e}"))?; - - debug!("重命名: {} -> {}", old_path, new_path); - Ok(()) + proxycast_services::file_browser_service::rename_file(old_path, new_path).await } /// Tauri 命令:复制文件名到剪贴板(返回文件名供前端处理) #[tauri::command] pub async fn get_file_name(path: String) -> Result { - let path_buf = PathBuf::from(&path); - path_buf - .file_name() - .and_then(|n| n.to_str()) - .map(|s| s.to_string()) - .ok_or_else(|| "无法获取文件名".to_string()) + proxycast_services::file_browser_service::get_file_name(path).await } /// Tauri 命令:在 Finder 中显示文件 #[tauri::command] pub async fn reveal_in_finder(path: String) -> Result<(), String> { - let path_buf = PathBuf::from(&path); - - if !path_buf.exists() { - return Err("文件或目录不存在".to_string()); - } - - #[cfg(target_os = "macos")] - { - std::process::Command::new("open") - .args(["-R", &path]) - .spawn() - .map_err(|e| format!("无法打开 Finder: {e}"))?; - } - - #[cfg(target_os = "windows")] - { - std::process::Command::new("explorer") - .args(["/select,", &path]) - .spawn() - .map_err(|e| format!("无法打开资源管理器: {}", e))?; - } - - #[cfg(target_os = "linux")] - { - // 尝试使用 xdg-open 打开父目录 - let parent = path_buf.parent().unwrap_or(&path_buf); - std::process::Command::new("xdg-open") - .arg(parent) - .spawn() - .map_err(|e| format!("无法打开文件管理器: {}", e))?; - } - - Ok(()) + proxycast_services::file_browser_service::reveal_in_finder(path).await } /// Tauri 命令:使用默认应用打开文件 #[tauri::command] pub async fn open_with_default_app(path: String) -> Result<(), String> { - let path_buf = PathBuf::from(&path); - - if !path_buf.exists() { - return Err("文件不存在".to_string()); - } - - #[cfg(target_os = "macos")] - { - std::process::Command::new("open") - .arg(&path) - .spawn() - .map_err(|e| format!("无法打开文件: {e}"))?; - } - - #[cfg(target_os = "windows")] - { - std::process::Command::new("cmd") - .args(["/C", "start", "", &path]) - .spawn() - .map_err(|e| format!("无法打开文件: {}", e))?; - } - - #[cfg(target_os = "linux")] - { - std::process::Command::new("xdg-open") - .arg(&path) - .spawn() - .map_err(|e| format!("无法打开文件: {}", e))?; - } - - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_list_home_directory() { - let result = list_directory("~"); - assert!(result.error.is_none()); - assert!(!result.entries.is_empty()); - } - - #[test] - fn test_is_hidden_file() { - assert!(is_hidden_file(".gitignore")); - assert!(is_hidden_file(".config")); - assert!(!is_hidden_file("readme.md")); - } - - #[test] - fn test_is_text_file() { - assert!(is_text_file(Some("txt"))); - assert!(is_text_file(Some("rs"))); - assert!(is_text_file(Some("json"))); - assert!(!is_text_file(Some("png"))); - assert!(!is_text_file(Some("exe"))); - assert!(!is_text_file(None)); - } + proxycast_services::file_browser_service::open_with_default_app(path).await } diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index aa8e9b8c1..22a8f911b 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -1,38 +1,10 @@ //! 业务服务模块 //! //! 核心业务逻辑已迁移到 proxycast-services crate。 -//! 本模块保留 Tauri 相关服务和重新导出。 - -// 从 proxycast-services crate 重新导出 -pub use proxycast_services::api_key_provider_service; -pub use proxycast_services::aster_session_store; -pub use proxycast_services::backup_service; -pub use proxycast_services::content_creator; -pub use proxycast_services::context_memory_service; -pub use proxycast_services::general_chat; -pub use proxycast_services::kiro_event_service; -pub use proxycast_services::live_sync; -pub use proxycast_services::machine_id_service; -pub use proxycast_services::material_service; -pub use proxycast_services::mcp_service; -pub use proxycast_services::mcp_sync; -pub use proxycast_services::model_registry_service; -pub use proxycast_services::model_service; -pub use proxycast_services::persona_service; -pub use proxycast_services::project_context_builder; -pub use proxycast_services::prompt_service; -pub use proxycast_services::prompt_sync; -pub use proxycast_services::provider_pool_service; -pub use proxycast_services::session_context_service; -pub use proxycast_services::skill_service; -pub use proxycast_services::switch; -pub use proxycast_services::template_service; -pub use proxycast_services::token_cache_service; -pub use proxycast_services::tool_hooks_service; -pub use proxycast_services::update_check_service; -pub use proxycast_services::usage_service; +//! 本模块保留 Tauri 相关服务。 // 保留在主 crate 的 Tauri 相关服务 pub mod file_browser_service; pub mod sysinfo_service; +pub mod update_check_service; pub mod update_window; diff --git a/src-tauri/src/services/sysinfo_service.rs b/src-tauri/src/services/sysinfo_service.rs index ff753322e..b1d16e921 100644 --- a/src-tauri/src/services/sysinfo_service.rs +++ b/src-tauri/src/services/sysinfo_service.rs @@ -1,201 +1,35 @@ -//! 系统信息服务 +//! 系统信息服务(Tauri 命令桥接层) //! -//! 提供 CPU、内存等系统信息的获取和订阅功能 -//! 用于系统监控视图的数据源 -//! -//! # 功能 -//! - 获取当前系统信息快照 -//! - 订阅系统信息更新(每秒推送) -//! -//! # 依赖 -//! - sysinfo crate +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块仅保留 Tauri 命令封装。 -use serde::{Deserialize, Serialize}; -use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; -use sysinfo::{CpuRefreshKind, MemoryRefreshKind, RefreshKind, System}; + use tauri::{AppHandle, Emitter}; -use tokio::sync::Mutex; -use tokio::time::{interval, Duration}; -use tracing::{debug, error, info}; -/// 系统信息数据点 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SysinfoData { - /// 时间戳(毫秒) - pub ts: u64, - /// 总 CPU 使用率(0-100) - pub cpu: f32, - /// 已用内存(GB) - #[serde(rename = "mem:used")] - pub mem_used: f64, - /// 总内存(GB) - #[serde(rename = "mem:total")] - pub mem_total: f64, - /// 各核心 CPU 使用率 - #[serde(flatten)] - pub cpu_cores: std::collections::HashMap, -} - -/// 系统信息服务状态 -pub struct SysinfoService { - /// 是否正在运行订阅 - running: AtomicBool, - /// 系统信息实例 - system: Mutex, -} - -impl SysinfoService { - /// 创建新的系统信息服务 - pub fn new() -> Self { - let system = System::new_with_specifics( - RefreshKind::new() - .with_cpu(CpuRefreshKind::everything()) - .with_memory(MemoryRefreshKind::everything()), - ); - Self { - running: AtomicBool::new(false), - system: Mutex::new(system), - } - } - - /// 获取当前系统信息快照 - pub async fn get_sysinfo(&self) -> SysinfoData { - let mut system = self.system.lock().await; - - // 刷新 CPU 和内存信息 - system.refresh_cpu_usage(); - system.refresh_memory(); - - let ts = std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_millis() as u64; - - // 计算总 CPU 使用率 - let cpu = system.global_cpu_usage(); - - // 内存信息(转换为 GB) - let mem_total = system.total_memory() as f64 / 1024.0 / 1024.0 / 1024.0; - let mem_used = system.used_memory() as f64 / 1024.0 / 1024.0 / 1024.0; - - // 各核心 CPU 使用率 - let mut cpu_cores = std::collections::HashMap::new(); - for (i, cpu) in system.cpus().iter().enumerate() { - cpu_cores.insert(format!("cpu:{i}"), cpu.cpu_usage()); - } - - SysinfoData { - ts, - cpu, - mem_used, - mem_total, - cpu_cores, - } - } - - /// 检查是否正在运行 - pub fn is_running(&self) -> bool { - self.running.load(Ordering::SeqCst) - } - - /// 停止订阅 - pub fn stop(&self) { - self.running.store(false, Ordering::SeqCst); - info!("系统信息订阅已停止"); - } -} - -impl Default for SysinfoService { - fn default() -> Self { - Self::new() - } -} - -/// 全局系统信息服务实例 -static SYSINFO_SERVICE: once_cell::sync::Lazy> = - once_cell::sync::Lazy::new(|| Arc::new(SysinfoService::new())); - -/// 获取全局系统信息服务 -pub fn get_sysinfo_service() -> Arc { - SYSINFO_SERVICE.clone() -} +pub use proxycast_services::sysinfo_service::{get_sysinfo_service, SysinfoEmitter}; +pub use proxycast_services::sysinfo_service::{SysinfoData, SysinfoService}; /// Tauri 命令:获取当前系统信息 #[tauri::command] pub async fn get_sysinfo() -> Result { - let service = get_sysinfo_service(); - Ok(service.get_sysinfo().await) + proxycast_services::sysinfo_service::get_sysinfo().await } /// Tauri 命令:开始订阅系统信息 /// 每秒向前端发送 sysinfo 事件 #[tauri::command] pub async fn subscribe_sysinfo(app: AppHandle) -> Result<(), String> { - let service = get_sysinfo_service(); - - // 如果已经在运行,直接返回 - if service.is_running() { - debug!("系统信息订阅已在运行"); - return Ok(()); - } - - service.running.store(true, Ordering::SeqCst); - info!("开始系统信息订阅"); - - // 启动后台任务 - let service_clone = service.clone(); - tokio::spawn(async move { - let mut ticker = interval(Duration::from_secs(1)); - - // 首次需要等待一下让 CPU 使用率计算准确 - ticker.tick().await; - - while service_clone.is_running() { - ticker.tick().await; - - if !service_clone.is_running() { - break; - } - - let data = service_clone.get_sysinfo().await; - - if let Err(e) = app.emit("sysinfo", &data) { - error!("发送系统信息事件失败: {}", e); - } - } - - info!("系统信息订阅任务结束"); + let emitter: SysinfoEmitter = Arc::new(move |data: &SysinfoData| { + app.emit("sysinfo", data) + .map_err(|e| format!("发送系统信息事件失败: {e}")) }); - Ok(()) + proxycast_services::sysinfo_service::subscribe_sysinfo(emitter).await } /// Tauri 命令:停止订阅系统信息 #[tauri::command] pub async fn unsubscribe_sysinfo() -> Result<(), String> { - let service = get_sysinfo_service(); - service.stop(); - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[tokio::test] - async fn test_get_sysinfo() { - let service = SysinfoService::new(); - - // 等待一下让 CPU 数据准确 - tokio::time::sleep(Duration::from_millis(100)).await; - - let data = service.get_sysinfo().await; - - assert!(data.ts > 0); - assert!(data.cpu >= 0.0 && data.cpu <= 100.0); - assert!(data.mem_total > 0.0); - assert!(data.mem_used >= 0.0); - assert!(data.mem_used <= data.mem_total); - } + proxycast_services::sysinfo_service::unsubscribe_sysinfo().await } diff --git a/src-tauri/src/services/update_check_service.rs b/src-tauri/src/services/update_check_service.rs new file mode 100644 index 000000000..a6261457f --- /dev/null +++ b/src-tauri/src/services/update_check_service.rs @@ -0,0 +1,8 @@ +//! 更新检查服务(桥接层) +//! +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块仅保留兼容导出。 + +pub use proxycast_services::update_check_service::{ + UpdateCheckService, UpdateCheckServiceState, UpdateCheckState, UpdateInfo, +}; diff --git a/src-tauri/src/session/mod.rs b/src-tauri/src/session/mod.rs deleted file mode 100644 index 192e77d0d..000000000 --- a/src-tauri/src/session/mod.rs +++ /dev/null @@ -1,21 +0,0 @@ -//! 会话管理模块 -//! -//! 提供以下功能: -//! - 稳定的 SessionId 生成(基于请求内容哈希) -//! - thoughtSignature 全局缓存 -//! - 会话粘性管理(会话与账号映射) -//! - 调度模式配置 -//! - 增强的限流处理(Duration 解析、指数退避) - -// 从 providers crate 重新导出 session_manager 和 signature_store -pub use proxycast_providers::session::SessionManager; -pub use proxycast_providers::session::{ - clear_thought_signature, get_thought_signature, has_valid_signature, store_thought_signature, - take_thought_signature, -}; - -// 从 core crate 重新导出 rate_limit、sticky_config、sticky_manager -pub use proxycast_core::session::{ - extract_retry_delay, parse_duration_string, AccountInfo, RateLimitReason, RateLimitRecord, - RateLimitTracker, SchedulingMode, StickySessionConfig, StickySessionManager, -}; diff --git a/src-tauri/src/skills/README.md b/src-tauri/src/skills/README.md index 3dea5ad05..bbdf70d4b 100644 --- a/src-tauri/src/skills/README.md +++ b/src-tauri/src/skills/README.md @@ -7,8 +7,8 @@ | 文件 | 说明 | |------|------| | `mod.rs` | 模块导出 | -| `llm_provider.rs` | ProxyCastLlmProvider 实现 | -| `execution_callback.rs` | TauriExecutionCallback 实现 | +| `llm_provider.rs` | 桥接层(纯逻辑已迁移到 `crates/skills/src/proxycast_llm_provider.rs`) | +| `execution_callback.rs` | TauriExecutionCallback 实现(保留在主 crate) | ## Skills 集成架构 @@ -77,9 +77,10 @@ agent/aster_state.rs └── reload_proxycast_skills() skills/ -├── llm_provider.rs -│ ├── ProviderPoolService (凭证池管理) -│ └── ApiKeyProviderService (API Key 服务) +├── llm_provider.rs (桥接) +│ └── crates/skills/src/proxycast_llm_provider.rs +│ ├── ProviderPoolService (凭证池管理) +│ └── ApiKeyProviderService (API Key 服务) └── execution_callback.rs └── tauri::AppHandle (事件发送) diff --git a/src-tauri/src/skills/llm_provider.rs b/src-tauri/src/skills/llm_provider.rs index 36bd86106..1b0b71ef0 100644 --- a/src-tauri/src/skills/llm_provider.rs +++ b/src-tauri/src/skills/llm_provider.rs @@ -1,541 +1,6 @@ -//! ProxyCast LLM Provider 实现 +//! ProxyCast LLM Provider(桥接层) //! -//! 使用 ProviderPoolService 选择凭证并调用 LLM API。 -//! trait 定义(LlmProvider, SkillError)已迁移到 proxycast-skills crate。 +//! 纯逻辑已迁移到 `proxycast-skills` crate, +//! 本模块保留兼容导出。 -use std::sync::Arc; - -use async_trait::async_trait; - -use crate::database::DbConnection; -use crate::models::anthropic::AnthropicMessagesRequest; -#[cfg(test)] -use crate::models::provider_pool_model::PoolProviderType; -use crate::models::provider_pool_model::{CredentialData, ProviderCredential}; -use crate::providers::{ClaudeCustomProvider, KiroProvider, OpenAICustomProvider}; -use crate::services::api_key_provider_service::ApiKeyProviderService; -use crate::services::provider_pool_service::ProviderPoolService; - -use proxycast_skills::{LlmProvider, SkillError}; - -/// ProxyCast LLM Provider -/// -/// 使用 ProviderPoolService 选择凭证并调用 LLM API。 -/// 实现 aster-rust 定义的 LlmProvider trait。 -pub struct ProxyCastLlmProvider { - /// 凭证池服务 - pool_service: Arc, - /// API Key Provider 服务(用于智能降级) - api_key_service: Arc, - /// 数据库连接 - db: DbConnection, - /// 偏好的 Provider 类型(可选) - preferred_provider: Option, -} - -impl ProxyCastLlmProvider { - /// 创建新的 ProxyCastLlmProvider 实例 - /// - /// # Arguments - /// * `pool_service` - 凭证池服务 - /// * `api_key_service` - API Key 服务 - /// * `db` - 数据库连接 - pub fn new( - pool_service: Arc, - api_key_service: Arc, - db: DbConnection, - ) -> Self { - Self { - pool_service, - api_key_service, - db, - preferred_provider: None, - } - } - - /// 创建带有偏好 Provider 的实例 - /// - /// # Arguments - /// * `pool_service` - 凭证池服务 - /// * `api_key_service` - API Key 服务 - /// * `db` - 数据库连接 - /// * `preferred_provider` - 偏好的 Provider 类型 - pub fn with_preferred_provider( - pool_service: Arc, - api_key_service: Arc, - db: DbConnection, - preferred_provider: String, - ) -> Self { - Self { - pool_service, - api_key_service, - db, - preferred_provider: Some(preferred_provider), - } - } - - /// 设置偏好的 Provider 类型 - pub fn set_preferred_provider(&mut self, provider: Option) { - self.preferred_provider = provider; - } - - /// 获取偏好的 Provider 类型 - pub fn preferred_provider(&self) -> Option<&str> { - self.preferred_provider.as_deref() - } - - /// 将 Skill 的 provider 字段映射到 PoolProviderType - /// - /// # Arguments - /// * `provider` - Provider 名称字符串 - /// - /// # Returns - /// 对应的 PoolProviderType,未知类型返回 None - #[cfg(test)] - fn map_skill_provider_to_pool_type(provider: &str) -> Option { - match provider.to_lowercase().as_str() { - "openai" | "gpt" => Some(PoolProviderType::OpenAI), - "anthropic" | "claude" => Some(PoolProviderType::Claude), - "gemini" | "google" => Some(PoolProviderType::Gemini), - "kiro" | "codewhisperer" => Some(PoolProviderType::Kiro), - "vertex" => Some(PoolProviderType::Vertex), - "codex" => Some(PoolProviderType::Codex), - _ => None, - } - } - - /// 根据凭证调用 LLM API - /// - /// # Arguments - /// * `credential` - 选中的凭证 - /// * `system_prompt` - 系统提示词 - /// * `user_message` - 用户消息 - /// * `model` - 模型名称 - /// - /// # Returns - /// LLM 响应文本或错误 - async fn call_llm_with_credential( - &self, - credential: &ProviderCredential, - system_prompt: &str, - user_message: &str, - model: &str, - ) -> Result { - match &credential.credential { - CredentialData::KiroOAuth { creds_file_path } => { - self.call_kiro_api(creds_file_path, system_prompt, user_message, model) - .await - } - CredentialData::ClaudeKey { api_key, base_url } => { - self.call_claude_api( - api_key, - base_url.as_deref(), - system_prompt, - user_message, - model, - ) - .await - } - CredentialData::OpenAIKey { api_key, base_url } => { - self.call_openai_api( - api_key, - base_url.as_deref(), - system_prompt, - user_message, - model, - ) - .await - } - CredentialData::AnthropicKey { api_key, base_url } => { - // Anthropic API Key 使用 Claude API - self.call_claude_api( - api_key, - base_url.as_deref(), - system_prompt, - user_message, - model, - ) - .await - } - _ => Err(SkillError::ProviderError(format!( - "不支持的凭证类型: {:?}", - credential.provider_type - ))), - } - } - - /// 调用 Kiro API - async fn call_kiro_api( - &self, - creds_file_path: &str, - system_prompt: &str, - user_message: &str, - model: &str, - ) -> Result { - use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; - use crate::models::anthropic::AnthropicMessage; - use crate::providers::traits::CredentialProvider; - use crate::server_utils::parse_cw_response; - - let mut kiro = KiroProvider::new(); - kiro.load_credentials_from_path(creds_file_path) - .await - .map_err(|e| SkillError::ProviderError(format!("加载 Kiro 凭证失败: {}", e)))?; - - // 确保 Token 有效 - if !kiro.is_token_valid() || kiro.is_token_expiring_soon() { - kiro.refresh_token() - .await - .map_err(|e| SkillError::ProviderError(format!("刷新 Token 失败: {}", e)))?; - } - - // 构建 Anthropic 请求 - let request = AnthropicMessagesRequest { - model: model.to_string(), - max_tokens: Some(4096), - system: Some(serde_json::Value::String(system_prompt.to_string())), - messages: vec![AnthropicMessage { - role: "user".to_string(), - content: serde_json::Value::String(user_message.to_string()), - }], - stream: false, - temperature: None, - tools: None, - tool_choice: None, - }; - - // 转换为 OpenAI 格式并调用 - let openai_request = convert_anthropic_to_openai(&request); - let resp = kiro - .call_api(&openai_request) - .await - .map_err(|e| SkillError::ProviderError(format!("Kiro API 调用失败: {}", e)))?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(SkillError::ProviderError(format!( - "Kiro API 返回错误: status={}, body={}", - status, body - ))); - } - - let bytes = resp - .bytes() - .await - .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; - let body = String::from_utf8_lossy(&bytes).to_string(); - let parsed = parse_cw_response(&body); - - Ok(parsed.content) - } - - /// 调用 Claude API - async fn call_claude_api( - &self, - api_key: &str, - base_url: Option<&str>, - system_prompt: &str, - user_message: &str, - model: &str, - ) -> Result { - use crate::models::anthropic::AnthropicMessage; - - let claude = - ClaudeCustomProvider::with_config(api_key.to_string(), base_url.map(|s| s.to_string())); - - // 构建 Anthropic 请求 - let request = AnthropicMessagesRequest { - model: model.to_string(), - max_tokens: Some(4096), - system: Some(serde_json::Value::String(system_prompt.to_string())), - messages: vec![AnthropicMessage { - role: "user".to_string(), - content: serde_json::Value::String(user_message.to_string()), - }], - stream: false, - temperature: None, - tools: None, - tool_choice: None, - }; - - let resp = claude - .call_api(&request) - .await - .map_err(|e| SkillError::ProviderError(format!("Claude API 调用失败: {}", e)))?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(SkillError::ProviderError(format!( - "Claude API 返回错误: status={}, body={}", - status, body - ))); - } - - let body = resp - .text() - .await - .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; - - // 解析 Anthropic 响应 - let json: serde_json::Value = serde_json::from_str(&body) - .map_err(|e| SkillError::ProviderError(format!("解析响应失败: {}", e)))?; - - // 提取文本内容 - let content = json["content"] - .as_array() - .and_then(|arr| arr.first()) - .and_then(|block| block["text"].as_str()) - .unwrap_or(""); - - Ok(content.to_string()) - } - - /// 调用 OpenAI API - async fn call_openai_api( - &self, - api_key: &str, - base_url: Option<&str>, - system_prompt: &str, - user_message: &str, - model: &str, - ) -> Result { - use crate::models::openai::{ChatCompletionRequest, ChatMessage, MessageContent}; - - let openai = - OpenAICustomProvider::with_config(api_key.to_string(), base_url.map(|s| s.to_string())); - - // 构建 OpenAI 请求 - let request = ChatCompletionRequest { - model: model.to_string(), - messages: vec![ - ChatMessage { - role: "system".to_string(), - content: Some(MessageContent::Text(system_prompt.to_string())), - tool_calls: None, - tool_call_id: None, - reasoning_content: None, - }, - ChatMessage { - role: "user".to_string(), - content: Some(MessageContent::Text(user_message.to_string())), - tool_calls: None, - tool_call_id: None, - reasoning_content: None, - }, - ], - max_tokens: Some(4096), - stream: false, - temperature: None, - top_p: None, - tools: None, - tool_choice: None, - reasoning_effort: None, - }; - - let resp = openai - .call_api(&request) - .await - .map_err(|e| SkillError::ProviderError(format!("OpenAI API 调用失败: {}", e)))?; - - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - return Err(SkillError::ProviderError(format!( - "OpenAI API 返回错误: status={}, body={}", - status, body - ))); - } - - let body = resp - .text() - .await - .map_err(|e| SkillError::ProviderError(format!("读取响应失败: {}", e)))?; - - // 解析 OpenAI 响应 - let json: serde_json::Value = serde_json::from_str(&body) - .map_err(|e| SkillError::ProviderError(format!("解析响应失败: {}", e)))?; - - // 提取文本内容 - let content = json["choices"] - .as_array() - .and_then(|arr| arr.first()) - .and_then(|choice| choice["message"]["content"].as_str()) - .unwrap_or(""); - - Ok(content.to_string()) - } -} - -#[async_trait] -impl LlmProvider for ProxyCastLlmProvider { - /// 调用 LLM 进行对话 - /// - /// # 实现说明 - /// 1. 使用 ProviderPoolService.select_credential_with_fallback() 选择凭证 - /// 2. 如果指定了 preferred_provider,优先选择该类型的凭证 - /// 3. 如果指定了 model,传递给底层 provider - /// 4. 如果没有可用凭证,返回 ProviderError - /// - /// # Requirements - /// - 1.2: 使用 ProviderPoolService 选择可用凭证 - /// - 1.3: 优先选择指定 provider 类型的凭证 - /// - 1.4: 将 model 参数传递给底层 provider - /// - 1.5: 没有可用凭证时返回 ProviderError - async fn chat( - &self, - system_prompt: &str, - user_message: &str, - model: Option<&str>, - ) -> Result { - // 确定要使用的 provider 类型 - let provider_type = self.preferred_provider.as_deref().unwrap_or("claude"); // 默认使用 Claude - - // 确定要使用的模型 - let model_name = model.unwrap_or("claude-sonnet-4-5-20250514"); - - tracing::info!( - "[ProxyCastLlmProvider] chat 调用: provider_type={}, model={}", - provider_type, - model_name - ); - - // 使用 ProviderPoolService 选择凭证(Requirements 1.2, 1.3) - let credential = self - .pool_service - .select_credential_with_fallback( - &self.db, - &self.api_key_service, - provider_type, - Some(model_name), - None, // provider_id_hint - None, // client_type - ) - .await - .map_err(|e| SkillError::ProviderError(format!("选择凭证失败: {}", e)))? - .ok_or_else(|| { - // Requirements 1.5: 没有可用凭证时返回 ProviderError - SkillError::ProviderError(format!( - "没有可用的凭证: provider_type={}, model={}", - provider_type, model_name - )) - })?; - - tracing::info!( - "[ProxyCastLlmProvider] 选中凭证: uuid={}, type={:?}", - &credential.uuid[..8], - credential.provider_type - ); - - // 调用 LLM API(Requirements 1.4: 传递 model 参数) - let result = self - .call_llm_with_credential(&credential, system_prompt, user_message, model_name) - .await; - - // 记录使用情况 - match &result { - Ok(_) => { - let _ = self.pool_service.record_usage(&self.db, &credential.uuid); - let _ = - self.pool_service - .mark_healthy(&self.db, &credential.uuid, Some(model_name)); - } - Err(e) => { - let _ = self.pool_service.mark_unhealthy( - &self.db, - &credential.uuid, - Some(&e.to_string()), - ); - } - } - - result - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_map_skill_provider_openai() { - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("openai"), - Some(PoolProviderType::OpenAI) - ); - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("gpt"), - Some(PoolProviderType::OpenAI) - ); - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("OPENAI"), - Some(PoolProviderType::OpenAI) - ); - } - - #[test] - fn test_map_skill_provider_claude() { - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("claude"), - Some(PoolProviderType::Claude) - ); - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("anthropic"), - Some(PoolProviderType::Claude) - ); - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("CLAUDE"), - Some(PoolProviderType::Claude) - ); - } - - #[test] - fn test_map_skill_provider_gemini() { - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("gemini"), - Some(PoolProviderType::Gemini) - ); - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("google"), - Some(PoolProviderType::Gemini) - ); - } - - #[test] - fn test_map_skill_provider_kiro() { - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("kiro"), - Some(PoolProviderType::Kiro) - ); - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("codewhisperer"), - Some(PoolProviderType::Kiro) - ); - } - - #[test] - fn test_map_skill_provider_unknown() { - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type("unknown_provider"), - None - ); - assert_eq!( - ProxyCastLlmProvider::map_skill_provider_to_pool_type(""), - None - ); - } - - #[test] - fn test_skill_error_display() { - let provider_err = SkillError::ProviderError("没有可用凭证".to_string()); - assert!(provider_err.to_string().contains("Provider error")); - assert!(provider_err.to_string().contains("没有可用凭证")); - - let exec_err = SkillError::ExecutionError("执行失败".to_string()); - assert!(exec_err.to_string().contains("Execution error")); - - let config_err = SkillError::ConfigError("配置错误".to_string()); - assert!(config_err.to_string().contains("Config error")); - } -} +pub use proxycast_skills::ProxyCastLlmProvider; diff --git a/src-tauri/src/skills/mod.rs b/src-tauri/src/skills/mod.rs index dc6c6d6f4..278340d65 100644 --- a/src-tauri/src/skills/mod.rs +++ b/src-tauri/src/skills/mod.rs @@ -1,21 +1,13 @@ //! Skills 集成模块 //! -//! trait 定义和纯逻辑已迁移到 proxycast-skills crate, -//! 本模块保留 Tauri 相关的实现。 +//! 纯逻辑已迁移到 `proxycast-skills` crate, +//! 本模块保留 Tauri 相关实现和兼容导出层。 mod execution_callback; mod llm_provider; -// 从 proxycast-skills crate re-export -pub use proxycast_skills::{ - events, ExecutionCallback, ExecutionCompletePayload, LlmProvider, SkillError, - StepCompletePayload, StepErrorPayload, StepStartPayload, -}; -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; + +// 兼容导出(实际实现位于 proxycast-skills crate) pub use llm_provider::ProxyCastLlmProvider; diff --git a/src-tauri/src/terminal/mod.rs b/src-tauri/src/terminal/mod.rs index eb9a4f4c9..185efd493 100644 --- a/src-tauri/src/terminal/mod.rs +++ b/src-tauri/src/terminal/mod.rs @@ -1,4 +1,4 @@ -//! 终端核心模块(重导出层) +//! 终端模块 //! //! 实际实现位于 `proxycast-terminal` crate。 //! 本模块提供 `TauriEmitter` newtype 桥接 Tauri 与终端 crate。 @@ -29,26 +29,3 @@ impl TerminalEventEmit for TauriEmitter { .map_err(|e| format!("获取应用数据目录失败: {e}")) } } - -// 重新导出 proxycast-terminal 的所有公共类型 -pub use proxycast_terminal::block_controller; -pub use proxycast_terminal::connections; -pub use proxycast_terminal::emit_helper; -pub use proxycast_terminal::emitter; -pub use proxycast_terminal::error; -pub use proxycast_terminal::events; -pub use proxycast_terminal::integration; -pub use proxycast_terminal::persistence; -pub use proxycast_terminal::pty_session; -pub use proxycast_terminal::session_manager; - -// 重新导出常用类型 -pub use proxycast_terminal::{ - resync_controller, BlockController, BlockControllerRuntimeStatus, BlockFile, BlockInputUnion, - BlockMeta, ControllerRegistry, ControllerStatusEvent, DynEmitter, NoOpEmitter, PtySession, - ResyncController, ResyncOptions, ResyncResult, RuntimeOpts, SessionMetadata, - SessionMetadataStore, SessionRecord, SessionStatus, ShellController, ShellProc, TermSize, - TerminalError, TerminalEventEmitter, TerminalOutputEvent, TerminalSessionManager, - TerminalStatusEvent, CONTROLLER_STATUS_EVENT, DEFAULT_COLS, DEFAULT_ROWS, - TERMINAL_RESET_SEQUENCE, TERMINAL_SOFT_RESET_SEQUENCE, -}; diff --git a/src-tauri/src/credential/tests.rs b/src-tauri/src/tests/credential_tests.rs similarity index 99% rename from src-tauri/src/credential/tests.rs rename to src-tauri/src/tests/credential_tests.rs index 0d8bfadca..fc845a1b5 100644 --- a/src-tauri/src/credential/tests.rs +++ b/src-tauri/src/tests/credential_tests.rs @@ -4,11 +4,10 @@ #![allow(dead_code)] -use crate::credential::{ - BalanceStrategy, Credential, CredentialData, CredentialPool, LoadBalancer, -}; use crate::ProviderType; use proptest::prelude::*; +use proxycast_core::credential::{Credential, CredentialData, CredentialPool}; +use proxycast_credential::{BalanceStrategy, LoadBalancer}; use std::collections::HashSet; use std::sync::Arc; @@ -357,7 +356,7 @@ proptest! { provider in arb_provider_type(), failure_threshold in 1u32..=5u32 ) { - use crate::credential::{CredentialStatus, HealthCheckConfig, HealthChecker}; + use proxycast_core::credential::{CredentialStatus, HealthCheckConfig, HealthChecker}; use std::time::Duration; // 创建带自定义阈值的健康检查器 @@ -430,7 +429,7 @@ proptest! { provider in arb_provider_type(), latency_ms in 1u64..1000u64 ) { - use crate::credential::{CredentialStatus, HealthChecker}; + use proxycast_core::credential::{CredentialStatus, HealthChecker}; let checker = HealthChecker::with_defaults(); let pool = CredentialPool::new(provider); @@ -488,7 +487,7 @@ proptest! { failures_before in 1u32..3u32, latency_ms in 1u64..1000u64 ) { - use crate::credential::HealthChecker; + use proxycast_core::credential::HealthChecker; let checker = HealthChecker::with_defaults(); let pool = CredentialPool::new(provider); @@ -541,7 +540,7 @@ proptest! { cooldown_index in 0usize..5usize ) { use chrono::{Duration, Utc}; - use crate::credential::CredentialStatus; + use proxycast_core::credential::CredentialStatus; let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); let pool = Arc::new(CredentialPool::new(provider)); @@ -623,7 +622,7 @@ proptest! { cred_count in 1usize..=5usize ) { use chrono::Duration; - use crate::credential::PoolError; + use proxycast_core::credential::PoolError; let lb = LoadBalancer::new(BalanceStrategy::RoundRobin); let pool = Arc::new(CredentialPool::new(provider)); @@ -668,11 +667,11 @@ proptest! { // ============ 凭证同步服务属性测试 ============ -use crate::config::{Config, ConfigManager}; -use crate::credential::CredentialSyncService; use crate::models::provider_pool_model::{ CredentialData as PoolCredentialData, PoolProviderType, ProviderCredential, }; +use proxycast_core::config::{Config, ConfigManager}; +use proxycast_credential::CredentialSyncService; use std::sync::RwLock; use tempfile::TempDir; @@ -1392,8 +1391,8 @@ proptest! { // ============ 配额管理器属性测试 ============ -use crate::config::QuotaExceededConfig; -use crate::credential::QuotaManager; +use proxycast_core::config::QuotaExceededConfig; +use proxycast_credential::QuotaManager; /// 生成随机的配额超限配置 fn arb_quota_config() -> impl Strategy { diff --git a/src-tauri/src/tests/mod.rs b/src-tauri/src/tests/mod.rs new file mode 100644 index 000000000..ed02fac65 --- /dev/null +++ b/src-tauri/src/tests/mod.rs @@ -0,0 +1,2 @@ +mod credential_tests; +mod processor_tests; diff --git a/src-tauri/src/processor/tests.rs b/src-tauri/src/tests/processor_tests.rs similarity index 99% rename from src-tauri/src/processor/tests.rs rename to src-tauri/src/tests/processor_tests.rs index 4b33d86c2..dd5003454 100644 --- a/src-tauri/src/processor/tests.rs +++ b/src-tauri/src/tests/processor_tests.rs @@ -1,8 +1,8 @@ //! 处理器模块测试 -use super::*; -use crate::services::provider_pool_service::ProviderPoolService; use crate::ProviderType; +use proxycast_processor::*; +use proxycast_services::provider_pool_service::ProviderPoolService; use std::sync::Arc; #[test] diff --git a/src-tauri/src/tray/format.rs b/src-tauri/src/tray/format.rs index bf7d27175..b8ef52d43 100644 --- a/src-tauri/src/tray/format.rs +++ b/src-tauri/src/tray/format.rs @@ -1,120 +1,8 @@ -//! 菜单文本格式化模块 +//! 菜单文本格式化模块(桥接层) //! -//! 提供托盘菜单文本的格式化函数 +//! 纯逻辑已迁移到 `proxycast-core` crate, +//! 本模块保留兼容导出。 -/// 格式化服务器状态文本 -/// -/// # 示例输出 -/// - "🟢 API 服务器: 运行中 (127.0.0.1:8080)" -/// - "⚪ API 服务器: 已停止" -pub fn format_server_status(running: bool, host: &str, port: u16) -> String { - if running { - format!("🟢 API 服务器: 运行中 ({host}:{port})") - } else { - "⚪ API 服务器: 已停止".to_string() - } -} - -/// 格式化凭证状态文本 -/// -/// # 示例输出 -/// - "🔑 可用凭证: 3/5" -pub fn format_credential_status(available: usize, total: usize) -> String { - format!("🔑 可用凭证: {available}/{total}") -} - -/// 格式化请求统计文本 -/// -/// # 示例输出 -/// - "📊 今日请求: 128 次" -pub fn format_request_count(count: u64) -> String { - format!("📊 今日请求: {count} 次") -} - -/// 格式化 API 地址 -/// -/// # 示例输出 -/// - "http://127.0.0.1:8080" -pub fn format_api_address(host: &str, port: u16) -> String { - format!("http://{host}:{port}") -} - -#[cfg(test)] -mod tests { - use super::*; - use proptest::prelude::*; - - proptest! { - /// **Feature: system-tray, Property 2: 菜单内容格式化正确性** - /// **Validates: Requirements 2.2, 2.3, 2.4** - #[test] - fn prop_menu_content_formatting( - host in "[a-z0-9.]{1,50}", - port in 1024u16..65535, - available in 0usize..100, - total in 0usize..100, - requests in 0u64..1000000 - ) { - // 测试服务器状态格式化 - 运行中 - let running_status = format_server_status(true, &host, port); - prop_assert!(running_status.contains(&host), "运行状态应包含 host"); - prop_assert!(running_status.contains(&port.to_string()), "运行状态应包含 port"); - prop_assert!(running_status.contains("运行中"), "运行状态应包含'运行中'"); - - // 测试服务器状态格式化 - 已停止 - let stopped_status = format_server_status(false, &host, port); - prop_assert!(stopped_status.contains("已停止"), "停止状态应包含'已停止'"); - - // 测试凭证状态格式化 - let cred_status = format_credential_status(available, total); - prop_assert!(cred_status.contains(&available.to_string()), "凭证状态应包含可用数"); - prop_assert!(cred_status.contains(&total.to_string()), "凭证状态应包含总数"); - - // 测试请求统计格式化 - let req_status = format_request_count(requests); - prop_assert!(req_status.contains(&requests.to_string()), "请求统计应包含请求次数"); - } - - /// **Feature: system-tray, Property 4: API 地址格式化正确性** - /// **Validates: Requirements 4.2** - #[test] - fn prop_api_address_formatting( - host in "[a-z0-9.]{1,50}", - port in 1024u16..65535 - ) { - let address = format_api_address(&host, port); - let expected = format!("http://{host}:{port}"); - prop_assert_eq!(address, expected, "API 地址格式应为 http://{{host}}:{{port}}"); - } - } - - #[test] - fn test_format_server_status_running() { - let status = format_server_status(true, "127.0.0.1", 8080); - assert_eq!(status, "🟢 API 服务器: 运行中 (127.0.0.1:8080)"); - } - - #[test] - fn test_format_server_status_stopped() { - let status = format_server_status(false, "127.0.0.1", 8080); - assert_eq!(status, "⚪ API 服务器: 已停止"); - } - - #[test] - fn test_format_credential_status() { - let status = format_credential_status(3, 5); - assert_eq!(status, "🔑 可用凭证: 3/5"); - } - - #[test] - fn test_format_request_count() { - let status = format_request_count(128); - assert_eq!(status, "📊 今日请求: 128 次"); - } - - #[test] - fn test_format_api_address() { - let address = format_api_address("127.0.0.1", 8080); - assert_eq!(address, "http://127.0.0.1:8080"); - } -} +pub use proxycast_core::tray_format::{ + format_api_address, format_credential_status, format_request_count, format_server_status, +}; diff --git a/src-tauri/src/tray/menu.rs b/src-tauri/src/tray/menu.rs index d4d4827c7..5130bc9f1 100644 --- a/src-tauri/src/tray/menu.rs +++ b/src-tauri/src/tray/menu.rs @@ -9,59 +9,8 @@ use tauri::{ AppHandle, Runtime, }; -/// 菜单项 ID 常量 -pub mod menu_ids { - /// 状态信息 - pub const STATUS_INFO: &str = "status_info"; - /// 凭证信息 - pub const CREDENTIAL_INFO: &str = "credential_info"; - /// 请求信息 - pub const REQUEST_INFO: &str = "request_info"; - /// 分隔符 1 - pub const SEPARATOR_1: &str = "sep_1"; - /// 启动服务器 - pub const START_SERVER: &str = "start_server"; - /// 停止服务器 - pub const STOP_SERVER: &str = "stop_server"; - /// 刷新所有 Token - pub const REFRESH_TOKENS: &str = "refresh_tokens"; - /// 健康检查 - pub const HEALTH_CHECK: &str = "health_check"; - /// 分隔符 2 - pub const SEPARATOR_2: &str = "sep_2"; - /// 打开主窗口 - pub const OPEN_WINDOW: &str = "open_window"; - /// 复制 API 地址 - pub const COPY_API_ADDRESS: &str = "copy_api_address"; - /// 打开日志目录 - pub const OPEN_LOG_DIR: &str = "open_log_dir"; - /// 分隔符 3 - pub const SEPARATOR_3: &str = "sep_3"; - /// 开机自启 - pub const AUTO_START: &str = "auto_start"; - /// 分隔符 4 - pub const SEPARATOR_4: &str = "sep_4"; - /// 退出 - pub const QUIT: &str = "quit"; - - /// 获取所有必需的菜单项 ID 列表 - pub fn all_required_ids() -> Vec<&'static str> { - vec![ - STATUS_INFO, - CREDENTIAL_INFO, - REQUEST_INFO, - START_SERVER, - STOP_SERVER, - REFRESH_TOKENS, - HEALTH_CHECK, - OPEN_WINDOW, - COPY_API_ADDRESS, - OPEN_LOG_DIR, - AUTO_START, - QUIT, - ] - } -} +pub use proxycast_core::tray_menu_meta::menu_ids; +pub use proxycast_core::tray_menu_meta::{get_menu_item_ids, parse_server_address}; /// 托盘菜单构建错误 #[derive(Debug, thiserror::Error)] @@ -253,33 +202,6 @@ pub fn build_tray_menu( .map_err(|e| MenuBuildError::MenuError(e.to_string())) } -/// 解析服务器地址字符串为 host 和 port -/// -/// 支持格式: -/// - "host:port" -> (host, port) -/// - "host" -> (host, 8080) -/// - "" -> ("127.0.0.1", 8080) -fn parse_server_address(address: &str) -> (String, u16) { - if address.is_empty() { - return ("127.0.0.1".to_string(), 8080); - } - - if let Some((host, port_str)) = address.rsplit_once(':') { - if let Ok(port) = port_str.parse::() { - return (host.to_string(), port); - } - } - - (address.to_string(), 8080) -} - -/// 获取菜单中包含的所有菜单项 ID -/// -/// 用于验证菜单构建的完整性 -pub fn get_menu_item_ids() -> Vec<&'static str> { - menu_ids::all_required_ids() -} - #[cfg(test)] mod tests { use super::*; diff --git a/src-tauri/src/tray/state.rs b/src-tauri/src/tray/state.rs index ff6abf6af..417d661fe 100644 --- a/src-tauri/src/tray/state.rs +++ b/src-tauri/src/tray/state.rs @@ -1,214 +1,8 @@ -//! 托盘状态模块 +//! 托盘状态模块(桥接层) //! -//! 定义托盘图标状态和状态快照结构 +//! 纯逻辑已迁移到 `proxycast-core` crate, +//! 本模块保留兼容导出。 -use serde::{Deserialize, Serialize}; - -/// 托盘图标状态枚举 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum TrayIconStatus { - /// 正常运行(绿色)- 服务器运行且凭证健康 - Running, - /// 警告状态(黄色)- 有凭证即将过期或余额不足 - Warning, - /// 错误状态(红色)- 服务器停止或所有凭证无效 - Error, - /// 停止状态(灰色)- 服务器未启动 - Stopped, -} - -impl Default for TrayIconStatus { - fn default() -> Self { - Self::Stopped - } -} - -/// 凭证健康状态 -#[derive(Debug, Clone, Default)] -pub struct CredentialHealth { - /// 凭证是否有效 - pub is_valid: bool, - /// 是否即将过期 - pub is_expiring_soon: bool, - /// 是否余额不足 - pub is_low_balance: bool, -} - -impl CredentialHealth { - /// 创建健康的凭证状态 - pub fn healthy() -> Self { - Self { - is_valid: true, - is_expiring_soon: false, - is_low_balance: false, - } - } - - /// 创建无效的凭证状态 - pub fn invalid() -> Self { - Self { - is_valid: false, - is_expiring_soon: false, - is_low_balance: false, - } - } - - /// 检查凭证是否有警告 - pub fn has_warning(&self) -> bool { - self.is_valid && (self.is_expiring_soon || self.is_low_balance) - } -} - -/// 托盘状态快照 -#[derive(Debug, Clone, Serialize)] -pub struct TrayStateSnapshot { - /// 图标状态 - pub icon_status: TrayIconStatus, - /// 服务器是否运行 - pub server_running: bool, - /// 服务器地址 - pub server_address: String, - /// 可用凭证数 - pub available_credentials: usize, - /// 总凭证数 - pub total_credentials: usize, - /// 今日请求数 - pub today_requests: u64, - /// 是否开机自启 - pub auto_start_enabled: bool, -} - -impl Default for TrayStateSnapshot { - fn default() -> Self { - Self { - icon_status: TrayIconStatus::Stopped, - server_running: false, - server_address: String::new(), - available_credentials: 0, - total_credentials: 0, - today_requests: 0, - auto_start_enabled: false, - } - } -} - -/// 根据服务器状态和凭证健康状态计算托盘图标状态 -/// -/// # 规则 -/// - 服务器未运行 -> Stopped -/// - 服务器运行 + 所有凭证无效 -> Error -/// - 服务器运行 + 有凭证警告 -> Warning -/// - 服务器运行 + 所有凭证健康 -> Running -pub fn calculate_icon_status( - server_running: bool, - credentials: &[CredentialHealth], -) -> TrayIconStatus { - if !server_running { - return TrayIconStatus::Stopped; - } - - // 如果没有凭证,视为错误状态 - if credentials.is_empty() { - return TrayIconStatus::Error; - } - - // 检查是否所有凭证都无效 - let all_invalid = credentials.iter().all(|c| !c.is_valid); - if all_invalid { - return TrayIconStatus::Error; - } - - // 检查是否有任何凭证有警告 - let has_warning = credentials.iter().any(|c| c.has_warning()); - if has_warning { - return TrayIconStatus::Warning; - } - - TrayIconStatus::Running -} - -#[cfg(test)] -mod tests { - use super::*; - use proptest::prelude::*; - - // 生成任意的 CredentialHealth - fn arb_credential_health() -> impl Strategy { - (any::(), any::(), any::()).prop_map( - |(is_valid, is_expiring_soon, is_low_balance)| CredentialHealth { - is_valid, - is_expiring_soon, - is_low_balance, - }, - ) - } - - proptest! { - /// **Feature: system-tray, Property 1: 状态到图标映射正确性** - /// **Validates: Requirements 1.1, 1.2, 1.3** - #[test] - fn prop_icon_status_mapping( - server_running in any::(), - credentials in prop::collection::vec(arb_credential_health(), 0..10) - ) { - let status = calculate_icon_status(server_running, &credentials); - - // 规则 1: 服务器未运行 -> Stopped - if !server_running { - prop_assert_eq!(status, TrayIconStatus::Stopped); - return Ok(()); - } - - // 规则 2: 没有凭证 -> Error - if credentials.is_empty() { - prop_assert_eq!(status, TrayIconStatus::Error); - return Ok(()); - } - - // 规则 3: 所有凭证无效 -> Error - let all_invalid = credentials.iter().all(|c| !c.is_valid); - if all_invalid { - prop_assert_eq!(status, TrayIconStatus::Error); - return Ok(()); - } - - // 规则 4: 有凭证警告 -> Warning - let has_warning = credentials.iter().any(|c| c.has_warning()); - if has_warning { - prop_assert_eq!(status, TrayIconStatus::Warning); - return Ok(()); - } - - // 规则 5: 其他情况 -> Running - prop_assert_eq!(status, TrayIconStatus::Running); - } - } - - #[test] - fn test_credential_health_healthy() { - let health = CredentialHealth::healthy(); - assert!(health.is_valid); - assert!(!health.is_expiring_soon); - assert!(!health.is_low_balance); - assert!(!health.has_warning()); - } - - #[test] - fn test_credential_health_invalid() { - let health = CredentialHealth::invalid(); - assert!(!health.is_valid); - assert!(!health.has_warning()); // 无效凭证不算警告 - } - - #[test] - fn test_credential_health_warning() { - let mut health = CredentialHealth::healthy(); - health.is_expiring_soon = true; - assert!(health.has_warning()); - - let mut health2 = CredentialHealth::healthy(); - health2.is_low_balance = true; - assert!(health2.has_warning()); - } -} +pub use proxycast_core::tray_state::{ + calculate_icon_status, CredentialHealth, TrayIconStatus, TrayStateSnapshot, +}; diff --git a/src-tauri/src/tray/sync.rs b/src-tauri/src/tray/sync.rs index feaab9959..da0fd3a85 100644 --- a/src-tauri/src/tray/sync.rs +++ b/src-tauri/src/tray/sync.rs @@ -229,51 +229,3 @@ impl TraySynchronizer { Ok(()) } } - -/// 从 ProviderPoolService 获取凭证健康状态 -/// -/// 将 ProviderPoolService 中的凭证状态转换为 CredentialHealth 列表 -pub fn get_credential_health_from_pool( - pool_credentials: &[(String, bool, bool, bool)], // (id, is_valid, is_expiring_soon, is_low_balance) -) -> Vec { - pool_credentials - .iter() - .map( - |(_, is_valid, is_expiring_soon, is_low_balance)| CredentialHealth { - is_valid: *is_valid, - is_expiring_soon: *is_expiring_soon, - is_low_balance: *is_low_balance, - }, - ) - .collect() -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_get_credential_health_from_pool() { - let pool_data = vec![ - ("cred1".to_string(), true, false, false), - ("cred2".to_string(), true, true, false), - ("cred3".to_string(), false, false, false), - ]; - - let health = get_credential_health_from_pool(&pool_data); - - assert_eq!(health.len(), 3); - assert!(health[0].is_valid); - assert!(!health[0].is_expiring_soon); - assert!(health[1].is_valid); - assert!(health[1].is_expiring_soon); - assert!(!health[2].is_valid); - } - - #[test] - fn test_get_credential_health_empty() { - let pool_data: Vec<(String, bool, bool, bool)> = vec![]; - let health = get_credential_health_from_pool(&pool_data); - assert!(health.is_empty()); - } -} diff --git a/src-tauri/src/voice/README.md b/src-tauri/src/voice/README.md index 8ee394ca0..1c28aed2e 100644 --- a/src-tauri/src/voice/README.md +++ b/src-tauri/src/voice/README.md @@ -7,12 +7,12 @@ | 文件 | 说明 | |------|------| | `mod.rs` | 模块入口,导出子模块 | -| `asr_service.rs` | ASR 服务,统一管理本地 Whisper 和云端 ASR | +| `asr_service.rs` | ASR 桥接层(纯逻辑已迁移到 `crates/services/src/voice_asr_service.rs`) | | `commands.rs` | Tauri 命令,供前端调用 | -| `config.rs` | 配置管理,读写语音输入配置 | -| `output_service.rs` | 文字输出服务,模拟键盘输入和剪贴板 | -| `processor.rs` | LLM 润色处理,调用本地 API 服务器 | -| `recording_service.rs` | 录音服务,使用独立线程 + channel 通信 | +| `config.rs` | 配置桥接层(纯逻辑已迁移到 `crates/services/src/voice_config_service.rs`) | +| `output_service.rs` | 输出桥接层(纯逻辑已迁移到 `crates/services/src/voice_output_service.rs`) | +| `processor.rs` | 润色桥接层(纯逻辑已迁移到 `crates/services/src/voice_processor_service.rs`) | +| `recording_service.rs` | 录音桥接层(纯逻辑已迁移到 `crates/services/src/voice_recording_service.rs`) | | `shortcut.rs` | 全局快捷键管理 | | `window.rs` | 悬浮窗管理 | @@ -42,10 +42,11 @@ ``` voice/ -├── 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 润色) +├── asr_service.rs ──→ proxycast-services (voice_asr_service) +├── recording_service.rs ──→ proxycast-services (voice_recording_service) +├── config.rs ──→ proxycast-services (voice_config_service) +├── output_service.rs ──→ proxycast-services (voice_output_service) +├── processor.rs ──→ proxycast-services (voice_processor_service) └── commands.rs ──→ 上述所有服务 ``` diff --git a/src-tauri/src/voice/asr_service.rs b/src-tauri/src/voice/asr_service.rs index 543e99784..4689286b4 100644 --- a/src-tauri/src/voice/asr_service.rs +++ b/src-tauri/src/voice/asr_service.rs @@ -1,319 +1,6 @@ -//! ASR 服务 +//! ASR 服务(桥接层) //! -//! 统一管理语音识别服务,支持本地 Whisper 和云端 ASR。 -//! -//! ## 功能 -//! - 本地 Whisper 识别(离线、隐私) -//! - OpenAI Whisper API -//! - 百度语音识别 -//! - 讯飞语音识别(WebSocket 流式) -//! -//! ## 模型文件路径 -//! Whisper 模型文件存储在:`~/Library/Application Support/proxycast/models/whisper/` -//! -//! 支持的模型: -//! - `ggml-tiny.bin` (~75MB) -//! - `ggml-base.bin` (~142MB) -//! - `ggml-small.bin` (~466MB) -//! - `ggml-medium.bin` (~1.5GB) -//! -//! ## 使用示例 -//! ```rust,ignore -//! let credential = AsrService::get_default_credential()?.unwrap(); -//! let text = AsrService::transcribe(&credential, &audio_data, 16000).await?; -//! ``` +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块保留兼容导出。 -#[cfg(feature = "local-whisper")] -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; - -impl AsrService { - /// 获取默认 ASR 凭证 - pub fn get_default_credential() -> Result, String> { - let config = load_config().map_err(|e| e.to_string())?; - Ok(config - .credential_pool - .asr - .into_iter() - .find(|c| c.is_default && !c.disabled)) - } - - /// 获取指定 ID 的 ASR 凭证 - pub fn get_credential(id: &str) -> Result, String> { - let config = load_config().map_err(|e| e.to_string())?; - Ok(config.credential_pool.asr.into_iter().find(|c| c.id == id)) - } - - /// 使用指定凭证进行语音识别 - /// - /// 当云端服务失败时,自动回退到本地 Whisper(需求 3.4) - pub async fn transcribe( - credential: &AsrCredentialEntry, - audio_data: &[u8], - sample_rate: u32, - ) -> Result { - // 如果是本地 Whisper,直接调用 - if matches!(credential.provider, AsrProviderType::WhisperLocal) { - return Self::transcribe_whisper_local(credential, audio_data, sample_rate).await; - } - - // 云端服务:先尝试云端,失败则回退到本地 Whisper - let cloud_result = match credential.provider { - AsrProviderType::OpenAI => { - Self::transcribe_openai(credential, audio_data, sample_rate).await - } - AsrProviderType::Baidu => { - Self::transcribe_baidu(credential, audio_data, sample_rate).await - } - AsrProviderType::Xunfei => { - Self::transcribe_xunfei(credential, audio_data, sample_rate).await - } - AsrProviderType::WhisperLocal => unreachable!(), // 已在上面处理 - }; - - // 云端成功,直接返回 - if cloud_result.is_ok() { - return cloud_result; - } - - // 云端失败,尝试回退到本地 Whisper - let cloud_error = cloud_result.unwrap_err(); - tracing::warn!( - "云端 ASR 服务 ({:?}) 失败: {},尝试回退到本地 Whisper", - credential.provider, - cloud_error - ); - - // 尝试获取本地 Whisper 凭证 - match Self::get_whisper_local_credential() { - Ok(Some(whisper_credential)) => { - tracing::info!("正在使用本地 Whisper 进行回退识别..."); - match Self::transcribe_whisper_local(&whisper_credential, audio_data, sample_rate) - .await - { - Ok(text) => { - tracing::info!("本地 Whisper 回退识别成功"); - Ok(text) - } - Err(whisper_error) => { - tracing::error!("本地 Whisper 回退也失败: {}", whisper_error); - // 返回原始云端错误,因为那是用户选择的服务 - Err(format!( - "云端服务失败: {cloud_error};本地 Whisper 回退也失败: {whisper_error}" - )) - } - } - } - Ok(None) => { - tracing::warn!("未找到本地 Whisper 凭证,无法回退"); - Err(format!( - "云端服务失败: {cloud_error};未配置本地 Whisper,无法回退" - )) - } - Err(e) => { - tracing::error!("获取本地 Whisper 凭证失败: {}", e); - Err(format!( - "云端服务失败: {cloud_error};获取本地 Whisper 凭证失败: {e}" - )) - } - } - } - - /// 获取本地 Whisper 凭证(用于回退) - fn get_whisper_local_credential() -> Result, String> { - let config = load_config().map_err(|e| e.to_string())?; - Ok(config - .credential_pool - .asr - .into_iter() - .find(|c| matches!(c.provider, AsrProviderType::WhisperLocal) && !c.disabled)) - } - - /// 本地 Whisper 识别 - #[cfg(feature = "local-whisper")] - async fn transcribe_whisper_local( - credential: &AsrCredentialEntry, - audio_data: &[u8], - sample_rate: u32, - ) -> Result { - // 获取 Whisper 配置 - let whisper_config = credential - .whisper_config - .as_ref() - .ok_or("Whisper 本地配置缺失")?; - - // 获取模型文件路径 - let model_path = Self::get_whisper_model_path(&whisper_config.model)?; - - // 将 PCM 字节转换为 i16 采样 - let audio = Self::build_audio_data(audio_data, sample_rate)?; - - // 检查录音时长 - if !audio.is_valid() { - return Err("录音时间过短(需要至少 0.5 秒)".to_string()); - } - - // 转换模型大小枚举 - let model = Self::convert_model_size(&whisper_config.model); - - // 创建 Whisper 识别器 - let transcriber = - voice_core::WhisperTranscriber::new(model_path, model, &credential.language) - .map_err(|e| format!("Whisper 模型加载失败: {e}"))?; - - // 执行识别 - let result = transcriber - .transcribe(&audio) - .map_err(|e| format!("Whisper 识别失败: {e}"))?; - - Ok(result.text) - } - - /// 本地 Whisper 识别(未启用 local-whisper feature 时的 stub) - #[cfg(not(feature = "local-whisper"))] - async fn transcribe_whisper_local( - _credential: &AsrCredentialEntry, - _audio_data: &[u8], - _sample_rate: u32, - ) -> Result { - Err("本地 Whisper 功能未启用。请使用云端 ASR 服务(OpenAI、百度、讯飞)".to_string()) - } - - /// 获取 Whisper 模型文件路径 - #[cfg(feature = "local-whisper")] - fn get_whisper_model_path(model_size: &WhisperModelSize) -> Result { - // 模型文件名 - let filename = match model_size { - WhisperModelSize::Tiny => "ggml-tiny.bin", - WhisperModelSize::Base => "ggml-base.bin", - WhisperModelSize::Small => "ggml-small.bin", - WhisperModelSize::Medium => "ggml-medium.bin", - }; - - // 模型存储目录:~/Library/Application Support/proxycast/models/whisper/ - let models_dir = dirs::data_dir() - .ok_or("无法获取数据目录")? - .join("proxycast") - .join("models") - .join("whisper"); - - let model_path = models_dir.join(filename); - - // 检查模型文件是否存在 - if !model_path.exists() { - return Err(format!( - "Whisper 模型文件不存在: {}\n请下载模型文件到: {}", - filename, - models_dir.display() - )); - } - - Ok(model_path) - } - - /// 转换模型大小枚举 - #[cfg(feature = "local-whisper")] - fn convert_model_size(size: &WhisperModelSize) -> voice_core::types::WhisperModel { - match size { - WhisperModelSize::Tiny => voice_core::types::WhisperModel::Tiny, - WhisperModelSize::Base => voice_core::types::WhisperModel::Base, - WhisperModelSize::Small => voice_core::types::WhisperModel::Small, - WhisperModelSize::Medium => voice_core::types::WhisperModel::Medium, - } - } - - /// OpenAI Whisper API 识别 - 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)?; - - 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()); - } - - let result = client - .transcribe(&audio) - .await - .map_err(|e| format!("OpenAI Whisper 识别失败: {e}"))?; - - Ok(result.text) - } - - /// 百度语音识别 - async fn transcribe_baidu( - credential: &AsrCredentialEntry, - audio_data: &[u8], - sample_rate: u32, - ) -> Result { - let config = credential.baidu_config.as_ref().ok_or("百度配置缺失")?; - let audio = Self::build_audio_data(audio_data, sample_rate)?; - - let client = BaiduClient::new(config.api_key.clone(), config.secret_key.clone()); - let result = client - .transcribe(&audio) - .await - .map_err(|e| format!("百度识别失败: {e}"))?; - - Ok(result.text) - } - - /// 讯飞语音识别 - /// - /// 使用 WebSocket 流式识别,支持实时语音转文字 - async fn transcribe_xunfei( - credential: &AsrCredentialEntry, - audio_data: &[u8], - sample_rate: u32, - ) -> Result { - let config = credential.xunfei_config.as_ref().ok_or("讯飞配置缺失")?; - let audio = Self::build_audio_data(audio_data, sample_rate)?; - - // 创建讯飞客户端 - // 讯飞语言代码转换:zh -> zh_cn, en -> en_us - let xunfei_language = match credential.language.as_str() { - "zh" => "zh_cn".to_string(), - "en" => "en_us".to_string(), - other => other.to_string(), - }; - - let client = XunfeiClient::new( - config.app_id.clone(), - config.api_key.clone(), - config.api_secret.clone(), - ) - .with_language(xunfei_language); - - let result = client - .transcribe(&audio) - .await - .map_err(|e| format!("讯飞识别失败: {e}"))?; - - Ok(result.text) - } - - /// 将 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()); - } - - Ok(audio) - } -} +pub use proxycast_services::voice_asr_service::AsrService; diff --git a/src-tauri/src/voice/commands.rs b/src-tauri/src/voice/commands.rs index 26b2567d3..8c9037799 100644 --- a/src-tauri/src/voice/commands.rs +++ b/src-tauri/src/voice/commands.rs @@ -1,12 +1,14 @@ //! 语音输入 Tauri 命令 //! -//! 提供前端调用的语音输入相关命令 +//! 提供前端调用的语音输入相关命令。 -use crate::config::{VoiceInputConfig, VoiceInstruction}; +use proxycast_core::config::{VoiceInputConfig, VoiceInstruction}; +use proxycast_services::voice_command_service; use tauri::{command, AppHandle}; use super::config; -use super::recording_service::AudioDeviceInfo; +use super::recording_service::{AudioDeviceInfo, RecordingServiceState}; +use tauri::State; /// 获取所有可用的麦克风设备 #[command] @@ -28,12 +30,10 @@ pub async fn save_voice_input_config( ) -> Result<(), String> { let old_config = config::load_voice_config()?; - // 如果快捷键变化,更新注册 if old_config.shortcut != voice_config.shortcut { super::shortcut::update(&app, &voice_config.shortcut)?; } - // 如果启用状态变化 if old_config.enabled != voice_config.enabled { if voice_config.enabled { super::shortcut::register(&app, &voice_config.shortcut)?; @@ -54,36 +54,13 @@ pub async fn get_voice_instructions() -> Result, String> { /// 保存指令 #[command] pub async fn save_voice_instruction(instruction: VoiceInstruction) -> Result<(), String> { - let mut voice_config = config::load_voice_config()?; - - // 查找是否已存在 - if let Some(idx) = voice_config - .instructions - .iter() - .position(|i| i.id == instruction.id) - { - voice_config.instructions[idx] = instruction; - } else { - voice_config.instructions.push(instruction); - } - - config::save_voice_config(voice_config) + config::save_voice_instruction(instruction) } /// 删除指令 #[command] pub async fn delete_voice_instruction(id: String) -> Result<(), String> { - let mut voice_config = config::load_voice_config()?; - - // 检查是否为预设指令 - if let Some(instruction) = voice_config.instructions.iter().find(|i| i.id == id) { - if instruction.is_preset { - return Err("无法删除预设指令".to_string()); - } - } - - voice_config.instructions.retain(|i| i.id != id); - config::save_voice_config(voice_config) + config::delete_voice_instruction(&id) } /// 打开语音输入窗口 @@ -98,14 +75,7 @@ pub async fn close_voice_window(app: AppHandle) -> Result<(), String> { super::window::close_voice_window(&app) } -/// 语音识别结果 -#[derive(serde::Serialize)] -pub struct TranscribeResult { - /// 识别文本 - pub text: String, - /// 使用的 ASR 服务 - pub provider: String, -} +pub use proxycast_services::voice_command_service::{PolishResult, TranscribeResult}; /// 执行语音识别 #[command] @@ -114,100 +84,8 @@ pub async fn transcribe_audio( sample_rate: u32, credential_id: Option, ) -> Result { - use super::asr_service::AsrService; - - tracing::info!( - "[语音识别] 开始识别,音频大小: {} 字节,采样率: {}", - audio_data.len(), - sample_rate - ); - - // 检查音频数据是否有效 - if audio_data.is_empty() { - tracing::error!("[语音识别] 音频数据为空!"); - return Err("音频数据为空,请检查麦克风权限".to_string()); - } - - // 检查音频数据是否全为静音(全零) - let non_zero_count = audio_data.iter().filter(|&&b| b != 0).count(); - let non_zero_ratio = non_zero_count as f32 / audio_data.len() as f32; - tracing::info!( - "[语音识别] 非零字节比例: {:.2}% ({}/{})", - non_zero_ratio * 100.0, - non_zero_count, - audio_data.len() - ); - - if non_zero_ratio < 0.01 { - tracing::warn!("[语音识别] 音频数据几乎全为静音,可能是麦克风权限问题或未正确录音"); - } - - // 获取凭证 - let credential = if let Some(id) = credential_id { - tracing::info!("[语音识别] 使用指定凭证: {}", id); - AsrService::get_credential(&id)?.ok_or_else(|| format!("凭证不存在: {id}"))? - } else { - tracing::info!("[语音识别] 获取默认凭证..."); - match AsrService::get_default_credential() { - Ok(Some(cred)) => { - tracing::info!( - "[语音识别] 找到默认凭证: id={}, provider={:?}", - cred.id, - cred.provider - ); - cred - } - Ok(None) => { - // 打印所有 ASR 凭证用于调试 - if let Ok(config) = crate::config::load_config() { - tracing::error!( - "[语音识别] 未找到默认凭证,当前 ASR 凭证数量: {}", - config.credential_pool.asr.len() - ); - for (i, c) in config.credential_pool.asr.iter().enumerate() { - tracing::error!( - "[语音识别] 凭证 {}: id={}, is_default={}, disabled={}", - i, - c.id, - c.is_default, - c.disabled - ); - } - } - return Err("未配置语音识别服务。请在设置 → 凭证池 → ASR 中添加讯飞、百度或 OpenAI Whisper 凭证。".to_string()); - } - Err(e) => { - tracing::error!("[语音识别] 获取默认凭证失败: {}", e); - return Err(format!("获取凭证失败: {e}")); - } - } - }; - - let provider_name = match credential.provider { - crate::config::AsrProviderType::WhisperLocal => "本地 Whisper", - crate::config::AsrProviderType::OpenAI => "OpenAI Whisper", - crate::config::AsrProviderType::Baidu => "百度语音", - crate::config::AsrProviderType::Xunfei => "讯飞语音", - }; - tracing::info!("[语音识别] 使用服务: {}", provider_name); - - // 执行识别 - let text = AsrService::transcribe(&credential, &audio_data, sample_rate).await?; - tracing::info!("[语音识别] 识别完成,文本长度: {} 字符", text.len()); - - Ok(TranscribeResult { - text, - provider: provider_name.to_string(), - }) -} - -/// 润色文本结果 -#[derive(serde::Serialize)] -pub struct PolishResult { - /// 润色后的文本 - pub text: String, - /// 使用的指令 - pub instruction_name: String, + voice_command_service::transcribe_audio(&audio_data, sample_rate, credential_id.as_deref()) + .await } /// 润色文本 @@ -216,73 +94,25 @@ pub async fn polish_voice_text( text: String, instruction_id: Option, ) -> Result { - let voice_config = config::load_voice_config()?; - - // 获取指令 - let instruction_id = - instruction_id.unwrap_or_else(|| voice_config.processor.default_instruction_id.clone()); - - let instruction = voice_config - .instructions - .iter() - .find(|i| i.id == instruction_id) - .ok_or_else(|| format!("指令不存在: {instruction_id}"))?; - - // 如果是原始输出,直接返回 - if instruction_id == "raw" { - return Ok(PolishResult { - text, - instruction_name: instruction.name.clone(), - }); - } - - // 调用 LLM 润色 - let polished = super::processor::polish_text( - &text, - instruction, - voice_config.processor.polish_provider.as_deref(), - voice_config.processor.polish_model.as_deref(), - ) - .await?; - - Ok(PolishResult { - text: polished, - instruction_name: instruction.name.clone(), - }) + voice_command_service::polish_voice_text(&text, instruction_id.as_deref()).await } /// 输出文本到系统 -/// -/// 根据配置的输出模式,将文字输出到当前焦点应用 #[command] pub async fn output_voice_text(text: String, mode: Option) -> Result<(), String> { - use crate::config::VoiceOutputMode; - - // 解析输出模式 - let output_mode = match mode.as_deref() { - Some("type") => VoiceOutputMode::Type, - Some("clipboard") => VoiceOutputMode::Clipboard, - Some("both") => VoiceOutputMode::Both, - None => { - // 使用配置的默认模式 - let config = config::load_voice_config()?; - config.output.mode - } - Some(other) => return Err(format!("未知的输出模式: {other}")), - }; - - // 执行输出 - super::output_service::output_text(&text, output_mode)?; - - tracing::info!("[语音输出] 文本已输出: {} 字符", text.chars().count()); - Ok(()) + voice_command_service::output_voice_text(&text, mode.as_deref()) } -// ============ 录音控制命令 ============ -// 使用独立线程 + channel 通信解决 cpal::Stream 不是 Send 的问题 - -use super::recording_service::RecordingServiceState; -use tauri::State; +/// 停止录音的返回结果 +#[derive(serde::Serialize)] +pub struct StopRecordingResult { + /// 音频数据(i16 样本的字节数组,小端序) + pub audio_data: Vec, + /// 采样率 + pub sample_rate: u32, + /// 录音时长(秒) + pub duration: f32, +} /// 开始录音 #[command] @@ -317,8 +147,7 @@ pub async fn stop_recording( audio.duration_secs ); - // 检查音频数据是否有效 - let non_zero_samples = audio.samples.iter().filter(|&&s| s != 0).count(); + let non_zero_samples = audio.samples.iter().filter(|&&sample| sample != 0).count(); let non_zero_ratio = non_zero_samples as f32 / audio.samples.len().max(1) as f32; tracing::info!( "[录音命令] 非零样本比例: {:.2}% ({}/{})", @@ -327,7 +156,6 @@ pub async fn stop_recording( audio.samples.len() ); - // 将 i16 样本转换为字节(小端序) let bytes = audio.to_pcm16le_bytes(); Ok(StopRecordingResult { @@ -337,23 +165,11 @@ pub async fn stop_recording( }) } -/// 停止录音的返回结果 -#[derive(serde::Serialize)] -pub struct StopRecordingResult { - /// 音频数据(i16 样本的字节数组,小端序) - pub audio_data: Vec, - /// 采样率 - pub sample_rate: u32, - /// 录音时长(秒) - pub duration: f32, -} - /// 取消录音 #[command] pub async fn cancel_recording( recording_service: State<'_, RecordingServiceState>, ) -> Result<(), String> { - // 使用 try_lock 避免阻塞,如果锁被占用则跳过 match recording_service.0.try_lock() { Some(mut service) => { service.cancel(); @@ -361,7 +177,6 @@ pub async fn cancel_recording( } None => { tracing::warn!("[录音命令] 取消录音时锁被占用,跳过"); - // 即使锁被占用,也尝试直接重置状态标志 } } Ok(()) diff --git a/src-tauri/src/voice/config.rs b/src-tauri/src/voice/config.rs index 663190226..0a6953b75 100644 --- a/src-tauri/src/voice/config.rs +++ b/src-tauri/src/voice/config.rs @@ -1,43 +1,11 @@ -//! 语音输入配置管理 +//! 语音输入配置管理(桥接层) //! -//! 加载和保存语音输入相关配置 +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块保留兼容导出。 -use crate::config::{ - load_config, save_config, AsrCredentialEntry, VoiceInputConfig, VoiceInstruction, +pub use proxycast_services::voice_config_service::{ + asr_provider_name, delete_voice_instruction, get_asr_credential, get_default_asr_credential, + get_enabled_asr_credential_by_provider, get_instruction, get_instructions, + list_asr_credentials, load_voice_config, resolve_output_mode, save_voice_config, + save_voice_instruction, }; - -/// 加载语音输入配置 -pub fn load_voice_config() -> Result { - let config = load_config().map_err(|e| e.to_string())?; - Ok(config.experimental.voice_input) -} - -/// 保存语音输入配置 -pub fn save_voice_config(voice_config: VoiceInputConfig) -> Result<(), String> { - let mut config = load_config().map_err(|e| e.to_string())?; - config.experimental.voice_input = voice_config; - save_config(&config).map_err(|e| e.to_string())?; - Ok(()) -} - -/// 获取默认 ASR 凭证 -pub fn get_default_asr_credential() -> Result, String> { - let config = load_config().map_err(|e| e.to_string())?; - Ok(config - .credential_pool - .asr - .into_iter() - .find(|c| c.is_default && !c.disabled)) -} - -/// 获取指令列表 -pub fn get_instructions() -> Result, String> { - let config = load_config().map_err(|e| e.to_string())?; - Ok(config.experimental.voice_input.instructions) -} - -/// 获取指定 ID 的指令 -pub fn get_instruction(id: &str) -> Result, String> { - let instructions = get_instructions()?; - Ok(instructions.into_iter().find(|i| i.id == id)) -} diff --git a/src-tauri/src/voice/output_service.rs b/src-tauri/src/voice/output_service.rs index f1d720f1b..aedb5d268 100644 --- a/src-tauri/src/voice/output_service.rs +++ b/src-tauri/src/voice/output_service.rs @@ -1,22 +1,6 @@ -//! 文字输出服务 +//! 文字输出服务(桥接层) //! -//! 提供模拟键盘输入和剪贴板输出功能 +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块保留兼容导出。 -use crate::config::VoiceOutputMode; -use voice_core::{OutputHandler, OutputMode}; - -/// 输出文字到系统 -/// -/// 根据配置的输出模式,将文字输出到当前焦点应用 -pub fn output_text(text: &str, mode: VoiceOutputMode) -> Result<(), String> { - 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}")) -} +pub use proxycast_services::voice_output_service::output_text; diff --git a/src-tauri/src/voice/processor.rs b/src-tauri/src/voice/processor.rs index 505e4bbd9..a02caee6d 100644 --- a/src-tauri/src/voice/processor.rs +++ b/src-tauri/src/voice/processor.rs @@ -1,168 +1,6 @@ -//! 语音处理器 +//! 语音处理器(桥接层) //! -//! 处理语音识别结果的 LLM 润色 +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块保留兼容导出。 -use crate::config::VoiceInstruction; - -/// 语音润色的 System Prompt -const VOICE_POLISH_SYSTEM_PROMPT: &str = r#"你是一个语音转文字的后处理助手。 - -## 背景说明 -用户通过麦克风录音,然后使用语音识别 API(如讯飞、百度等)将语音转换为文字。由于录音环境、口音、语速等因素,识别结果可能存在以下问题: - -1. **重复字词**:语音识别的流式返回机制可能导致字词重复,如"你你好好"实际是"你好" -2. **语气词和填充词**:如"嗯"、"啊"、"那个"、"就是"、"然后"等口语习惯 -3. **同音字错误**:识别引擎可能选错同音字,如"准备"识别成"准被" -4. **环境噪音干扰**:背景噪音可能被误识别为无意义的字词 -5. **断句错误**:缺少标点或标点位置不当 -6. **多余内容**:录音开始或结束时的杂音可能被识别为无关文字 - -## 你的任务 -根据上下文语义,智能还原用户真正想表达的内容: -- 去除明显的重复字词 -- 去除无意义的语气词和填充词 -- 根据语义修正可能的同音字错误 -- 过滤掉噪音产生的无意义内容 -- 添加合适的标点符号 -- 保持用户的原意,不要添加或臆测内容 - -## 输出要求 -只输出处理后的文本,不要添加任何解释、说明或前缀。"#; - -/// 处理文本(应用指令模板) -pub fn process_text(text: &str, instruction: &VoiceInstruction) -> String { - // 替换模板中的占位符 - instruction.prompt.replace("{{text}}", text) -} - -/// 使用 LLM 润色文本 -/// -/// 通过本地 API 服务器调用 LLM 进行文本润色 -pub async fn polish_text( - text: &str, - instruction: &VoiceInstruction, - _provider: Option<&str>, - model: Option<&str>, -) -> Result { - // 如果是原始输出指令,直接返回 - if instruction.id == "raw" { - return Ok(text.to_string()); - } - - // 构建 prompt - let prompt = process_text(text, instruction); - - // 调用本地 API 服务器 - let result = call_local_llm(&prompt, model, &instruction.id).await?; - Ok(result) -} - -/// 调用本地 API 服务器进行 LLM 推理 -async fn call_local_llm( - prompt: &str, - model: Option<&str>, - instruction_id: &str, -) -> Result { - use crate::config::load_config; - - // 加载配置获取 API 地址和密钥 - let config = load_config().map_err(|e| e.to_string())?; - let base_url = format!("http://{}:{}", config.server.host, config.server.port); - let api_key = &config.server.api_key; - - // 使用配置的模型,如果没有配置则使用 deepseek-chat - let model_name = model.filter(|m| !m.is_empty()).unwrap_or("deepseek-chat"); - - tracing::info!( - "[语音润色] 使用模型: {}, 指令: {}", - model_name, - instruction_id - ); - - // 构建请求 - #[derive(serde::Serialize)] - struct Message { - role: String, - content: String, - } - - #[derive(serde::Serialize)] - struct ChatRequest { - model: String, - messages: Vec, - max_tokens: u32, - temperature: f32, - } - - // 根据指令类型决定是否使用 system prompt - let messages = if instruction_id == "default" { - // 默认润色使用专门的 system prompt - vec![ - Message { - role: "system".to_string(), - content: VOICE_POLISH_SYSTEM_PROMPT.to_string(), - }, - Message { - role: "user".to_string(), - content: prompt.to_string(), - }, - ] - } else { - // 其他指令(翻译、邮件等)直接使用 user message - vec![Message { - role: "user".to_string(), - content: prompt.to_string(), - }] - }; - - let request = ChatRequest { - model: model_name.to_string(), - messages, - max_tokens: 2048, - temperature: 0.3, - }; - - // 发送请求 - let client = reqwest::Client::new(); - let response = client - .post(format!("{base_url}/v1/chat/completions")) - .header("Authorization", format!("Bearer {api_key}")) - .header("Content-Type", "application/json") - .json(&request) - .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!("LLM API 错误: {status} - {body}")); - } - - // 解析响应 - #[derive(serde::Deserialize)] - struct Choice { - message: ResponseMessage, - } - - #[derive(serde::Deserialize)] - struct ResponseMessage { - content: Option, - } - - #[derive(serde::Deserialize)] - struct ChatResponse { - choices: Vec, - } - - let result: ChatResponse = response - .json() - .await - .map_err(|e| format!("解析响应失败: {e}"))?; - - result - .choices - .first() - .and_then(|c| c.message.content.clone()) - .ok_or_else(|| "LLM 返回空内容".to_string()) -} +pub use proxycast_services::voice_processor_service::{polish_text, process_text}; diff --git a/src-tauri/src/voice/recording_service.rs b/src-tauri/src/voice/recording_service.rs index a44175a2f..2fe446afe 100644 --- a/src-tauri/src/voice/recording_service.rs +++ b/src-tauri/src/voice/recording_service.rs @@ -1,35 +1,9 @@ //! 录音服务桥接层 //! -//! 录音核心逻辑已迁移到 `voice-core` 的 `threaded_recorder` 模块。 -//! 本模块保留 Tauri State 包装和向后兼容导出路径。 +//! 纯逻辑已迁移到 `proxycast-services` crate, +//! 本模块保留兼容导出。 -use parking_lot::Mutex; -use std::sync::Arc; - -pub use voice_core::{AudioDeviceInfo, RecordingCommand, RecordingResponse, RecordingService}; - -/// 获取所有可用的麦克风设备 -pub fn list_audio_devices() -> Result, String> { - voice_core::list_audio_devices().map_err(|e| e.to_string()) -} - -/// 全局录音服务状态(Tauri State 包装) -pub struct RecordingServiceState(pub Arc>); - -impl RecordingServiceState { - /// 创建新的录音服务状态 - pub fn new() -> Self { - Self(Arc::new(Mutex::new(RecordingService::new()))) - } -} - -impl Default for RecordingServiceState { - fn default() -> Self { - Self::new() - } -} - -/// 创建录音服务状态 -pub fn create_recording_service_state() -> RecordingServiceState { - RecordingServiceState::new() -} +pub use proxycast_services::voice_recording_service::{ + create_recording_service_state, list_audio_devices, AudioDeviceInfo, RecordingCommand, + RecordingResponse, RecordingService, RecordingServiceState, +}; diff --git a/src-tauri/src/websocket/mod.rs b/src-tauri/src/websocket/mod.rs deleted file mode 100644 index 586fdf24e..000000000 --- a/src-tauri/src/websocket/mod.rs +++ /dev/null @@ -1,19 +0,0 @@ -//! WebSocket 支持模块(重导出层) -//! -//! 实际实现位于 `proxycast-websocket` crate。 - -// 重新导出 proxycast-websocket 的所有公共模块 -pub use proxycast_websocket::handler; -pub use proxycast_websocket::lifecycle; -pub use proxycast_websocket::processor; -pub use proxycast_websocket::stream; - -// 重新导出常用类型 -pub use proxycast_websocket::{ - MessageProcessor, WsApiRequest, WsApiResponse, WsConfig, WsConnection, WsConnectionManager, - WsEndpoint, WsError, WsMessage, WsStats, WsStatsSnapshot, WsStreamChunk, WsStreamEnd, -}; - -// 保持 types 子模块兼容 -pub use proxycast_core::websocket::types; -pub use proxycast_core::websocket::{KiroTokenInfo, WsKiroEvent}; diff --git a/src-tauri/src/workspace/mod.rs b/src-tauri/src/workspace/mod.rs deleted file mode 100644 index e1b6d29e6..000000000 --- a/src-tauri/src/workspace/mod.rs +++ /dev/null @@ -1,3 +0,0 @@ -//! Workspace 管理模块 - 从 proxycast-core 重新导出 - -pub use proxycast_core::workspace::*; diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 1c0c2873c..e9ce036d4 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.60.0", + "version": "0.61.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev",