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