From 0686737331b68e24b2f5329d6c4c2e2574b81d63 Mon Sep 17 00:00:00 2001 From: coso Date: Fri, 27 Feb 2026 16:37:53 +0800 Subject: [PATCH] feat: release v0.73.0 with full pending changes Co-Authored-By: Claude Opus 4.6 (1M context) --- RELEASE_NOTES.md | 64 +- package.json | 2 +- src-tauri/Cargo.lock | 31 +- src-tauri/Cargo.toml | 4 +- src-tauri/crates/agent/Cargo.toml | 1 + .../crates/agent/src/aster_state_support.rs | 9 +- src-tauri/crates/agent/src/event_converter.rs | 45 +- src-tauri/crates/agent/src/hooks.rs | 593 +++++++++++ src-tauri/crates/agent/src/lib.rs | 6 + src-tauri/crates/agent/src/prompt/builder.rs | 88 +- .../agent/src/prompt/instruction_discovery.rs | 513 ++++++++++ src-tauri/crates/agent/src/prompt/mod.rs | 5 + src-tauri/crates/agent/src/shell_security.rs | 251 +++++ .../crates/agent/src/subagent_scheduler.rs | 312 +++++- .../crates/agent/src/tool_permissions.rs | 387 ++++++++ src-tauri/crates/core/src/config/mod.rs | 6 +- src-tauri/crates/core/src/config/types.rs | 137 +++ .../processor/src/conversation_summarizer.rs | 401 +++++++- src-tauri/crates/processor/src/steps/mod.rs | 1 + .../crates/processor/src/steps/registry.rs | 290 ++++++ .../services/src/context_memory_service.rs | 366 ++++++- src-tauri/crates/skills/src/lib.rs | 4 +- src-tauri/crates/skills/src/skill_loader.rs | 20 + src-tauri/crates/skills/src/skill_matcher.rs | 390 ++++++++ src-tauri/src/agent/aster_agent.rs | 4 +- src-tauri/src/agent/mod.rs | 2 +- src-tauri/src/agent/subagent_scheduler.rs | 20 +- src-tauri/src/app/bootstrap.rs | 21 +- src-tauri/src/app/runner.rs | 4 + src-tauri/src/app/state.rs | 22 +- src-tauri/src/commands/agent_cmd.rs | 9 +- src-tauri/src/commands/aster_agent_cmd.rs | 29 +- src-tauri/src/commands/context_memory.rs | 14 +- .../commands/ecommerce_review_reply_cmd.rs | 3 + .../src/commands/memory_management_cmd.rs | 131 ++- src-tauri/src/commands/persona_cmd.rs | 13 +- src-tauri/src/commands/skill_exec_cmd.rs | 25 +- src-tauri/src/commands/subagent_cmd.rs | 19 +- src-tauri/src/commands/unified_chat_cmd.rs | 27 +- src-tauri/src/commands/unified_memory_cmd.rs | 37 +- src-tauri/src/config/tests.rs | 3 + src-tauri/src/services/auto_memory_service.rs | 343 +++++++ .../services/memory_import_parser_service.rs | 264 +++++ .../services/memory_profile_prompt_service.rs | 173 ++++ .../services/memory_rules_loader_service.rs | 243 +++++ .../memory_source_resolver_service.rs | 664 +++++++++++++ src-tauri/src/services/mod.rs | 5 + src-tauri/tauri.conf.json | 2 +- .../agent/chat/components/MessageList.tsx | 22 + .../chat/hooks/useAsterAgentChat.test.tsx | 56 ++ .../agent/chat/hooks/useAsterAgentChat.ts | 46 + src/components/agent/chat/types.ts | 3 + src/components/memory/MemoryPage.tsx | 236 ++++- src/components/memory/UnifiedMemoryPage.tsx | 2 +- src/components/memory/UnifiedMemoryTest.tsx | 2 +- .../memory/memoryLayerMetrics.test.ts | 152 +++ src/components/memory/memoryLayerMetrics.ts | 92 ++ src/components/settings-v2/_layout/index.tsx | 19 +- .../settings-v2/general/appearance/index.tsx | 1 - .../settings-v2/general/memory/index.test.tsx | 242 +++++ .../settings-v2/general/memory/index.tsx | 922 ++++++++++++++++++ .../settings-v2/hooks/useSettingsCategory.ts | 12 +- .../settings-v2/legacy-import-guard.test.ts | 1 + .../system/channels/AIChannelsList.tsx | 26 +- .../system/channels/ConnectionTestButton.tsx | 2 +- .../channels/NotificationChannelFormModal.tsx | 1 - .../channels/NotificationChannelsList.tsx | 26 +- .../system/channels/SendTestMessageButton.tsx | 2 +- .../system/external-tools/index.tsx | 240 ----- src/hooks/useTauri.ts | 150 +++ src/lib/api/agent.ts | 18 + src/lib/api/unifiedMemory.ts | 24 +- src/types/settings.ts | 6 +- 73 files changed, 7844 insertions(+), 462 deletions(-) create mode 100644 src-tauri/crates/agent/src/hooks.rs create mode 100644 src-tauri/crates/agent/src/prompt/instruction_discovery.rs create mode 100644 src-tauri/crates/agent/src/shell_security.rs create mode 100644 src-tauri/crates/agent/src/tool_permissions.rs create mode 100644 src-tauri/crates/processor/src/steps/registry.rs create mode 100644 src-tauri/crates/skills/src/skill_matcher.rs create mode 100644 src-tauri/src/services/auto_memory_service.rs create mode 100644 src-tauri/src/services/memory_import_parser_service.rs create mode 100644 src-tauri/src/services/memory_profile_prompt_service.rs create mode 100644 src-tauri/src/services/memory_rules_loader_service.rs create mode 100644 src-tauri/src/services/memory_source_resolver_service.rs create mode 100644 src/components/memory/memoryLayerMetrics.test.ts create mode 100644 src/components/memory/memoryLayerMetrics.ts create mode 100644 src/components/settings-v2/general/memory/index.test.tsx create mode 100644 src/components/settings-v2/general/memory/index.tsx delete mode 100644 src/components/settings-v2/system/external-tools/index.tsx diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index 533c4209a..7d560dd90 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -1,45 +1,39 @@ -## ProxyCast v0.72.0 +## ProxyCast v0.73.0 -发布日期:2026-02-26 +发布日期:2026-02-27 ### ✨ 新功能 -#### 渠道管理重构 -- 重写渠道设置页面:移除旧的「AI 模型提供商」和「消息通知渠道」双 tab 布局,改为 Telegram / Discord / 飞书 三个 Bot 渠道 tab,每个 tab 内联表单配置 -- 新增后端 ChannelsConfig 类型:在 Rust 配置层新增 `ChannelsConfig`、`TelegramBotConfig`、`DiscordBotConfig`、`FeishuBotConfig` 结构体,支持 YAML 序列化/反序列化 -- Telegram Bot 配置:支持 Enable 开关、Bot Token(密码输入+显示切换)、允许的用户 ID 列表、默认模型选择 -- Discord Bot 配置:支持 Enable 开关、Bot Token、允许的服务器 ID 列表、默认模型选择 -- 飞书 Bot 配置:支持 Enable 开关、App ID、App Secret、Verification Token(可选)、Encrypt Key(可选)、默认模型选择 -- 默认模型选择器:复用现有 Provider Pool 数据,下拉列出所有已配置 Provider 的模型 -- 脏状态检测:修改表单后底部固定栏显示「未保存的更改」提示,支持保存和取消操作 +#### 记忆管理系统 +- 新增多层记忆架构:支持组织策略、项目记忆、用户记忆、项目本地记忆四层配置 +- 新增记忆画像(MemoryProfile):可配置学习状态、擅长领域、解释风格、难题偏好 +- 新增记忆设置页面(settings-v2/general/memory),支持记忆来源、自动记忆、画像等配置 +- 新增记忆层级指标统计(memoryLayerMetrics),量化各层记忆贡献 +- 新增 memory profile prompt 服务,将记忆画像自动合并到系统提示词 -#### Agent Chat 改进 -- ChatSidebar 精简(减少约 300 行冗余代码) -- CharacterMention 角色提及组件功能增强 -- Inputbar 新增 SkillBadge 组件和相关 hooks -- 新增 Agent Chat 集成测试 +#### Agent 增强 +- Agent 支持上下文准备轨迹(ContextTrace)事件,前端可展示上下文注入过程 +- 新增 instruction discovery 模块,自动发现项目级指令文件 +- 新增 shell security 和 tool permissions 模块 +- 新增 hooks 模块,支持 Agent 生命周期钩子 +- SessionConfigBuilder 支持 include_context_trace 配置 -#### 内容创作增强 -- 新增 `content-creator/canvas/shared/` 共享组件目录 -- Document、Music、Novel、Poster、Script、Video 画布均有功能增强 -- 视频工作区 PromptInput、VideoCanvas、VideoWorkspace 组件优化 +#### 技能与处理器 +- 新增 skill matcher 模块,优化技能匹配逻辑 +- 新增 processor steps registry,统一步骤注册管理 + +#### 渠道管理 +- 新增 ChannelsConfig 配置类型与渠道管理 UI 组件 + +### 🐛 修复 +- 修复 workspace_mismatch 错误:会话切换 workspace 时自动更新 working_dir,不再阻断用户操作 +- 修复前端 lint 错误:清理未使用的导入和不必要的 try/catch 包装 +- 修复 Config 测试中缺少 channels 字段导致编译失败的问题 ### 🔧 优化与重构 - -#### 设置页面迁移 -- 删除旧版 `src/components/settings/` 下 13 个组件(AboutSection、ConnectionsSettings、DeveloperSettings、ExperimentalSettings、ExtensionsSettings、ExternalToolsSettings、GeneralSettings、LanguageSelector、ProxySettings、SettingsPage、UpdateNotification 等) -- settings-v2 布局和导航结构优化 - -#### 其他改进 -- 通用聊天 ChatPanel 和 CompactModelSelector 组件优化 -- 图像生成 ImageGenPage 功能增强 -- input-kit ModelSelector 组件改进 -- Smart Input ChatInput 和 SmartInputWindow 优化 -- 终端 AI TerminalAIInput 和 TerminalAIPanel 改进 -- 工具页面、工作台、记忆管理、插件系统、资源管理页面更新 -- 外观设置页优化 +- 优化 unified memory API 和前端调用 +- 移除废弃的 external-tools 设置页面 ### 📦 技术细节 -- 62 个文件变更,+1551 行,-3217 行(净减少 1666 行代码) -- Rust 后端新增渠道配置类型,前端 TypeScript 类型同步更新 -- 旧版设置页面完全迁移至 settings-v2 架构 +- 54 个文件变更,+2279 行,-410 行 +- 新增 10 个文件,涵盖记忆管理、Agent 安全、技能匹配等模块 diff --git a/package.json b/package.json index 631fc55d5..1d75ffbf0 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.72.0", + "version": "0.73.0", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 817cf61df..4ed3af220 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -6685,7 +6685,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.72.0" +version = "0.73.0" dependencies = [ "anyhow", "arboard", @@ -6785,7 +6785,7 @@ dependencies = [ [[package]] name = "proxycast-agent" -version = "0.72.0" +version = "0.73.0" dependencies = [ "aster-core", "async-trait", @@ -6795,6 +6795,7 @@ dependencies = [ "proxycast-mcp", "proxycast-providers", "proxycast-services", + "regex", "rmcp", "serde", "serde_json", @@ -6808,7 +6809,7 @@ dependencies = [ [[package]] name = "proxycast-config" -version = "0.72.0" +version = "0.73.0" dependencies = [ "async-trait", "parking_lot", @@ -6824,7 +6825,7 @@ dependencies = [ [[package]] name = "proxycast-core" -version = "0.72.0" +version = "0.73.0" dependencies = [ "aster-models", "async-trait", @@ -6864,7 +6865,7 @@ dependencies = [ [[package]] name = "proxycast-credential" -version = "0.72.0" +version = "0.73.0" dependencies = [ "axum 0.7.9", "base64 0.22.1", @@ -6899,7 +6900,7 @@ dependencies = [ [[package]] name = "proxycast-infra" -version = "0.72.0" +version = "0.73.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -6919,7 +6920,7 @@ dependencies = [ [[package]] name = "proxycast-mcp" -version = "0.72.0" +version = "0.73.0" dependencies = [ "async-trait", "glob", @@ -6950,7 +6951,7 @@ dependencies = [ [[package]] name = "proxycast-processor" -version = "0.72.0" +version = "0.73.0" dependencies = [ "async-trait", "parking_lot", @@ -6969,7 +6970,7 @@ dependencies = [ [[package]] name = "proxycast-providers" -version = "0.72.0" +version = "0.73.0" dependencies = [ "anyhow", "async-stream", @@ -7021,7 +7022,7 @@ dependencies = [ [[package]] name = "proxycast-server" -version = "0.72.0" +version = "0.73.0" dependencies = [ "async-stream", "axum 0.7.9", @@ -7063,7 +7064,7 @@ dependencies = [ [[package]] name = "proxycast-server-utils" -version = "0.72.0" +version = "0.73.0" dependencies = [ "axum 0.7.9", "futures", @@ -7078,7 +7079,7 @@ dependencies = [ [[package]] name = "proxycast-services" -version = "0.72.0" +version = "0.73.0" dependencies = [ "anyhow", "aster-core", @@ -7119,7 +7120,7 @@ dependencies = [ [[package]] name = "proxycast-skills" -version = "0.72.0" +version = "0.73.0" dependencies = [ "async-trait", "dirs 5.0.1", @@ -7135,7 +7136,7 @@ dependencies = [ [[package]] name = "proxycast-terminal" -version = "0.72.0" +version = "0.73.0" dependencies = [ "async-trait", "base64 0.22.1", @@ -7162,7 +7163,7 @@ dependencies = [ [[package]] name = "proxycast-websocket" -version = "0.72.0" +version = "0.73.0" dependencies = [ "axum 0.7.9", "chrono", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index c36c9b4f6..e13d77223 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.72.0" +version = "0.73.0" edition = "2021" authors = ["coso"] repository = "https://github.com/aiclientproxy/proxycast" @@ -189,7 +189,7 @@ version = "2.4" [package] name = "proxycast" -version = "0.72.0" +version = "0.73.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/crates/agent/Cargo.toml b/src-tauri/crates/agent/Cargo.toml index ec789faa6..30587c909 100644 --- a/src-tauri/crates/agent/Cargo.toml +++ b/src-tauri/crates/agent/Cargo.toml @@ -22,6 +22,7 @@ chrono.workspace = true dirs.workspace = true uuid.workspace = true thiserror.workspace = true +regex.workspace = true [dev-dependencies] tempfile.workspace = true diff --git a/src-tauri/crates/agent/src/aster_state_support.rs b/src-tauri/crates/agent/src/aster_state_support.rs index b19e60f84..94e3a4acb 100644 --- a/src-tauri/crates/agent/src/aster_state_support.rs +++ b/src-tauri/crates/agent/src/aster_state_support.rs @@ -102,6 +102,7 @@ pub struct SessionConfigBuilder { id: String, max_turns: Option, system_prompt: Option, + include_context_trace: Option, } impl SessionConfigBuilder { @@ -110,6 +111,7 @@ impl SessionConfigBuilder { id: id.into(), max_turns: None, system_prompt: None, + include_context_trace: None, } } @@ -123,6 +125,11 @@ impl SessionConfigBuilder { self } + pub fn include_context_trace(mut self, include: bool) -> Self { + self.include_context_trace = Some(include); + self + } + pub fn build(self) -> SessionConfig { SessionConfig { id: self.id, @@ -130,7 +137,7 @@ impl SessionConfigBuilder { max_turns: self.max_turns, retry_config: None, system_prompt: self.system_prompt, - include_context_trace: None, + include_context_trace: self.include_context_trace, } } } diff --git a/src-tauri/crates/agent/src/event_converter.rs b/src-tauri/crates/agent/src/event_converter.rs index 93d8c151b..549556007 100644 --- a/src-tauri/crates/agent/src/event_converter.rs +++ b/src-tauri/crates/agent/src/event_converter.rs @@ -264,6 +264,10 @@ pub enum TauriAgentEvent { #[serde(rename = "model_change")] ModelChange { model: String, mode: String }, + /// 上下文准备轨迹 + #[serde(rename = "context_trace")] + ContextTrace { steps: Vec }, + /// 完成(单次响应完成) #[serde(rename = "done")] Done { @@ -323,6 +327,13 @@ pub struct TauriTokenUsage { pub output_tokens: u32, } +/// 上下文准备轨迹步骤 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TauriContextTraceStep { + pub stage: String, + pub detail: String, +} + /// 简化的消息结构 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TauriMessage { @@ -391,10 +402,15 @@ pub fn convert_agent_event(event: AgentEvent) -> Vec { tracing::debug!("History replaced"); vec![] } - AgentEvent::ContextTrace { steps } => { - tracing::debug!("Context trace received, steps: {}", steps.len()); - vec![] - } + AgentEvent::ContextTrace { steps } => vec![TauriAgentEvent::ContextTrace { + steps: steps + .into_iter() + .map(|step| TauriContextTraceStep { + stage: step.stage, + detail: step.detail, + }) + .collect(), + }], } } @@ -687,6 +703,27 @@ mod tests { } } + #[test] + fn test_convert_context_trace() { + let event = AgentEvent::ContextTrace { + steps: vec![aster::context::ContextTraceStep { + stage: "memory_injection".to_string(), + detail: "query_len=10,injected=2".to_string(), + }], + }; + + let events = convert_agent_event(event); + assert_eq!(events.len(), 1); + match &events[0] { + TauriAgentEvent::ContextTrace { steps } => { + assert_eq!(steps.len(), 1); + assert_eq!(steps[0].stage, "memory_injection"); + assert_eq!(steps[0].detail, "query_len=10,injected=2"); + } + _ => panic!("Expected ContextTrace event"), + } + } + #[test] fn test_extract_tool_result_text_should_handle_nested_content_and_error() { let payload = serde_json::json!({ diff --git a/src-tauri/crates/agent/src/hooks.rs b/src-tauri/crates/agent/src/hooks.rs new file mode 100644 index 000000000..9abdd6134 --- /dev/null +++ b/src-tauri/crates/agent/src/hooks.rs @@ -0,0 +1,593 @@ +//! Agent Hook 系统 +//! +//! 提供轻量级事件钩子,允许在工具调用、提交等操作前后执行自定义 shell 命令。 + +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use std::path::Path; +use tokio::process::Command; + +/// Hook 事件类型 +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] +pub enum HookEvent { + BeforeToolCall, + AfterToolCall, + BeforePromptSubmit, + AfterPromptSubmit, + AfterCommit, + OnError, + SessionStart, + SessionEnd, + SubagentStart, + SubagentStop, + PreCompact, + PermissionRequest, +} + +/// Hook 匹配条件 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct HookMatcher { + /// 匹配特定工具名(支持正则:/pattern/) + pub tool: Option, + /// 匹配特定模式的内容 + pub content_pattern: Option, +} + +/// 单个 Hook 定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct HookDefinition { + pub event: HookEvent, + #[serde(default)] + pub matcher: HookMatcher, + /// 要执行的 shell 命令 + pub command: String, + /// 超时时间(秒) + #[serde(default = "default_timeout")] + pub timeout_secs: u64, + /// Hook 失败是否阻止原操作 + #[serde(default)] + pub blocking: bool, + /// 是否异步后台执行(不等待结果) + #[serde(default)] + pub async_exec: bool, +} + +fn default_timeout() -> u64 { + 10 +} + +/// Hook 执行结果 +#[derive(Debug)] +pub struct HookResult { + pub success: bool, + pub stdout: String, + pub stderr: String, + pub blocked: bool, + /// 注入到对话上下文的额外信息 + pub additional_context: Option, +} + +/// Hook 执行上下文 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct HookContext { + pub tool_name: Option, + pub content: Option, + pub metadata: HashMap, +} + +/// Hook 配置文件结构(旧格式) +#[derive(Debug, Deserialize)] +struct HookConfig { + hooks: Vec, +} + +/// 新格式:按事件分组 +#[derive(Debug, Deserialize)] +#[serde(untagged)] +enum HookConfigFormat { + /// 旧格式:{ "hooks": [...] } + Legacy(HookConfig), + /// 新格式:{ "hooks": { "BeforeToolCall": [...], ... } } + Grouped(GroupedHookConfig), +} + +#[derive(Debug, Deserialize)] +struct GroupedHookConfig { + hooks: HashMap>, +} + +#[derive(Debug, Deserialize)] +struct GroupedHookEntry { + command: String, + #[serde(default)] + matcher: HookMatcher, + #[serde(default = "default_timeout")] + timeout_secs: u64, + #[serde(default)] + blocking: bool, + #[serde(default)] + async_exec: bool, +} + +/// Hook 管理器 +pub struct HookManager { + hooks: Vec, +} + +impl HookManager { + pub fn new() -> Self { + Self { hooks: Vec::new() } + } + + /// 从配置文件加载 hooks + pub fn load_from_config(config_path: &Path) -> Result> { + let content = std::fs::read_to_string(config_path)?; + let format: HookConfigFormat = serde_json::from_str(&content)?; + let hooks = match format { + HookConfigFormat::Legacy(config) => config.hooks, + HookConfigFormat::Grouped(grouped) => { + let mut hooks = Vec::new(); + for (event, entries) in grouped.hooks { + for entry in entries { + hooks.push(HookDefinition { + event: event.clone(), + matcher: entry.matcher, + command: entry.command, + timeout_secs: entry.timeout_secs, + blocking: entry.blocking, + async_exec: entry.async_exec, + }); + } + } + hooks + } + }; + Ok(Self { hooks }) + } + + /// 注册一个 hook + pub fn register(&mut self, hook: HookDefinition) { + self.hooks.push(hook); + } + + /// 触发指定事件的所有匹配 hooks + pub async fn trigger(&self, event: HookEvent, context: &HookContext) -> Vec { + let matching: Vec<&HookDefinition> = self + .hooks + .iter() + .filter(|h| h.event == event && Self::matches(h, context)) + .collect(); + + let mut results = Vec::with_capacity(matching.len()); + for hook in matching { + results.push(Self::execute_hook(hook, context).await); + } + results + } + + /// 检查是否有任何 hook 阻止了操作 + pub fn is_blocked(results: &[HookResult]) -> bool { + results.iter().any(|r| r.blocked) + } + + fn matches(hook: &HookDefinition, context: &HookContext) -> bool { + if let Some(ref tool_pattern) = hook.matcher.tool { + match &context.tool_name { + Some(name) => { + if tool_pattern.starts_with('/') + && tool_pattern.ends_with('/') + && tool_pattern.len() > 2 + { + let pattern = &tool_pattern[1..tool_pattern.len() - 1]; + match regex::Regex::new(pattern) { + Ok(re) => { + if !re.is_match(name) { + return false; + } + } + Err(_) => return false, + } + } else if name != tool_pattern { + return false; + } + } + None => return false, + } + } + if let Some(ref content_pattern) = hook.matcher.content_pattern { + match &context.content { + Some(content) => { + if !content.contains(content_pattern) { + return false; + } + } + None => return false, + } + } + true + } + + async fn execute_hook(hook: &HookDefinition, context: &HookContext) -> HookResult { + let context_json = serde_json::to_string(context).unwrap_or_default(); + + let child = Command::new("sh") + .arg("-c") + .arg(&hook.command) + .env( + "HOOK_EVENT", + serde_json::to_string(&hook.event).unwrap_or_default(), + ) + .env("HOOK_TOOL_NAME", context.tool_name.as_deref().unwrap_or("")) + .env("HOOK_CONTEXT", &context_json) + .stdin(std::process::Stdio::piped()) + .stdout(std::process::Stdio::piped()) + .stderr(std::process::Stdio::piped()) + .spawn(); + + let mut child = match child { + Ok(c) => c, + Err(e) => { + return HookResult { + success: false, + stdout: String::new(), + stderr: format!("执行失败: {e}"), + blocked: hook.blocking, + additional_context: None, + }; + } + }; + + // 通过 stdin 写入完整上下文 JSON + if let Some(mut stdin) = child.stdin.take() { + use tokio::io::AsyncWriteExt; + let _ = stdin.write_all(context_json.as_bytes()).await; + drop(stdin); + } + + // 异步后台执行,不等待结果 + if hook.async_exec { + tokio::spawn(async move { + let _ = child.wait().await; + }); + return HookResult { + success: true, + stdout: String::new(), + stderr: String::new(), + blocked: false, + additional_context: None, + }; + } + + let result = tokio::time::timeout( + std::time::Duration::from_secs(hook.timeout_secs), + child.wait_with_output(), + ) + .await; + + match result { + Ok(Ok(output)) => { + let success = output.status.success(); + let stdout = String::from_utf8_lossy(&output.stdout).into_owned(); + let stderr = String::from_utf8_lossy(&output.stderr).into_owned(); + let additional_context = if success { + serde_json::from_str::(&stdout) + .ok() + .and_then(|v| { + v.get("additional_context") + .and_then(|c| c.as_str().map(String::from)) + }) + } else { + None + }; + HookResult { + success, + stdout, + stderr, + blocked: hook.blocking && !success, + additional_context, + } + } + Ok(Err(e)) => HookResult { + success: false, + stdout: String::new(), + stderr: format!("执行失败: {e}"), + blocked: hook.blocking, + additional_context: None, + }, + Err(_) => HookResult { + success: false, + stdout: String::new(), + stderr: format!("超时 ({}s)", hook.timeout_secs), + blocked: hook.blocking, + additional_context: None, + }, + } + } +} + +impl Default for HookManager { + fn default() -> Self { + Self::new() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn make_context(tool: Option<&str>, content: Option<&str>) -> HookContext { + HookContext { + tool_name: tool.map(String::from), + content: content.map(String::from), + metadata: HashMap::new(), + } + } + + fn make_hook(event: HookEvent, command: &str, blocking: bool) -> HookDefinition { + HookDefinition { + event, + matcher: HookMatcher::default(), + command: command.to_string(), + timeout_secs: 5, + blocking, + async_exec: false, + } + } + + #[test] + fn test_new_manager_is_empty() { + let mgr = HookManager::new(); + assert!(mgr.hooks.is_empty()); + } + + #[test] + fn test_register_hook() { + let mut mgr = HookManager::new(); + mgr.register(make_hook(HookEvent::BeforeToolCall, "echo hi", false)); + assert_eq!(mgr.hooks.len(), 1); + } + + #[test] + fn test_matcher_no_constraints() { + let hook = make_hook(HookEvent::BeforeToolCall, "echo hi", false); + let ctx = make_context(None, None); + assert!(HookManager::matches(&hook, &ctx)); + } + + #[test] + fn test_matcher_tool_match() { + let mut hook = make_hook(HookEvent::BeforeToolCall, "echo hi", false); + hook.matcher.tool = Some("read_file".to_string()); + + let ctx_match = make_context(Some("read_file"), None); + assert!(HookManager::matches(&hook, &ctx_match)); + + let ctx_no_match = make_context(Some("write_file"), None); + assert!(!HookManager::matches(&hook, &ctx_no_match)); + + let ctx_none = make_context(None, None); + assert!(!HookManager::matches(&hook, &ctx_none)); + } + + #[test] + fn test_matcher_content_pattern() { + let mut hook = make_hook(HookEvent::BeforeToolCall, "echo hi", false); + hook.matcher.content_pattern = Some("secret".to_string()); + + let ctx_match = make_context(None, Some("this has secret inside")); + assert!(HookManager::matches(&hook, &ctx_match)); + + let ctx_no_match = make_context(None, Some("nothing here")); + assert!(!HookManager::matches(&hook, &ctx_no_match)); + } + + #[test] + fn test_is_blocked() { + let results = vec![ + HookResult { + success: true, + stdout: String::new(), + stderr: String::new(), + blocked: false, + additional_context: None, + }, + HookResult { + success: false, + stdout: String::new(), + stderr: String::new(), + blocked: true, + additional_context: None, + }, + ]; + assert!(HookManager::is_blocked(&results)); + + let results_ok = vec![HookResult { + success: true, + stdout: String::new(), + stderr: String::new(), + blocked: false, + additional_context: None, + }]; + assert!(!HookManager::is_blocked(&results_ok)); + } + + #[tokio::test] + async fn test_trigger_executes_matching_hooks() { + let mut mgr = HookManager::new(); + mgr.register(make_hook(HookEvent::BeforeToolCall, "echo hello", false)); + mgr.register(make_hook(HookEvent::AfterToolCall, "echo world", false)); + + let ctx = make_context(None, None); + let results = mgr.trigger(HookEvent::BeforeToolCall, &ctx).await; + assert_eq!(results.len(), 1); + assert!(results[0].success); + assert!(results[0].stdout.contains("hello")); + } + + #[tokio::test] + async fn test_trigger_blocking_hook_failure() { + let mut mgr = HookManager::new(); + mgr.register(make_hook(HookEvent::OnError, "exit 1", true)); + + let ctx = make_context(None, None); + let results = mgr.trigger(HookEvent::OnError, &ctx).await; + assert_eq!(results.len(), 1); + assert!(!results[0].success); + assert!(results[0].blocked); + assert!(HookManager::is_blocked(&results)); + } + + #[tokio::test] + async fn test_trigger_timeout() { + let mut mgr = HookManager::new(); + let mut hook = make_hook(HookEvent::BeforeToolCall, "sleep 30", true); + hook.timeout_secs = 1; + mgr.register(hook); + + let ctx = make_context(None, None); + let results = mgr.trigger(HookEvent::BeforeToolCall, &ctx).await; + assert_eq!(results.len(), 1); + assert!(!results[0].success); + assert!(results[0].blocked); + assert!(results[0].stderr.contains("超时")); + } + + #[test] + fn test_load_from_config() { + let dir = tempfile::tempdir().unwrap(); + let config_path = dir.path().join("hooks.json"); + let config = r#"{ + "hooks": [ + { + "event": "BeforeToolCall", + "command": "echo test", + "blocking": false + } + ] + }"#; + std::fs::write(&config_path, config).unwrap(); + + let mgr = HookManager::load_from_config(&config_path).unwrap(); + assert_eq!(mgr.hooks.len(), 1); + assert_eq!(mgr.hooks[0].event, HookEvent::BeforeToolCall); + assert_eq!(mgr.hooks[0].timeout_secs, 10); // default + } + + #[test] + fn test_load_from_config_invalid_path() { + let result = HookManager::load_from_config(Path::new("/nonexistent/hooks.json")); + assert!(result.is_err()); + } + + // --- 新增测试 --- + + #[test] + fn test_regex_tool_matching() { + let mut hook = make_hook(HookEvent::BeforeToolCall, "echo hi", false); + hook.matcher.tool = Some("/^read_.*/".to_string()); + + let ctx_match = make_context(Some("read_file"), None); + assert!(HookManager::matches(&hook, &ctx_match)); + + let ctx_match2 = make_context(Some("read_dir"), None); + assert!(HookManager::matches(&hook, &ctx_match2)); + + let ctx_no_match = make_context(Some("write_file"), None); + assert!(!HookManager::matches(&hook, &ctx_no_match)); + + let ctx_none = make_context(None, None); + assert!(!HookManager::matches(&hook, &ctx_none)); + } + + #[test] + fn test_regex_invalid_pattern() { + let mut hook = make_hook(HookEvent::BeforeToolCall, "echo hi", false); + hook.matcher.tool = Some("/[invalid/".to_string()); + + let ctx = make_context(Some("anything"), None); + assert!(!HookManager::matches(&hook, &ctx)); + } + + #[tokio::test] + async fn test_async_exec_hook() { + let mut mgr = HookManager::new(); + let mut hook = make_hook(HookEvent::BeforeToolCall, "sleep 10", false); + hook.async_exec = true; + mgr.register(hook); + + let ctx = make_context(None, None); + let start = std::time::Instant::now(); + let results = mgr.trigger(HookEvent::BeforeToolCall, &ctx).await; + let elapsed = start.elapsed(); + + assert_eq!(results.len(), 1); + assert!(results[0].success); + assert!(results[0].stdout.is_empty()); + assert!(results[0].additional_context.is_none()); + // 异步执行应该立即返回,不会等待 sleep 10 + assert!(elapsed.as_secs() < 2); + } + + #[test] + fn test_new_hook_events() { + // 验证新事件类型可以正确序列化/反序列化 + let events = vec![ + HookEvent::SessionStart, + HookEvent::SessionEnd, + HookEvent::SubagentStart, + HookEvent::SubagentStop, + HookEvent::PreCompact, + HookEvent::PermissionRequest, + ]; + for event in &events { + let json = serde_json::to_string(event).unwrap(); + let deserialized: HookEvent = serde_json::from_str(&json).unwrap(); + assert_eq!(&deserialized, event); + } + } + + #[test] + fn test_grouped_config_format() { + let dir = tempfile::tempdir().unwrap(); + let config_path = dir.path().join("hooks.json"); + let config = r#"{ + "hooks": { + "BeforeToolCall": [ + { + "command": "echo before", + "blocking": true + } + ], + "AfterToolCall": [ + { + "command": "echo after1" + }, + { + "command": "echo after2", + "matcher": { "tool": "read_file" } + } + ] + } + }"#; + std::fs::write(&config_path, config).unwrap(); + + let mgr = HookManager::load_from_config(&config_path).unwrap(); + assert_eq!(mgr.hooks.len(), 3); + + let before_hooks: Vec<_> = mgr + .hooks + .iter() + .filter(|h| h.event == HookEvent::BeforeToolCall) + .collect(); + assert_eq!(before_hooks.len(), 1); + assert!(before_hooks[0].blocking); + assert_eq!(before_hooks[0].command, "echo before"); + + let after_hooks: Vec<_> = mgr + .hooks + .iter() + .filter(|h| h.event == HookEvent::AfterToolCall) + .collect(); + assert_eq!(after_hooks.len(), 2); + } +} diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs index c7f857a4f..31e1b4932 100644 --- a/src-tauri/crates/agent/src/lib.rs +++ b/src-tauri/crates/agent/src/lib.rs @@ -8,11 +8,14 @@ pub mod aster_state; pub mod aster_state_support; pub mod credential_bridge; pub mod event_converter; +pub mod hooks; pub mod lsp_bridge; pub mod mcp_bridge; pub mod prompt; pub mod session_store; +pub mod shell_security; pub mod subagent_scheduler; +pub mod tool_permissions; pub mod tools; pub use ask_bridge::{create_ask_callback, extract_response as extract_ask_response}; @@ -31,7 +34,10 @@ pub use prompt::SystemPromptBuilder; pub use session_store::{ create_session_sync, get_session_sync, list_sessions_sync, SessionDetail, SessionInfo, }; +pub use shell_security::ShellSecurityChecker; pub use subagent_scheduler::{ ProxyCastScheduler, ProxyCastSubAgentExecutor, SchedulerEventEmitter, SubAgentProgressEvent, + SubAgentRole, }; +pub use tool_permissions::{DynamicPermissionCheck, PermissionBehavior}; pub use tools::{BrowserAction, BrowserTool, BrowserToolError, BrowserToolResult}; diff --git a/src-tauri/crates/agent/src/prompt/builder.rs b/src-tauri/crates/agent/src/prompt/builder.rs index 6b8426a46..e52dae5ca 100644 --- a/src-tauri/crates/agent/src/prompt/builder.rs +++ b/src-tauri/crates/agent/src/prompt/builder.rs @@ -2,9 +2,10 @@ //! //! 组装完整的模块化系统提示词 +use super::instruction_discovery::{discover_instructions, merge_instructions}; use super::templates::*; use chrono::Utc; -use std::path::Path; +use std::path::{Path, PathBuf}; /// System Prompt 构建选项 #[derive(Debug, Clone, Default)] @@ -46,6 +47,10 @@ impl SystemPromptOptions { /// System Prompt 构建器 pub struct SystemPromptBuilder { options: SystemPromptOptions, + /// 启用指令发现的工作目录 + instruction_discovery_dir: Option, + /// Skill 描述(注入到 system prompt) + skill_prompt: Option, } impl Default for SystemPromptBuilder { @@ -59,12 +64,18 @@ impl SystemPromptBuilder { pub fn new() -> Self { Self { options: SystemPromptOptions::default_all(), + instruction_discovery_dir: None, + skill_prompt: None, } } /// 使用自定义选项创建构建器 pub fn with_options(options: SystemPromptOptions) -> Self { - Self { options } + Self { + options, + instruction_discovery_dir: None, + skill_prompt: None, + } } /// 设置工作目录 @@ -79,6 +90,20 @@ impl SystemPromptBuilder { self } + /// 启用层级化指令发现(从 AGENT.md 文件加载) + pub fn with_instruction_discovery(mut self, working_dir: impl AsRef) -> Self { + self.instruction_discovery_dir = Some(working_dir.as_ref().to_path_buf()); + self + } + + /// 设置 Skills 描述文本(注入到 system prompt) + pub fn with_skill_prompt(mut self, skill_prompt: String) -> Self { + if !skill_prompt.is_empty() { + self.skill_prompt = Some(skill_prompt); + } + self + } + /// 构建完整的 System Prompt pub fn build(&self) -> String { let mut parts: Vec<&str> = Vec::new(); @@ -122,7 +147,23 @@ impl SystemPromptBuilder { prompt.push_str(&env_info); } - // 添加自定义指令 + // 添加层级化指令(优先级低于 custom_instructions) + if let Some(ref dir) = self.instruction_discovery_dir { + let layers = discover_instructions(dir); + let merged = merge_instructions(&layers); + if !merged.is_empty() { + prompt.push_str("\n\n# 项目指令\n\n"); + prompt.push_str(&merged); + } + } + + // Skill 描述 + if let Some(ref skill_prompt) = self.skill_prompt { + prompt.push_str("\n\n"); + prompt.push_str(skill_prompt); + } + + // 添加自定义指令(最高优先级) if let Some(ref custom) = self.options.custom_instructions { prompt.push_str("\n\n# 附加指令\n\n"); prompt.push_str(custom); @@ -154,6 +195,8 @@ impl SystemPromptBuilder { #[cfg(test)] mod tests { use super::*; + use std::fs; + use tempfile::TempDir; #[test] fn test_build_default_prompt() { @@ -176,4 +219,43 @@ mod tests { let prompt = SystemPromptBuilder::new().working_dir("/tmp/test").build(); assert!(prompt.contains("/tmp/test")); } + + #[test] + fn test_build_with_instruction_discovery() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write( + tmp.path().join("AGENT.md"), + "# 测试项目指令\n使用 Rust 编写", + ) + .unwrap(); + + let prompt = SystemPromptBuilder::new() + .with_instruction_discovery(tmp.path()) + .build(); + assert!(prompt.contains("测试项目指令")); + assert!(prompt.contains("使用 Rust 编写")); + } + + #[test] + fn test_instruction_discovery_before_custom() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write(tmp.path().join("AGENT.md"), "DISCOVERED").unwrap(); + + let prompt = SystemPromptBuilder::new() + .with_instruction_discovery(tmp.path()) + .custom_instructions("CUSTOM") + .build(); + + let disc_pos = prompt.find("DISCOVERED").unwrap(); + let custom_pos = prompt.find("CUSTOM").unwrap(); + assert!(disc_pos < custom_pos, "发现的指令应在自定义指令之前"); + } + + #[test] + fn test_no_instruction_discovery_by_default() { + let prompt = SystemPromptBuilder::new().build(); + assert!(!prompt.contains("项目指令")); + } } diff --git a/src-tauri/crates/agent/src/prompt/instruction_discovery.rs b/src-tauri/crates/agent/src/prompt/instruction_discovery.rs new file mode 100644 index 000000000..1674c5bee --- /dev/null +++ b/src-tauri/crates/agent/src/prompt/instruction_discovery.rs @@ -0,0 +1,513 @@ +//! 层级化 AGENT.md 指令发现机制 +//! +//! 从文件系统发现并加载多层级的 AGENT.md 指令文件, +//! 按优先级从低到高:全局 -> 项目根 -> 当前目录 + +use std::collections::HashSet; +use std::path::{Path, PathBuf}; +use std::sync::RwLock; +use std::time::{Duration, Instant}; + +/// 支持的指令文件名列表(按优先级排序) +const INSTRUCTION_FILENAMES: &[&str] = &[ + "AGENT.md", + ".agent.md", + "agent.md", + ".proxycast/AGENT.md", + ".proxycast/instructions.md", +]; + +// 保留旧常量供测试使用(第一优先级文件名) +#[cfg(test)] +const INSTRUCTION_FILENAME: &str = "AGENT.md"; + +/// 指令来源,按优先级从低到高 +#[derive(Debug, Clone, PartialEq)] +pub enum InstructionSource { + /// ~/.proxycast/AGENT.md + Global, + /// 项目根目录/AGENT.md + Project, + /// 当前工作目录/AGENT.md(当不同于项目根时) + Directory, +} + +/// 单层指令 +#[derive(Debug, Clone)] +pub struct InstructionLayer { + pub source: InstructionSource, + pub content: String, + pub path: PathBuf, +} + +/// 在指定目录查找第一个存在的指令文件 +fn find_instruction_file(dir: &Path) -> Option { + for filename in INSTRUCTION_FILENAMES { + let path = dir.join(filename); + if path.is_file() { + return Some(path); + } + } + None +} + +/// 从文件系统发现并加载层级化指令 +/// 返回按优先级排序的指令列表(低优先级在前) +pub fn discover_instructions(working_dir: &Path) -> Vec { + let mut layers = Vec::new(); + + // 1. 全局: ~/.proxycast/ 下查找指令文件 + if let Some(home) = dirs::home_dir() { + let global_dir = home.join(".proxycast"); + // 全局层只查找 AGENT.md(不递归子目录模式) + let global_path = global_dir.join("AGENT.md"); + if let Some(layer) = load_layer(&global_path, InstructionSource::Global) { + layers.push(layer); + } + } + + // 2. 项目根: 从 working_dir 向上查找 .git 确定项目根 + let project_root = find_project_root(working_dir); + if let Some(ref root) = project_root { + if let Some(path) = find_instruction_file(root) { + if let Some(layer) = load_layer(&path, InstructionSource::Project) { + layers.push(layer); + } + } + } + + // 3. 目录级: working_dir 下查找指令文件(仅当不同于项目根时) + let is_same_as_root = project_root + .as_deref() + .map_or(false, |root| root == working_dir); + if !is_same_as_root { + if let Some(path) = find_instruction_file(working_dir) { + if let Some(layer) = load_layer(&path, InstructionSource::Directory) { + layers.push(layer); + } + } + } + + layers +} + +/// 合并多层指令为最终文本 +pub fn merge_instructions(layers: &[InstructionLayer]) -> String { + if layers.is_empty() { + return String::new(); + } + + layers + .iter() + .map(|layer| { + let label = match layer.source { + InstructionSource::Global => "全局指令", + InstructionSource::Project => "项目指令", + InstructionSource::Directory => "目录指令", + }; + format!( + "\n{}", + label, + layer.path.display(), + layer.content + ) + }) + .collect::>() + .join("\n\n") +} + +/// 从 path 向上查找包含 .git 的目录作为项目根 +fn find_project_root(path: &Path) -> Option { + let mut current = if path.is_file() { + path.parent()?.to_path_buf() + } else { + path.to_path_buf() + }; + loop { + if current.join(".git").exists() { + return Some(current); + } + if !current.pop() { + return None; + } + } +} + +/// 尝试加载单个指令文件(含 @include 展开) +fn load_layer(path: &Path, source: InstructionSource) -> Option { + let content = std::fs::read_to_string(path).ok()?; + let base_dir = path.parent().unwrap_or(Path::new(".")); + let mut visited = HashSet::new(); + visited.insert(path.to_path_buf()); + let expanded = process_includes(&content, base_dir, &mut visited); + let expanded = expanded.trim().to_string(); + if expanded.is_empty() { + return None; + } + Some(InstructionLayer { + source, + content: expanded, + path: path.to_path_buf(), + }) +} + +// --------------------------------------------------------------------------- +// @include 指令处理 +// --------------------------------------------------------------------------- + +/// 处理 @include 指令,递归展开引用的文件 +fn process_includes(content: &str, base_dir: &Path, visited: &mut HashSet) -> String { + let mut result = String::new(); + for line in content.lines() { + let trimmed = line.trim(); + if let Some(path_str) = trimmed.strip_prefix('@') { + // 跳过空路径 + if path_str.is_empty() { + result.push_str(line); + result.push('\n'); + continue; + } + // 解析路径(支持 @./path、@~/path、@/absolute/path) + let include_path = resolve_include_path(path_str.trim(), base_dir); + if let Some(ref path) = include_path { + if visited.contains(path) { + result.push_str(&format!("\n", path.display())); + continue; + } + if is_binary_file(path) { + result.push_str(&format!("\n", path.display())); + continue; + } + if let Ok(included_content) = std::fs::read_to_string(path) { + visited.insert(path.clone()); + let expanded = process_includes( + &included_content, + path.parent().unwrap_or(base_dir), + visited, + ); + result.push_str(&expanded); + if !expanded.ends_with('\n') { + result.push('\n'); + } + } else { + result.push_str(&format!("\n", path.display())); + } + } else { + // 不是有效的 include 路径,保留原文 + result.push_str(line); + result.push('\n'); + } + } else { + result.push_str(line); + result.push('\n'); + } + } + result +} + +/// 解析 include 路径 +fn resolve_include_path(path_str: &str, base_dir: &Path) -> Option { + let unescaped = path_str.replace("\\ ", " "); + if unescaped.starts_with("./") || unescaped.starts_with("../") { + Some(base_dir.join(&unescaped)) + } else if unescaped.starts_with('~') { + dirs::home_dir().map(|home| home.join(&unescaped[2..])) + } else if unescaped.starts_with('/') { + Some(PathBuf::from(&unescaped)) + } else { + // 相对路径 + Some(base_dir.join(&unescaped)) + } +} + +/// 判断是否为二进制文件(按扩展名) +fn is_binary_file(path: &Path) -> bool { + const BINARY_EXTENSIONS: &[&str] = &[ + "png", "jpg", "jpeg", "gif", "bmp", "ico", "svg", "woff", "woff2", "ttf", "eot", "zip", + "tar", "gz", "bz2", "xz", "7z", "exe", "dll", "so", "dylib", "pdf", "doc", "docx", "xls", + "xlsx", "mp3", "mp4", "avi", "mov", "wav", "wasm", "o", "a", "lib", + ]; + path.extension() + .and_then(|ext| ext.to_str()) + .map(|ext| BINARY_EXTENSIONS.contains(&ext.to_lowercase().as_str())) + .unwrap_or(false) +} + +// --------------------------------------------------------------------------- +// 缓存 +// --------------------------------------------------------------------------- + +struct CachedInstruction { + layers: Vec, + cached_at: Instant, +} + +// InstructionLayer 没有实现 Clone,手动实现缓存的 clone +impl CachedInstruction { + fn clone_layers(&self) -> Vec { + self.layers.clone() + } +} + +static CACHE: std::sync::LazyLock>> = + std::sync::LazyLock::new(|| RwLock::new(std::collections::HashMap::new())); + +/// 带缓存的指令发现(TTL 默认 60 秒) +pub fn discover_instructions_cached(working_dir: &Path, ttl: Duration) -> Vec { + let key = working_dir.to_path_buf(); + + // 检查缓存 + if let Ok(cache) = CACHE.read() { + if let Some(cached) = cache.get(&key) { + if cached.cached_at.elapsed() < ttl { + return cached.clone_layers(); + } + } + } + + // 缓存未命中或过期,重新发现 + let layers = discover_instructions(working_dir); + + if let Ok(mut cache) = CACHE.write() { + cache.insert( + key, + CachedInstruction { + layers: layers.clone(), + cached_at: Instant::now(), + }, + ); + } + + layers +} + +/// 清除指令缓存 +pub fn clear_instruction_cache() { + if let Ok(mut cache) = CACHE.write() { + cache.clear(); + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + use tempfile::TempDir; + + #[test] + fn test_discover_no_files() { + let tmp = TempDir::new().unwrap(); + let layers = discover_instructions(tmp.path()); + // 没有 AGENT.md,也没有 .git,不应发现任何指令 + // (全局指令取决于用户环境,这里只验证不会 panic) + assert!(layers + .iter() + .all(|l| l.source != InstructionSource::Project + && l.source != InstructionSource::Directory)); + } + + #[test] + fn test_discover_project_root() { + let tmp = TempDir::new().unwrap(); + // 模拟项目根 + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write( + tmp.path().join(INSTRUCTION_FILENAME), + "# Project Instructions", + ) + .unwrap(); + + let layers = discover_instructions(tmp.path()); + let project_layers: Vec<_> = layers + .iter() + .filter(|l| l.source == InstructionSource::Project) + .collect(); + assert_eq!(project_layers.len(), 1); + assert_eq!(project_layers[0].content, "# Project Instructions"); + } + + #[test] + fn test_discover_directory_layer() { + let tmp = TempDir::new().unwrap(); + // 项目根在 tmp + fs::create_dir(tmp.path().join(".git")).unwrap(); + // 子目录有自己的 AGENT.md + let subdir = tmp.path().join("src"); + fs::create_dir(&subdir).unwrap(); + fs::write(subdir.join(INSTRUCTION_FILENAME), "# Dir Instructions").unwrap(); + + let layers = discover_instructions(&subdir); + let dir_layers: Vec<_> = layers + .iter() + .filter(|l| l.source == InstructionSource::Directory) + .collect(); + assert_eq!(dir_layers.len(), 1); + assert_eq!(dir_layers[0].content, "# Dir Instructions"); + } + + #[test] + fn test_discover_no_duplicate_when_at_project_root() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write(tmp.path().join(INSTRUCTION_FILENAME), "# Root").unwrap(); + + let layers = discover_instructions(tmp.path()); + // working_dir == project_root 时不应出现 Directory 层 + let dir_layers: Vec<_> = layers + .iter() + .filter(|l| l.source == InstructionSource::Directory) + .collect(); + assert_eq!(dir_layers.len(), 0); + } + + #[test] + fn test_discover_empty_file_skipped() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write(tmp.path().join(INSTRUCTION_FILENAME), " \n ").unwrap(); + + let layers = discover_instructions(tmp.path()); + let project_layers: Vec<_> = layers + .iter() + .filter(|l| l.source == InstructionSource::Project) + .collect(); + assert_eq!(project_layers.len(), 0); + } + + #[test] + fn test_merge_instructions() { + let layers = vec![ + InstructionLayer { + source: InstructionSource::Global, + content: "global rule".to_string(), + path: PathBuf::from("/home/.proxycast/AGENT.md"), + }, + InstructionLayer { + source: InstructionSource::Project, + content: "project rule".to_string(), + path: PathBuf::from("/project/AGENT.md"), + }, + ]; + let merged = merge_instructions(&layers); + assert!(merged.contains("全局指令")); + assert!(merged.contains("global rule")); + assert!(merged.contains("项目指令")); + assert!(merged.contains("project rule")); + // 全局在前,项目在后 + assert!(merged.find("global rule").unwrap() < merged.find("project rule").unwrap()); + } + + #[test] + fn test_merge_empty() { + assert_eq!(merge_instructions(&[]), ""); + } + + #[test] + fn test_priority_order() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write(tmp.path().join(INSTRUCTION_FILENAME), "project").unwrap(); + let subdir = tmp.path().join("sub"); + fs::create_dir(&subdir).unwrap(); + fs::write(subdir.join(INSTRUCTION_FILENAME), "directory").unwrap(); + + let layers = discover_instructions(&subdir); + let non_global: Vec<_> = layers + .iter() + .filter(|l| l.source != InstructionSource::Global) + .collect(); + assert_eq!(non_global.len(), 2); + assert_eq!(non_global[0].source, InstructionSource::Project); + assert_eq!(non_global[1].source, InstructionSource::Directory); + } + + // --- 新增测试 --- + + #[test] + fn test_multi_filename_support() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + // 使用 .agent.md(第二优先级) + fs::write(tmp.path().join(".agent.md"), "dotfile agent").unwrap(); + + let layers = discover_instructions(tmp.path()); + let project: Vec<_> = layers + .iter() + .filter(|l| l.source == InstructionSource::Project) + .collect(); + assert_eq!(project.len(), 1); + assert!(project[0].content.contains("dotfile agent")); + } + + #[test] + fn test_multi_filename_priority() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + // 同时存在 AGENT.md 和 .agent.md,应优先使用 AGENT.md + fs::write(tmp.path().join("AGENT.md"), "primary agent").unwrap(); + fs::write(tmp.path().join(".agent.md"), "secondary agent").unwrap(); + + let layers = discover_instructions(tmp.path()); + let project: Vec<_> = layers + .iter() + .filter(|l| l.source == InstructionSource::Project) + .collect(); + assert_eq!(project.len(), 1); + assert!(project[0].content.contains("primary agent")); + } + + #[test] + fn test_include_directive() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write(tmp.path().join("extra.md"), "included content").unwrap(); + fs::write(tmp.path().join("AGENT.md"), "main\n@./extra.md\nend").unwrap(); + + let layers = discover_instructions(tmp.path()); + let project: Vec<_> = layers + .iter() + .filter(|l| l.source == InstructionSource::Project) + .collect(); + assert_eq!(project.len(), 1); + assert!(project[0].content.contains("included content")); + } + + #[test] + fn test_include_circular_reference() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write(tmp.path().join("a.md"), "@./b.md").unwrap(); + fs::write(tmp.path().join("b.md"), "@./a.md").unwrap(); + fs::write(tmp.path().join("AGENT.md"), "@./a.md").unwrap(); + + let layers = discover_instructions(tmp.path()); + // 不应该无限循环 + assert!(!layers.is_empty()); + } + + #[test] + fn test_binary_file_skip() { + assert!(is_binary_file(Path::new("image.png"))); + assert!(is_binary_file(Path::new("archive.zip"))); + assert!(!is_binary_file(Path::new("readme.md"))); + assert!(!is_binary_file(Path::new("code.rs"))); + } + + #[test] + fn test_cached_discovery() { + let tmp = TempDir::new().unwrap(); + fs::create_dir(tmp.path().join(".git")).unwrap(); + fs::write(tmp.path().join("AGENT.md"), "cached test").unwrap(); + + let layers1 = discover_instructions_cached(tmp.path(), Duration::from_secs(60)); + let layers2 = discover_instructions_cached(tmp.path(), Duration::from_secs(60)); + assert_eq!(layers1.len(), layers2.len()); + + // 清除缓存 + clear_instruction_cache(); + } +} diff --git a/src-tauri/crates/agent/src/prompt/mod.rs b/src-tauri/crates/agent/src/prompt/mod.rs index 687fa8ed7..a0e2fef37 100644 --- a/src-tauri/crates/agent/src/prompt/mod.rs +++ b/src-tauri/crates/agent/src/prompt/mod.rs @@ -8,7 +8,12 @@ //! - builder - 提示词构建器 pub mod builder; +pub mod instruction_discovery; pub mod templates; pub use builder::SystemPromptBuilder; +pub use instruction_discovery::{ + clear_instruction_cache, discover_instructions, discover_instructions_cached, + merge_instructions, InstructionLayer, InstructionSource, +}; pub use templates::*; diff --git a/src-tauri/crates/agent/src/shell_security.rs b/src-tauri/crates/agent/src/shell_security.rs new file mode 100644 index 000000000..6680eea0f --- /dev/null +++ b/src-tauri/crates/agent/src/shell_security.rs @@ -0,0 +1,251 @@ +//! Shell 命令安全检查 +//! +//! 对 bash/shell 工具的命令进行安全分析,检测危险操作。 + +use crate::tool_permissions::{DynamicPermissionCheck, PermissionBehavior, ToolRiskLevel}; + +/// 危险 shell 操作符 +const DANGEROUS_OPERATORS: &[&str] = &["&&", "||", ";", "|", ">", ">>", "$(", "`"]; + +/// 危险命令模式 +const DANGEROUS_COMMANDS: &[&str] = &[ + "rm -rf /", + "rm -rf ~", + "rm -rf .", + "mkfs", + "dd if=", + ":(){:|:&};:", + "chmod -R 777 /", + "wget|sh", + "curl|sh", + "curl|bash", + "wget|bash", + "> /dev/sda", + "mv / ", +]; + +/// 只读命令白名单 +const READONLY_COMMANDS: &[&str] = &[ + "ls", + "cat", + "head", + "tail", + "grep", + "find", + "wc", + "git status", + "git log", + "git diff", + "git branch", + "pwd", + "echo", + "which", + "type", + "file", + "stat", + "tree", + "du", + "df", + "env", + "printenv", + "uname", + "date", + "whoami", + "hostname", + "id", +]; + +/// Shell 安全检查结果 +#[derive(Debug, Clone)] +pub struct ShellSecurityResult { + pub safe: bool, + pub risk_level: ToolRiskLevel, + pub detected_operators: Vec, + pub is_readonly: bool, + pub reason: Option, +} + +/// Shell 安全检查器 +pub struct ShellSecurityChecker; + +impl ShellSecurityChecker { + /// 检查命令安全性 + pub fn check(command: &str) -> ShellSecurityResult { + let trimmed = command.trim(); + + // 检测危险命令 + for dangerous in DANGEROUS_COMMANDS { + if trimmed.contains(dangerous) { + return ShellSecurityResult { + safe: false, + risk_level: ToolRiskLevel::Destructive, + detected_operators: vec![], + is_readonly: false, + reason: Some(format!("检测到危险命令模式: {}", dangerous)), + }; + } + } + + let is_readonly = Self::is_readonly(trimmed); + let detected_operators = Self::detect_dangerous_operators(trimmed); + + let risk_level = if is_readonly { + ToolRiskLevel::ReadOnly + } else if detected_operators.is_empty() { + ToolRiskLevel::Reversible + } else { + ToolRiskLevel::Destructive + }; + + ShellSecurityResult { + safe: risk_level != ToolRiskLevel::Destructive, + risk_level, + detected_operators, + is_readonly, + reason: None, + } + } + + /// 是否为只读命令 + pub fn is_readonly(command: &str) -> bool { + let trimmed = command.trim(); + // 取第一个命令(管道前) + let first_cmd = trimmed.split('|').next().unwrap_or(trimmed).trim(); + // 取命令名(第一个 token) + let cmd_name = first_cmd.split_whitespace().next().unwrap_or(""); + + READONLY_COMMANDS.iter().any(|ro| { + if ro.contains(' ') { + // 多词命令(如 "git status"),前缀匹配 + first_cmd.starts_with(ro) + } else { + cmd_name == *ro + } + }) + } + + /// 检测危险操作符 + pub fn detect_dangerous_operators(command: &str) -> Vec { + DANGEROUS_OPERATORS + .iter() + .filter(|op| command.contains(**op)) + .map(|op| op.to_string()) + .collect() + } +} + +/// 为 bash 工具实现动态权限检查 +impl DynamicPermissionCheck for ShellSecurityChecker { + fn check_permissions(&self, tool_name: &str, input: &serde_json::Value) -> PermissionBehavior { + // 只检查 bash/shell 类工具 + if tool_name != "bash" && tool_name != "shell" && tool_name != "execute_command" { + return PermissionBehavior::Allow; + } + + let command = input.get("command").and_then(|v| v.as_str()).unwrap_or(""); + + if command.is_empty() { + return PermissionBehavior::Allow; + } + + let result = Self::check(command); + + if !result.safe { + let reason = result.reason.unwrap_or_else(|| { + format!("检测到危险操作符: {}", result.detected_operators.join(", ")) + }); + return PermissionBehavior::Deny { reason }; + } + + if result.is_readonly { + PermissionBehavior::Allow + } else { + PermissionBehavior::Ask { + message: format!("Shell 命令需要确认: {}", command), + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_readonly_commands() { + assert!(ShellSecurityChecker::is_readonly("ls -la")); + assert!(ShellSecurityChecker::is_readonly("git status")); + assert!(ShellSecurityChecker::is_readonly("cat file.txt")); + assert!(ShellSecurityChecker::is_readonly("grep pattern file")); + assert!(ShellSecurityChecker::is_readonly("pwd")); + } + + #[test] + fn test_non_readonly_commands() { + assert!(!ShellSecurityChecker::is_readonly("rm file.txt")); + assert!(!ShellSecurityChecker::is_readonly("cargo build")); + assert!(!ShellSecurityChecker::is_readonly("npm install")); + } + + #[test] + fn test_dangerous_commands() { + let result = ShellSecurityChecker::check("rm -rf /"); + assert!(!result.safe); + assert_eq!(result.risk_level, ToolRiskLevel::Destructive); + + let result = ShellSecurityChecker::check("mkfs.ext4 /dev/sda1"); + assert!(!result.safe); + } + + #[test] + fn test_safe_commands() { + let result = ShellSecurityChecker::check("ls -la"); + assert!(result.safe); + assert!(result.is_readonly); + assert_eq!(result.risk_level, ToolRiskLevel::ReadOnly); + } + + #[test] + fn test_detect_operators() { + let ops = ShellSecurityChecker::detect_dangerous_operators("echo hello && rm file"); + assert!(ops.contains(&"&&".to_string())); + } + + #[test] + fn test_dynamic_permission_check_readonly() { + let checker = ShellSecurityChecker; + let input = serde_json::json!({"command": "ls -la"}); + assert_eq!( + checker.check_permissions("bash", &input), + PermissionBehavior::Allow + ); + } + + #[test] + fn test_dynamic_permission_check_dangerous() { + let checker = ShellSecurityChecker; + let input = serde_json::json!({"command": "rm -rf /"}); + match checker.check_permissions("bash", &input) { + PermissionBehavior::Deny { .. } => {} + other => panic!("Expected Deny, got {:?}", other), + } + } + + #[test] + fn test_dynamic_permission_check_non_bash() { + let checker = ShellSecurityChecker; + let input = serde_json::json!({"command": "rm -rf /"}); + assert_eq!( + checker.check_permissions("read_file", &input), + PermissionBehavior::Allow + ); + } + + #[test] + fn test_reversible_command() { + let result = ShellSecurityChecker::check("cargo build"); + assert!(result.safe); + assert!(!result.is_readonly); + assert_eq!(result.risk_level, ToolRiskLevel::Reversible); + } +} diff --git a/src-tauri/crates/agent/src/subagent_scheduler.rs b/src-tauri/crates/agent/src/subagent_scheduler.rs index 000ba48b6..326d3ea9b 100644 --- a/src-tauri/crates/agent/src/subagent_scheduler.rs +++ b/src-tauri/crates/agent/src/subagent_scheduler.rs @@ -15,6 +15,7 @@ use aster::agents::subagent_scheduler::{ }; use aster::conversation::message::Message; use chrono::Utc; +use serde::{Deserialize, Serialize}; use tokio::sync::RwLock; use tracing::{debug, info, warn}; @@ -24,6 +25,110 @@ use proxycast_core::database::DbConnection; /// 调度器事件发射器 pub type SchedulerEventEmitter = Arc; +// --------------------------------------------------------------------------- +// SubAgentRole +// --------------------------------------------------------------------------- + +/// SubAgent 角色,决定可用的工具集 +/// +/// 遵循最小权限原则:默认 Explorer(只读),需要写入时显式升级。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub enum SubAgentRole { + /// 只读探索:Read, Grep, Glob, LSP 查询 + Explorer, + /// 规划分析:Read + 输出计划文档 + Planner, + /// 全能执行:所有工具(不限制) + Executor, +} + +impl SubAgentRole { + /// 返回该角色允许使用的工具名称列表 + /// + /// 空列表表示不限制(Executor 角色) + pub fn allowed_tools(&self) -> Vec<&'static str> { + match self { + Self::Explorer => vec!["read_file", "grep", "glob", "list_directory", "lsp_query"], + Self::Planner => vec!["read_file", "grep", "glob", "list_directory", "write_file"], + Self::Executor => vec![], // 空表示不限制 + } + } + + /// 返回该角色的最大对话轮次 + pub fn max_turns(&self) -> usize { + match self { + Self::Explorer => 15, + Self::Planner => 10, + Self::Executor => 30, + } + } + + /// 返回该角色的结果最大长度(字符数) + /// 0 表示不限制 + pub fn max_result_length(&self) -> usize { + match self { + Self::Explorer => 2000, + Self::Planner => 4000, + Self::Executor => 0, // 不限制 + } + } + + /// 该角色是否允许使用指定工具 + pub fn is_tool_allowed(&self, tool_name: &str) -> bool { + let allowed = self.allowed_tools(); + allowed.is_empty() || allowed.contains(&tool_name) + } + + /// 将角色的工具限制应用到 SubAgentTask 上 + /// + /// 如果任务已经设置了 allowed_tools,取交集;否则直接设置。 + /// Executor 角色不做任何修改。 + pub fn apply_to_task(&self, mut task: SubAgentTask) -> SubAgentTask { + let role_tools = self.allowed_tools(); + if role_tools.is_empty() { + // Executor: 不限制 + return task; + } + + let role_set: std::collections::HashSet<&str> = role_tools.into_iter().collect(); + + if let Some(ref existing) = task.allowed_tools { + // 取交集:任务自身限制 ∩ 角色限制 + let filtered: Vec = existing + .iter() + .filter(|t| role_set.contains(t.as_str())) + .cloned() + .collect(); + task.allowed_tools = Some(filtered); + } else { + task.allowed_tools = Some(role_set.into_iter().map(String::from).collect()); + } + + task + } +} + +impl Default for SubAgentRole { + fn default() -> Self { + Self::Explorer // 默认最小权限 + } +} + +impl std::fmt::Display for SubAgentRole { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Explorer => write!(f, "explorer"), + Self::Planner => write!(f, "planner"), + Self::Executor => write!(f, "executor"), + } + } +} + +// --------------------------------------------------------------------------- +// ProxyCastSubAgentExecutor +// --------------------------------------------------------------------------- + /// ProxyCast SubAgent 执行器 /// /// 实现 aster-rust 的 SubAgentExecutor trait, @@ -37,6 +142,8 @@ pub struct ProxyCastSubAgentExecutor { default_model: String, /// 默认 Provider 类型 default_provider: String, + /// SubAgent 角色 + role: SubAgentRole, } impl ProxyCastSubAgentExecutor { @@ -47,6 +154,7 @@ impl ProxyCastSubAgentExecutor { db, default_model: "claude-sonnet-4-20250514".to_string(), default_provider: "anthropic".to_string(), + role: SubAgentRole::default(), } } @@ -62,6 +170,17 @@ impl ProxyCastSubAgentExecutor { self } + /// 设置 SubAgent 角色 + pub fn with_role(mut self, role: SubAgentRole) -> Self { + self.role = role; + self + } + + /// 获取当前角色 + pub fn role(&self) -> SubAgentRole { + self.role + } + /// 从凭证池选择凭证 async fn select_credential(&self, task: &SubAgentTask) -> SchedulerResult { let model = task.model.as_deref().unwrap_or(&self.default_model); @@ -96,7 +215,7 @@ impl SubAgentExecutor for ProxyCastSubAgentExecutor { context: &AgentContext, ) -> SchedulerResult { let start_time = Utc::now(); - info!("执行 SubAgent 任务: {}", task.id); + info!("执行 SubAgent 任务: {} (角色: {})", task.id, self.role); let provider_config = self.select_credential(task).await?; debug!("使用凭证: {}", provider_config.credential_uuid); @@ -115,6 +234,16 @@ impl SubAgentExecutor for ProxyCastSubAgentExecutor { let response = response_msg.as_concat_text(); + // 按角色限制结果长度 + let max_len = self.role.max_result_length(); + let response = if max_len > 0 && response.chars().count() > max_len { + let original_len = response.len(); + let truncated: String = response.chars().take(max_len).collect(); + format!("{}\n\n[结果已截断,原始 {} 字符]", truncated, original_len) + } else { + response + }; + let end_time = Utc::now(); let duration = (end_time - start_time).to_std().unwrap_or(Duration::ZERO); @@ -146,12 +275,18 @@ impl SubAgentExecutor for ProxyCastSubAgentExecutor { } } +// --------------------------------------------------------------------------- +// ProxyCastScheduler +// --------------------------------------------------------------------------- + /// ProxyCast SubAgent 调度器 pub struct ProxyCastScheduler { /// 内部调度器 scheduler: Arc>>>, /// 数据库连接 db: DbConnection, + /// 默认角色 + default_role: SubAgentRole, } impl ProxyCastScheduler { @@ -160,9 +295,16 @@ impl ProxyCastScheduler { Self { scheduler: Arc::new(RwLock::new(None)), db, + default_role: SubAgentRole::default(), } } + /// 设置默认角色 + pub fn with_default_role(mut self, role: SubAgentRole) -> Self { + self.default_role = role; + self + } + /// 初始化调度器(不附带事件回调) pub async fn init(&self, config: Option) { self.init_with_event_emitter(config, None).await; @@ -174,7 +316,7 @@ impl ProxyCastScheduler { config: Option, event_emitter: Option, ) { - let executor = ProxyCastSubAgentExecutor::new(self.db.clone()); + let executor = ProxyCastSubAgentExecutor::new(self.db.clone()).with_role(self.default_role); let config = config.unwrap_or_default(); let scheduler = if let Some(emitter) = event_emitter { @@ -189,20 +331,41 @@ impl ProxyCastScheduler { }; *self.scheduler.write().await = Some(scheduler); - info!("ProxyCast SubAgent 调度器初始化完成"); + info!( + "ProxyCast SubAgent 调度器初始化完成 (默认角色: {})", + self.default_role + ); } /// 执行任务 + /// + /// 根据调度器的默认角色自动对每个任务应用工具限制。 pub async fn execute( &self, tasks: Vec, parent_context: Option<&AgentContext>, + ) -> SchedulerResult { + self.execute_with_role(tasks, parent_context, self.default_role) + .await + } + + /// 使用指定角色执行任务 + /// + /// 角色的工具限制会应用到每个任务上(与任务自身的 allowed_tools 取交集)。 + pub async fn execute_with_role( + &self, + tasks: Vec, + parent_context: Option<&AgentContext>, + role: SubAgentRole, ) -> SchedulerResult { let scheduler = self.scheduler.read().await; let scheduler = scheduler .as_ref() .ok_or_else(|| SchedulerError::ContextError("调度器未初始化".to_string()))?; + // 应用角色工具限制 + let tasks: Vec = tasks.into_iter().map(|t| role.apply_to_task(t)).collect(); + scheduler.execute(tasks, parent_context).await } @@ -214,6 +377,10 @@ impl ProxyCastScheduler { } } +// --------------------------------------------------------------------------- +// SubAgentProgressEvent +// --------------------------------------------------------------------------- + /// Tauri 事件:SubAgent 进度 #[derive(Debug, Clone, serde::Serialize)] #[serde(rename_all = "camelCase")] @@ -230,6 +397,8 @@ pub struct SubAgentProgressEvent { pub percentage: f64, /// 当前任务 pub current_tasks: Vec, + /// SubAgent 角色 + pub role: Option, } impl From for SubAgentProgressEvent { @@ -241,6 +410,143 @@ impl From for SubAgentProgressEvent { running: progress.running, percentage: progress.percentage, current_tasks: progress.current_tasks, + role: None, } } } + +impl SubAgentProgressEvent { + /// 附加角色信息 + pub fn with_role(mut self, role: SubAgentRole) -> Self { + self.role = Some(role.to_string()); + self + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_role_default_is_explorer() { + assert_eq!(SubAgentRole::default(), SubAgentRole::Explorer); + } + + #[test] + fn test_explorer_allowed_tools() { + let role = SubAgentRole::Explorer; + let tools = role.allowed_tools(); + assert!(tools.contains(&"read_file")); + assert!(tools.contains(&"grep")); + assert!(tools.contains(&"glob")); + assert!(tools.contains(&"list_directory")); + assert!(tools.contains(&"lsp_query")); + assert!(!tools.contains(&"write_file")); + } + + #[test] + fn test_planner_allowed_tools() { + let role = SubAgentRole::Planner; + let tools = role.allowed_tools(); + assert!(tools.contains(&"read_file")); + assert!(tools.contains(&"write_file")); + assert!(!tools.contains(&"lsp_query")); + } + + #[test] + fn test_executor_no_restriction() { + let role = SubAgentRole::Executor; + assert!(role.allowed_tools().is_empty()); + assert!(role.is_tool_allowed("anything")); + } + + #[test] + fn test_is_tool_allowed() { + let explorer = SubAgentRole::Explorer; + assert!(explorer.is_tool_allowed("read_file")); + assert!(!explorer.is_tool_allowed("write_file")); + assert!(!explorer.is_tool_allowed("execute_command")); + } + + #[test] + fn test_apply_to_task_explorer() { + let role = SubAgentRole::Explorer; + let task = SubAgentTask::new("t1", "explore", "test prompt"); + let task = role.apply_to_task(task); + + let allowed = task.allowed_tools.unwrap(); + assert!(allowed.contains(&"read_file".to_string())); + assert!(!allowed.contains(&"write_file".to_string())); + } + + #[test] + fn test_apply_to_task_executor_no_change() { + let role = SubAgentRole::Executor; + let task = SubAgentTask::new("t1", "code", "test prompt"); + let task = role.apply_to_task(task); + + assert!(task.allowed_tools.is_none()); + } + + #[test] + fn test_apply_to_task_intersection() { + let role = SubAgentRole::Explorer; + // 任务自身只允许 read_file 和 write_file + let task = SubAgentTask::new("t1", "explore", "test") + .with_allowed_tools(vec!["read_file", "write_file"]); + let task = role.apply_to_task(task); + + // Explorer 不允许 write_file,交集只剩 read_file + let allowed = task.allowed_tools.unwrap(); + assert_eq!(allowed, vec!["read_file".to_string()]); + } + + #[test] + fn test_role_display() { + assert_eq!(SubAgentRole::Explorer.to_string(), "explorer"); + assert_eq!(SubAgentRole::Planner.to_string(), "planner"); + assert_eq!(SubAgentRole::Executor.to_string(), "executor"); + } + + #[test] + fn test_role_serde_roundtrip() { + let role = SubAgentRole::Planner; + let json = serde_json::to_string(&role).unwrap(); + assert_eq!(json, "\"planner\""); + let deserialized: SubAgentRole = serde_json::from_str(&json).unwrap(); + assert_eq!(deserialized, role); + } + + #[test] + fn test_role_max_turns() { + assert_eq!(SubAgentRole::Explorer.max_turns(), 15); + assert_eq!(SubAgentRole::Planner.max_turns(), 10); + assert_eq!(SubAgentRole::Executor.max_turns(), 30); + } + + #[test] + fn test_role_max_result_length() { + assert_eq!(SubAgentRole::Explorer.max_result_length(), 2000); + assert_eq!(SubAgentRole::Planner.max_result_length(), 4000); + assert_eq!(SubAgentRole::Executor.max_result_length(), 0); + } + + #[test] + fn test_progress_event_with_role() { + let event = SubAgentProgressEvent { + total: 3, + completed: 1, + failed: 0, + running: 1, + percentage: 33.3, + current_tasks: vec!["task-1".to_string()], + role: None, + }; + let event = event.with_role(SubAgentRole::Explorer); + assert_eq!(event.role, Some("explorer".to_string())); + } +} diff --git a/src-tauri/crates/agent/src/tool_permissions.rs b/src-tauri/crates/agent/src/tool_permissions.rs new file mode 100644 index 000000000..f93ec4773 --- /dev/null +++ b/src-tauri/crates/agent/src/tool_permissions.rs @@ -0,0 +1,387 @@ +//! Tool 权限分级系统 +//! +//! 按操作的可逆性和影响范围对工具进行风险分级,决定是否需要用户确认。 + +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use std::collections::HashSet; + +/// 工具风险等级 +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)] +pub enum ToolRiskLevel { + /// 只读操作,无副作用 + ReadOnly, + /// 可逆操作(如编辑文件、创建分支) + Reversible, + /// 破坏性操作(如删除文件、force push) + Destructive, +} + +/// 权限检查结果(对标 Claude Code 的 allow/deny/ask) +#[derive(Debug, Clone, PartialEq)] +pub enum PermissionBehavior { + Allow, + Deny { reason: String }, + Ask { message: String }, +} + +/// 动态权限检查 trait(工具可根据输入内容判断风险) +pub trait DynamicPermissionCheck: Send + Sync { + fn check_permissions(&self, tool_name: &str, input: &serde_json::Value) -> PermissionBehavior; +} + +/// 工具权限元数据 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ToolPermissionMeta { + pub tool_name: String, + pub risk_level: ToolRiskLevel, + pub description: String, + /// 是否需要用户确认 + pub requires_confirmation: bool, +} + +/// 工具权限检查器 +pub struct ToolPermissionChecker { + permissions: HashMap, + auto_approve_level: ToolRiskLevel, + /// 会话内用户已允许的工具(tool_name → 允许次数) + session_allowed: HashMap, + /// 会话内用户已拒绝的工具 + session_denied: HashSet, + /// 动态权限检查器 + dynamic_checker: Option>, +} + +impl ToolPermissionChecker { + pub fn new() -> Self { + let mut checker = Self { + permissions: HashMap::new(), + auto_approve_level: ToolRiskLevel::ReadOnly, + session_allowed: HashMap::new(), + session_denied: HashSet::new(), + dynamic_checker: None, + }; + for meta in Self::default_permissions() { + checker.permissions.insert(meta.tool_name.clone(), meta); + } + checker + } + + /// 注册工具的权限元数据 + pub fn register_tool(&mut self, meta: ToolPermissionMeta) { + self.permissions.insert(meta.tool_name.clone(), meta); + } + + /// 检查工具是否需要用户确认 + pub fn needs_confirmation(&self, tool_name: &str) -> bool { + match self.permissions.get(tool_name) { + Some(meta) => meta.requires_confirmation && meta.risk_level > self.auto_approve_level, + // 未知工具默认需要确认 + None => true, + } + } + + /// 获取工具的风险等级 + pub fn risk_level(&self, tool_name: &str) -> ToolRiskLevel { + self.permissions + .get(tool_name) + .map(|m| m.risk_level) + // 未知工具默认为破坏性 + .unwrap_or(ToolRiskLevel::Destructive) + } + + /// 设置自动批准的风险等级 + pub fn set_auto_approve_level(&mut self, level: ToolRiskLevel) { + self.auto_approve_level = level; + } + + /// 设置动态权限检查器 + pub fn set_dynamic_checker(&mut self, checker: Box) { + self.dynamic_checker = Some(checker); + } + + /// 完整的权限决策链 + pub fn check_permission( + &self, + tool_name: &str, + input: Option<&serde_json::Value>, + ) -> PermissionBehavior { + // 1. 会话级记忆 + if let Some(allowed) = self.has_session_decision(tool_name) { + return if allowed { + PermissionBehavior::Allow + } else { + PermissionBehavior::Deny { + reason: format!("工具 {} 在本次会话中已被拒绝", tool_name), + } + }; + } + // 2. 动态检查(如 shell 安全) + if let Some(input) = input { + if let Some(checker) = &self.dynamic_checker { + let result = checker.check_permissions(tool_name, input); + if result != PermissionBehavior::Allow { + return result; + } + } + } + // 3. 静态分级 + if self.needs_confirmation(tool_name) { + PermissionBehavior::Ask { + message: format!("工具 {} 需要确认执行", tool_name), + } + } else { + PermissionBehavior::Allow + } + } + + /// 记录用户的允许决策 + pub fn record_allow(&mut self, tool_name: &str) { + let count = self + .session_allowed + .entry(tool_name.to_string()) + .or_insert(0); + *count += 1; + self.session_denied.remove(tool_name); + } + + /// 记录用户的拒绝决策 + pub fn record_deny(&mut self, tool_name: &str) { + self.session_denied.insert(tool_name.to_string()); + self.session_allowed.remove(tool_name); + } + + /// 检查是否有会话级记忆 + pub fn has_session_decision(&self, tool_name: &str) -> Option { + if self.session_allowed.contains_key(tool_name) { + Some(true) + } else if self.session_denied.contains(tool_name) { + Some(false) + } else { + None + } + } + + /// 清除会话记忆 + pub fn clear_session_memory(&mut self) { + self.session_allowed.clear(); + self.session_denied.clear(); + } + + /// 返回默认的工具权限映射 + pub fn default_permissions() -> Vec { + let read_only = &[ + ("read_file", "读取文件内容"), + ("grep", "搜索文件内容"), + ("glob", "按模式查找文件"), + ("list_directory", "列出目录内容"), + ("lsp_query", "LSP 查询"), + ]; + let reversible = &[ + ("write_file", "写入文件"), + ("edit_file", "编辑文件"), + ("create_file", "创建文件"), + ("git_commit", "Git 提交"), + ("git_branch", "Git 分支操作"), + ]; + let destructive = &[ + ("bash", "执行 Shell 命令"), + ("git_push", "Git 推送"), + ("git_force_push", "Git 强制推送"), + ("delete_file", "删除文件"), + ]; + + let mut perms = Vec::new(); + for &(name, desc) in read_only { + perms.push(ToolPermissionMeta { + tool_name: name.to_string(), + risk_level: ToolRiskLevel::ReadOnly, + description: desc.to_string(), + requires_confirmation: false, + }); + } + for &(name, desc) in reversible { + perms.push(ToolPermissionMeta { + tool_name: name.to_string(), + risk_level: ToolRiskLevel::Reversible, + description: desc.to_string(), + requires_confirmation: true, + }); + } + for &(name, desc) in destructive { + perms.push(ToolPermissionMeta { + tool_name: name.to_string(), + risk_level: ToolRiskLevel::Destructive, + description: desc.to_string(), + requires_confirmation: true, + }); + } + perms + } +} + +impl Default for ToolPermissionChecker { + fn default() -> Self { + Self::new() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_default_permissions_loaded() { + let checker = ToolPermissionChecker::new(); + assert_eq!(checker.risk_level("read_file"), ToolRiskLevel::ReadOnly); + assert_eq!(checker.risk_level("edit_file"), ToolRiskLevel::Reversible); + assert_eq!( + checker.risk_level("git_force_push"), + ToolRiskLevel::Destructive + ); + } + + #[test] + fn test_unknown_tool_defaults_destructive() { + let checker = ToolPermissionChecker::new(); + assert_eq!( + checker.risk_level("unknown_tool"), + ToolRiskLevel::Destructive + ); + assert!(checker.needs_confirmation("unknown_tool")); + } + + #[test] + fn test_read_only_no_confirmation() { + let checker = ToolPermissionChecker::new(); + assert!(!checker.needs_confirmation("read_file")); + assert!(!checker.needs_confirmation("grep")); + } + + #[test] + fn test_destructive_needs_confirmation() { + let checker = ToolPermissionChecker::new(); + assert!(checker.needs_confirmation("delete_file")); + assert!(checker.needs_confirmation("git_force_push")); + } + + #[test] + fn test_auto_approve_level_reversible() { + let mut checker = ToolPermissionChecker::new(); + checker.set_auto_approve_level(ToolRiskLevel::Reversible); + // Reversible 工具不再需要确认 + assert!(!checker.needs_confirmation("edit_file")); + assert!(!checker.needs_confirmation("write_file")); + // Destructive 仍需确认 + assert!(checker.needs_confirmation("delete_file")); + } + + #[test] + fn test_auto_approve_level_destructive() { + let mut checker = ToolPermissionChecker::new(); + checker.set_auto_approve_level(ToolRiskLevel::Destructive); + assert!(!checker.needs_confirmation("delete_file")); + assert!(!checker.needs_confirmation("git_force_push")); + } + + #[test] + fn test_register_custom_tool() { + let mut checker = ToolPermissionChecker::new(); + checker.register_tool(ToolPermissionMeta { + tool_name: "my_tool".to_string(), + risk_level: ToolRiskLevel::Reversible, + description: "自定义工具".to_string(), + requires_confirmation: false, + }); + assert_eq!(checker.risk_level("my_tool"), ToolRiskLevel::Reversible); + assert!(!checker.needs_confirmation("my_tool")); + } + + #[test] + fn test_register_overrides_default() { + let mut checker = ToolPermissionChecker::new(); + // 将 bash 从 Destructive 降级为 Reversible + checker.register_tool(ToolPermissionMeta { + tool_name: "bash".to_string(), + risk_level: ToolRiskLevel::Reversible, + description: "受限 Shell".to_string(), + requires_confirmation: false, + }); + assert_eq!(checker.risk_level("bash"), ToolRiskLevel::Reversible); + } + + #[test] + fn test_risk_level_ordering() { + assert!(ToolRiskLevel::ReadOnly < ToolRiskLevel::Reversible); + assert!(ToolRiskLevel::Reversible < ToolRiskLevel::Destructive); + } + + #[test] + fn test_default_permissions_count() { + let perms = ToolPermissionChecker::default_permissions(); + assert_eq!(perms.len(), 14); // 5 read + 5 reversible + 4 destructive + } + + #[test] + fn test_session_allow_memory() { + let mut checker = ToolPermissionChecker::new(); + assert_eq!(checker.has_session_decision("bash"), None); + checker.record_allow("bash"); + assert_eq!(checker.has_session_decision("bash"), Some(true)); + } + + #[test] + fn test_session_deny_memory() { + let mut checker = ToolPermissionChecker::new(); + checker.record_deny("bash"); + assert_eq!(checker.has_session_decision("bash"), Some(false)); + } + + #[test] + fn test_session_deny_overrides_allow() { + let mut checker = ToolPermissionChecker::new(); + checker.record_allow("bash"); + checker.record_deny("bash"); + assert_eq!(checker.has_session_decision("bash"), Some(false)); + } + + #[test] + fn test_clear_session_memory() { + let mut checker = ToolPermissionChecker::new(); + checker.record_allow("bash"); + checker.record_deny("read_file"); + checker.clear_session_memory(); + assert_eq!(checker.has_session_decision("bash"), None); + assert_eq!(checker.has_session_decision("read_file"), None); + } + + #[test] + fn test_check_permission_allow() { + let checker = ToolPermissionChecker::new(); + // read_file 是 ReadOnly,auto_approve_level 也是 ReadOnly + assert_eq!( + checker.check_permission("read_file", None), + PermissionBehavior::Allow + ); + } + + #[test] + fn test_check_permission_ask() { + let checker = ToolPermissionChecker::new(); + // bash 是 Destructive,需要确认 + match checker.check_permission("bash", None) { + PermissionBehavior::Ask { .. } => {} + other => panic!("Expected Ask, got {:?}", other), + } + } + + #[test] + fn test_check_permission_session_override() { + let mut checker = ToolPermissionChecker::new(); + checker.record_allow("bash"); + assert_eq!( + checker.check_permission("bash", None), + PermissionBehavior::Allow + ); + } +} diff --git a/src-tauri/crates/core/src/config/mod.rs b/src-tauri/crates/core/src/config/mod.rs index 306586c34..1a891092a 100644 --- a/src-tauri/crates/core/src/config/mod.rs +++ b/src-tauri/crates/core/src/config/mod.rs @@ -21,12 +21,14 @@ pub use import::{ImportOptions, ImportService, ValidationResult}; pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde}; pub use types::{ generate_secure_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, AsrCredentialEntry, - AsrProviderType, AssistantConfig, AssistantProfile, BaiduConfig, ChatAppearanceConfig, Config, + AsrProviderType, AssistantConfig, AssistantProfile, BaiduConfig, ChannelsConfig, + ChatAppearanceConfig, Config, ContentCreatorConfig, ConversationSettings, CredentialEntry, CredentialPoolConfig, CustomProviderConfig, DeliveryConfig, EndpointProvidersConfig, ExperimentalFeatures, GeminiApiKeyEntry, HeartbeatExecutionMode, HeartbeatSecurityConfig, HeartbeatSettings, HintRouteSettingsEntry, HintRouterSettings, ImageGenConfig, InjectionRuleConfig, - InjectionSettings, LoggingConfig, MemoryConfig, ModelInfo, ModelsConfig, NativeAgentConfig, + InjectionSettings, LoggingConfig, MemoryAutoConfig, MemoryConfig, MemoryProfileConfig, + MemoryResolveConfig, MemorySourcesConfig, ModelInfo, ModelsConfig, NativeAgentConfig, NavigationConfig, OpenAIAsrConfig, PairingSettings, ProviderConfig, ProviderModelsConfig, ProvidersConfig, QuotaExceededConfig, RateLimitSettings, RemoteManagementConfig, RetrySettings, RoutingConfig, ScreenshotChatConfig, ServerConfig, TaskSchedule, TlsConfig, UpdateCheckConfig, diff --git a/src-tauri/crates/core/src/config/types.rs b/src-tauri/crates/core/src/config/types.rs index e45105eb0..738e12305 100644 --- a/src-tauri/crates/core/src/config/types.rs +++ b/src-tauri/crates/core/src/config/types.rs @@ -1818,6 +1818,131 @@ pub struct ChatAppearanceConfig { pub append_selected_text_to_recommendation: Option, } +/// 记忆管理配置 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)] +pub struct MemoryProfileConfig { + /// 当前学习/工作状态(单选) + #[serde(default)] + pub current_status: Option, + /// 擅长领域(多选) + #[serde(default)] + pub strengths: Vec, + /// 偏好的解释风格(多选) + #[serde(default)] + pub explanation_style: Vec, + /// 遇到难题时的偏好(多选) + #[serde(default)] + pub challenge_preference: Vec, +} + +/// 记忆来源配置 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct MemorySourcesConfig { + /// 组织级策略文件(可选) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub managed_policy_path: Option, + /// 项目级记忆文件相对路径列表(会按目录层级向上查找) + #[serde(default)] + pub project_memory_paths: Vec, + /// 项目规则目录相对路径列表(会按目录层级向上查找) + #[serde(default)] + pub project_rule_dirs: Vec, + /// 用户级记忆文件(可选) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub user_memory_path: Option, + /// 项目本地私有记忆文件(可选) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub project_local_memory_path: Option, +} + +impl Default for MemorySourcesConfig { + fn default() -> Self { + Self { + managed_policy_path: None, + project_memory_paths: vec!["AGENTS.md".to_string(), ".agents/AGENTS.md".to_string()], + project_rule_dirs: vec![".agents/rules".to_string()], + user_memory_path: Some("~/.proxycast/AGENTS.md".to_string()), + project_local_memory_path: Some("AGENTS.local.md".to_string()), + } + } +} + +/// 自动记忆配置 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct MemoryAutoConfig { + /// 是否启用自动记忆 + #[serde(default = "default_memory_auto_enabled")] + pub enabled: bool, + /// MEMORY 入口文件名 + #[serde(default = "default_memory_auto_entrypoint")] + pub entrypoint: String, + /// 启动时加载 MEMORY 入口的最大行数 + #[serde(default = "default_memory_auto_max_loaded_lines")] + pub max_loaded_lines: u32, + /// 自动记忆根目录(可选,默认 ~/.proxycast/projects//memory) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub root_dir: Option, +} + +fn default_memory_auto_enabled() -> bool { + true +} + +fn default_memory_auto_entrypoint() -> String { + "MEMORY.md".to_string() +} + +fn default_memory_auto_max_loaded_lines() -> u32 { + 200 +} + +impl Default for MemoryAutoConfig { + fn default() -> Self { + Self { + enabled: default_memory_auto_enabled(), + entrypoint: default_memory_auto_entrypoint(), + max_loaded_lines: default_memory_auto_max_loaded_lines(), + root_dir: None, + } + } +} + +/// 记忆解析行为配置 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct MemoryResolveConfig { + /// 额外参与记忆解析的目录 + #[serde(default)] + pub additional_dirs: Vec, + /// 是否跟随 @import 引用 + #[serde(default = "default_memory_follow_imports")] + pub follow_imports: bool, + /// 最大导入深度 + #[serde(default = "default_memory_import_max_depth")] + pub import_max_depth: u8, + /// 是否从 additional_dirs 加载记忆文件 + #[serde(default)] + pub load_additional_dirs_memory: bool, +} + +fn default_memory_follow_imports() -> bool { + true +} + +fn default_memory_import_max_depth() -> u8 { + 5 +} + +impl Default for MemoryResolveConfig { + fn default() -> Self { + Self { + additional_dirs: Vec::new(), + follow_imports: default_memory_follow_imports(), + import_max_depth: default_memory_import_max_depth(), + load_additional_dirs_memory: false, + } + } +} + /// 记忆管理配置 #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)] pub struct MemoryConfig { @@ -1833,6 +1958,18 @@ pub struct MemoryConfig { /// 自动清理过期记忆 #[serde(default)] pub auto_cleanup: Option, + /// 记忆偏好画像 + #[serde(default)] + pub profile: Option, + /// 记忆来源配置 + #[serde(default)] + pub sources: MemorySourcesConfig, + /// 自动记忆配置 + #[serde(default)] + pub auto: MemoryAutoConfig, + /// 记忆解析行为配置 + #[serde(default)] + pub resolve: MemoryResolveConfig, } /// 语音服务配置 diff --git a/src-tauri/crates/processor/src/conversation_summarizer.rs b/src-tauri/crates/processor/src/conversation_summarizer.rs index ad5f10596..84904d6a7 100644 --- a/src-tauri/crates/processor/src/conversation_summarizer.rs +++ b/src-tauri/crates/processor/src/conversation_summarizer.rs @@ -5,11 +5,29 @@ use serde::{Deserialize, Serialize}; +/// 判断字符是否为 CJK(中日韩)字符 +fn is_cjk(c: char) -> bool { + matches!(c, + '\u{4E00}'..='\u{9FFF}' | // CJK Unified Ideographs + '\u{3400}'..='\u{4DBF}' | // CJK Unified Ideographs Extension A + '\u{F900}'..='\u{FAFF}' | // CJK Compatibility Ideographs + '\u{3000}'..='\u{303F}' | // CJK Symbols and Punctuation + '\u{FF00}'..='\u{FFEF}' // Halfwidth and Fullwidth Forms + ) +} + +/// 简单的 token 估算(中文约 1.5 token/字,英文约 0.75 token/word) +pub fn estimate_tokens(text: &str) -> usize { + let cjk_chars = text.chars().filter(|c| is_cjk(*c)).count(); + let non_cjk_len = text.len().saturating_sub(cjk_chars); + (cjk_chars as f64 * 1.5) as usize + (non_cjk_len as f64 * 0.25) as usize +} + /// 摘要配置 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SummaryConfig { /// 是否启用 - #[serde(default)] + #[serde(default = "default_enabled")] pub enabled: bool, /// 触发摘要的消息数阈值 #[serde(default = "default_threshold")] @@ -20,8 +38,23 @@ pub struct SummaryConfig { /// 摘要最大要点数 #[serde(default = "default_max_points")] pub max_summary_points: usize, + /// 系统消息永不压缩 + #[serde(default = "default_true")] + pub preserve_system_messages: bool, + /// 工具调用结果只保留摘要 + #[serde(default = "default_true")] + pub summarize_tool_results: bool, + /// 保留最近 N 轮完整对话(一轮 = user + assistant) + #[serde(default = "default_keep_turns")] + pub keep_recent_turns: usize, + /// Token 触发阈值(优先于消息数阈值) + #[serde(default = "default_token_threshold")] + pub token_threshold: Option, } +fn default_enabled() -> bool { + true +} fn default_threshold() -> usize { 50 } @@ -31,14 +64,27 @@ fn default_keep_recent() -> usize { fn default_max_points() -> usize { 12 } +fn default_true() -> bool { + true +} +fn default_keep_turns() -> usize { + 10 +} +fn default_token_threshold() -> Option { + Some(80000) +} impl Default for SummaryConfig { fn default() -> Self { Self { - enabled: false, + enabled: true, threshold_messages: default_threshold(), keep_recent_messages: default_keep_recent(), max_summary_points: default_max_points(), + preserve_system_messages: true, + summarize_tool_results: true, + keep_recent_turns: default_keep_turns(), + token_threshold: default_token_threshold(), } } } @@ -52,6 +98,10 @@ pub struct SummaryRequest { pub system_prompt: String, /// 需要摘要的消息(作为 user 消息发送) pub messages_to_summarize: String, + /// 被摘要的消息数 + pub messages_to_compact: usize, + /// 当前估算 token 数 + pub current_tokens: usize, } /// 摘要结果 @@ -76,15 +126,32 @@ impl ConversationSummarizer { } /// 判断是否需要摘要 - pub fn should_summarize(&self, message_count: usize) -> bool { - self.config.enabled && message_count > self.config.threshold_messages + pub fn should_summarize(&self, messages: &[serde_json::Value]) -> bool { + if !self.config.enabled { + return false; + } + + // 优先检查 token 阈值 + if let Some(token_threshold) = self.config.token_threshold { + let total_text: String = messages + .iter() + .filter_map(|m| m.get("content").and_then(|c| c.as_str())) + .collect::>() + .join(""); + let total_tokens = estimate_tokens(&total_text); + if total_tokens >= token_threshold { + return true; + } + } + + messages.len() > self.config.threshold_messages } /// 构建摘要请求 /// /// 将需要摘要的旧消息格式化为 LLM 请求 pub fn build_summary_request(&self, messages: &[serde_json::Value]) -> Option { - if !self.should_summarize(messages.len()) { + if !self.should_summarize(messages) { return None; } @@ -112,6 +179,20 @@ impl ConversationSummarizer { .get("role") .and_then(|r| r.as_str()) .unwrap_or("unknown"); + + // 工具调用结果用紧凑格式 + if self.config.summarize_tool_results { + if let Some(tool_name) = extract_tool_name(msg) { + let content = extract_content_text(msg); + let truncated = if content.len() > 200 { + format!("{}...(truncated)", &content[..200]) + } else { + content + }; + return format!("[{role}][tool:{tool_name}]: {truncated}"); + } + } + let content = extract_content_text(msg); format!("[{role}]: {content}") }) @@ -130,9 +211,13 @@ impl ConversationSummarizer { self.config.max_summary_points ); + let current_tokens = estimate_tokens(&messages_text); + Some(SummaryRequest { system_prompt, messages_to_summarize: messages_text, + messages_to_compact: to_summarize, + current_tokens, }) } @@ -187,6 +272,54 @@ impl ConversationSummarizer { } } +/// 从消息中提取工具名称(如果是工具调用或工具结果) +fn extract_tool_name(msg: &serde_json::Value) -> Option { + // tool_use 格式(Anthropic) + if let Some(content) = msg.get("content").and_then(|c| c.as_array()) { + for item in content { + if item.get("type").and_then(|t| t.as_str()) == Some("tool_use") { + return item.get("name").and_then(|n| n.as_str()).map(String::from); + } + if item.get("type").and_then(|t| t.as_str()) == Some("tool_result") { + return item + .get("tool_use_id") + .and_then(|n| n.as_str()) + .map(String::from); + } + } + } + // function_call 格式(OpenAI) + if let Some(fc) = msg.get("function_call") { + return fc.get("name").and_then(|n| n.as_str()).map(String::from); + } + // tool_calls 格式(OpenAI) + if let Some(tcs) = msg.get("tool_calls").and_then(|t| t.as_array()) { + if let Some(first) = tcs.first() { + return first + .get("function") + .and_then(|f| f.get("name")) + .and_then(|n| n.as_str()) + .map(String::from); + } + } + None +} + +/// 将 SubAgent 的完整结果压缩为摘要 +pub fn summarize_subagent_result(result: &str, max_length: usize) -> String { + if result.len() <= max_length { + return result.to_string(); + } + // 保留开头和结尾各占一半 + let half = max_length / 2; + let start = &result[..half]; + let end = &result[result.len() - half..]; + format!( + "{start}\n\n... [省略 {} 字符] ...\n\n{end}", + result.len() - max_length + ) +} + /// 从消息中提取文本内容 /// /// 兼容 OpenAI 格式(content 为字符串)和 Anthropic 格式(content 为数组) @@ -208,6 +341,101 @@ fn extract_content_text(msg: &serde_json::Value) -> String { } } +/// 在完整摘要前,先截断过长的工具输出 +/// max_tool_output_tokens: 单个工具输出的最大 token 数 +pub fn microcompact(messages: &mut [serde_json::Value], max_tool_output_tokens: usize) { + for msg in messages.iter_mut() { + if !is_tool_result(msg) { + continue; + } + let content = match extract_tool_content_text(msg) { + Some(text) => text, + None => continue, + }; + let tokens = estimate_tokens(&content); + if tokens > max_tool_output_tokens { + let truncated = truncate_to_tokens(&content, max_tool_output_tokens); + set_tool_content_text( + msg, + &format!("{}\n\n[输出已截断,原始约 {} tokens]", truncated, tokens), + ); + } + } +} + +/// 将文本截断到大约指定的 token 数 +fn truncate_to_tokens(text: &str, max_tokens: usize) -> String { + let mut current_tokens = 0.0f64; + let max = max_tokens as f64; + let mut last_valid_idx = 0; + for (idx, ch) in text.char_indices() { + let char_tokens = if is_cjk(ch) { 1.5 } else { 0.25 }; + current_tokens += char_tokens; + if current_tokens >= max { + break; + } + last_valid_idx = idx + ch.len_utf8(); + } + text[..last_valid_idx].to_string() +} + +/// 检查消息是否为工具结果 +fn is_tool_result(msg: &serde_json::Value) -> bool { + msg.get("role").and_then(|r| r.as_str()) == Some("user") + && msg.get("content").map_or(false, |c| { + if let Some(arr) = c.as_array() { + arr.iter() + .any(|item| item.get("type").and_then(|t| t.as_str()) == Some("tool_result")) + } else { + false + } + }) +} + +/// 提取工具结果消息的文本内容 +fn extract_tool_content_text(msg: &serde_json::Value) -> Option { + if let Some(content) = msg.get("content") { + if let Some(s) = content.as_str() { + return Some(s.to_string()); + } + if let Some(arr) = content.as_array() { + let texts: Vec<&str> = arr + .iter() + .filter_map(|item| { + if item.get("type").and_then(|t| t.as_str()) == Some("tool_result") { + item.get("content").and_then(|c| c.as_str()) + } else if item.get("type").and_then(|t| t.as_str()) == Some("text") { + item.get("text").and_then(|t| t.as_str()) + } else { + None + } + }) + .collect(); + if !texts.is_empty() { + return Some(texts.join("\n")); + } + } + } + None +} + +/// 设置工具结果消息的文本内容 +fn set_tool_content_text(msg: &mut serde_json::Value, text: &str) { + if let Some(content) = msg.get_mut("content") { + if content.is_string() { + *content = serde_json::Value::String(text.to_string()); + } else if let Some(arr) = content.as_array_mut() { + for item in arr.iter_mut() { + if item.get("type").and_then(|t| t.as_str()) == Some("tool_result") { + if let Some(c) = item.get_mut("content") { + *c = serde_json::Value::String(text.to_string()); + } + } + } + } + } +} + #[cfg(test)] mod tests { use super::*; @@ -216,10 +444,18 @@ mod tests { #[test] fn test_default_config() { let config = SummaryConfig::default(); - assert!(!config.enabled); + assert!(config.enabled); assert_eq!(config.threshold_messages, 50); assert_eq!(config.keep_recent_messages, 20); assert_eq!(config.max_summary_points, 12); + assert!(config.preserve_system_messages); + assert!(config.summarize_tool_results); + assert_eq!(config.keep_recent_turns, 10); + assert_eq!(config.token_threshold, Some(80000)); + } + + fn make_messages(n: usize) -> Vec { + (0..n).map(|i| json!({"role": if i % 2 == 0 { "user" } else { "assistant" }, "content": format!("msg {i}")})).collect() } #[test] @@ -229,7 +465,7 @@ mod tests { threshold_messages: 5, ..Default::default() }); - assert!(!s.should_summarize(100)); + assert!(!s.should_summarize(&make_messages(100))); } #[test] @@ -239,8 +475,8 @@ mod tests { threshold_messages: 50, ..Default::default() }); - assert!(!s.should_summarize(30)); - assert!(!s.should_summarize(50)); // 等于阈值不触发 + assert!(!s.should_summarize(&make_messages(30))); + assert!(!s.should_summarize(&make_messages(50))); // 等于阈值不触发 } #[test] @@ -250,8 +486,8 @@ mod tests { threshold_messages: 5, ..Default::default() }); - assert!(s.should_summarize(6)); - assert!(s.should_summarize(100)); + assert!(s.should_summarize(&make_messages(6))); + assert!(s.should_summarize(&make_messages(100))); } #[test] @@ -277,6 +513,7 @@ mod tests { threshold_messages: 2, keep_recent_messages: 1, max_summary_points: 5, + ..Default::default() }); let msgs = vec![ json!({"role": "system", "content": "You are helpful."}), @@ -375,4 +612,146 @@ mod tests { let msg3 = json!({"role": "user", "content": 42}); assert_eq!(extract_content_text(&msg3), ""); } + + #[test] + fn test_extract_tool_name_anthropic() { + let msg = json!({ + "role": "assistant", + "content": [ + {"type": "tool_use", "name": "read_file", "id": "t1", "input": {}} + ] + }); + assert_eq!(extract_tool_name(&msg), Some("read_file".to_string())); + } + + #[test] + fn test_extract_tool_name_openai() { + let msg = json!({ + "role": "assistant", + "tool_calls": [ + {"id": "t1", "type": "function", "function": {"name": "grep", "arguments": "{}"}} + ] + }); + assert_eq!(extract_tool_name(&msg), Some("grep".to_string())); + } + + #[test] + fn test_extract_tool_name_none() { + let msg = json!({"role": "user", "content": "hello"}); + assert_eq!(extract_tool_name(&msg), None); + } + + #[test] + fn test_summarize_subagent_result_short() { + let result = "short result"; + assert_eq!(summarize_subagent_result(result, 100), "short result"); + } + + #[test] + fn test_summarize_subagent_result_long() { + let result = "a".repeat(500); + let summary = summarize_subagent_result(&result, 200); + assert!(summary.len() < 500); + assert!(summary.contains("省略")); + assert!(summary.contains("300 字符")); + } + + #[test] + fn test_tool_result_compact_format() { + let s = ConversationSummarizer::new(SummaryConfig { + enabled: true, + threshold_messages: 2, + keep_recent_messages: 1, + summarize_tool_results: true, + ..Default::default() + }); + let msgs = vec![ + json!({"role": "assistant", "content": [ + {"type": "tool_use", "name": "bash", "id": "t1", "input": {}} + ]}), + json!({"role": "user", "content": "old msg"}), + json!({"role": "assistant", "content": "old reply"}), + json!({"role": "user", "content": "recent"}), + ]; + let req = s.build_summary_request(&msgs).unwrap(); + assert!(req.messages_to_summarize.contains("[tool:bash]")); + } + + #[test] + fn test_estimate_tokens_english() { + let text = "Hello world this is a test"; + let tokens = estimate_tokens(text); + assert!(tokens > 0); + assert!(tokens < 30); // 26 chars * 0.25 ≈ 6-7 + } + + #[test] + fn test_estimate_tokens_chinese() { + let text = "你好世界这是测试"; + let tokens = estimate_tokens(text); + assert!(tokens >= 8); // 8 CJK chars * 1.5 = 12 + } + + #[test] + fn test_estimate_tokens_mixed() { + let text = "Hello 你好 World 世界"; + let tokens = estimate_tokens(text); + assert!(tokens > 0); + } + + #[test] + fn test_microcompact_truncates_long_tool_output() { + let long_output = "x".repeat(10000); + let mut messages = vec![json!({ + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "t1", + "content": long_output + } + ] + })]; + microcompact(&mut messages, 100); + let content = extract_tool_content_text(&messages[0]).unwrap(); + assert!(content.contains("输出已截断")); + assert!(content.len() < long_output.len()); + } + + #[test] + fn test_microcompact_preserves_short_output() { + let short_output = "short result"; + let mut messages = vec![json!({ + "role": "user", + "content": [ + { + "type": "tool_result", + "tool_use_id": "t1", + "content": short_output + } + ] + })]; + microcompact(&mut messages, 1000); + let content = extract_tool_content_text(&messages[0]).unwrap(); + assert_eq!(content, short_output); + } + + #[test] + fn test_truncate_to_tokens() { + let text = "a".repeat(1000); + let truncated = truncate_to_tokens(&text, 100); + assert!(truncated.len() < 1000); + } + + #[test] + fn test_should_summarize_token_threshold() { + let s = ConversationSummarizer::new(SummaryConfig { + enabled: true, + threshold_messages: 1000, // 高消息阈值 + token_threshold: Some(10), // 低 token 阈值 + ..Default::default() + }); + let msgs = vec![json!({"role": "user", "content": "a".repeat(100)})]; + assert!(s.should_summarize(&msgs)); + } } diff --git a/src-tauri/crates/processor/src/steps/mod.rs b/src-tauri/crates/processor/src/steps/mod.rs index 47b119d8a..1d8f653e4 100644 --- a/src-tauri/crates/processor/src/steps/mod.rs +++ b/src-tauri/crates/processor/src/steps/mod.rs @@ -6,6 +6,7 @@ mod auth; mod injection; mod plugin; mod provider; +pub mod registry; mod routing; mod telemetry; mod traits; diff --git a/src-tauri/crates/processor/src/steps/registry.rs b/src-tauri/crates/processor/src/steps/registry.rs new file mode 100644 index 000000000..6f582e414 --- /dev/null +++ b/src-tauri/crates/processor/src/steps/registry.rs @@ -0,0 +1,290 @@ +//! 动态 Pipeline 步骤注册表 +//! +//! 允许在运行时注册、移除自定义 Pipeline 步骤,并按阶段和优先级排序。 + +use super::traits::PipelineStep; + +/// Pipeline 阶段 +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub enum PipelinePhase { + /// 在指定步骤之前执行 + Before(String), + /// 在指定步骤之后执行 + After(String), + /// 替换指定步骤 + Replace(String), + /// 在所有步骤之前 + First, + /// 在所有步骤之后 + Last, +} + +/// 注册的步骤条目 +struct RegisteredStep { + step: Box, + phase: PipelinePhase, + priority: i32, +} + +/// Pipeline 步骤注册表 +pub struct StepRegistry { + core_steps: Vec>, + dynamic_steps: Vec, +} + +impl StepRegistry { + pub fn new(core_steps: Vec>) -> Self { + Self { + core_steps, + dynamic_steps: Vec::new(), + } + } + + /// 运行时注册自定义步骤 + pub fn register(&mut self, step: Box, phase: PipelinePhase, priority: i32) { + self.dynamic_steps.push(RegisteredStep { + step, + phase, + priority, + }); + } + + /// 移除动态注册的步骤 + pub fn unregister(&mut self, step_name: &str) -> bool { + let before = self.dynamic_steps.len(); + self.dynamic_steps.retain(|s| s.step.name() != step_name); + self.dynamic_steps.len() < before + } + + /// 按正确顺序返回所有步骤的引用 + pub fn ordered_steps(&self) -> Vec<&dyn PipelineStep> { + // 收集被替换的核心步骤名 + let replaced: std::collections::HashSet<&str> = self + .dynamic_steps + .iter() + .filter_map(|s| match &s.phase { + PipelinePhase::Replace(name) => Some(name.as_str()), + _ => None, + }) + .collect(); + + let mut result: Vec<&dyn PipelineStep> = Vec::new(); + + // First 阶段(按 priority 排序) + let mut firsts: Vec<&RegisteredStep> = self + .dynamic_steps + .iter() + .filter(|s| s.phase == PipelinePhase::First) + .collect(); + firsts.sort_by_key(|s| s.priority); + result.extend(firsts.iter().map(|s| s.step.as_ref())); + + // 核心步骤 + Before/After/Replace + for core in &self.core_steps { + let core_name = core.name(); + + // Before 此核心步骤的动态步骤 + let mut befores: Vec<&RegisteredStep> = self + .dynamic_steps + .iter() + .filter(|s| matches!(&s.phase, PipelinePhase::Before(n) if n == core_name)) + .collect(); + befores.sort_by_key(|s| s.priority); + result.extend(befores.iter().map(|s| s.step.as_ref())); + + if replaced.contains(core_name) { + // 用替换步骤代替核心步骤 + let mut replacements: Vec<&RegisteredStep> = self + .dynamic_steps + .iter() + .filter(|s| matches!(&s.phase, PipelinePhase::Replace(n) if n == core_name)) + .collect(); + replacements.sort_by_key(|s| s.priority); + result.extend(replacements.iter().map(|s| s.step.as_ref())); + } else { + result.push(core.as_ref()); + } + + // After 此核心步骤的动态步骤 + let mut afters: Vec<&RegisteredStep> = self + .dynamic_steps + .iter() + .filter(|s| matches!(&s.phase, PipelinePhase::After(n) if n == core_name)) + .collect(); + afters.sort_by_key(|s| s.priority); + result.extend(afters.iter().map(|s| s.step.as_ref())); + } + + // Last 阶段 + let mut lasts: Vec<&RegisteredStep> = self + .dynamic_steps + .iter() + .filter(|s| s.phase == PipelinePhase::Last) + .collect(); + lasts.sort_by_key(|s| s.priority); + result.extend(lasts.iter().map(|s| s.step.as_ref())); + + result + } +} + +#[cfg(test)] +mod tests { + use super::super::traits::StepError; + use super::*; + use async_trait::async_trait; + use proxycast_core::processor::RequestContext; + + struct DummyStep { + name: String, + } + + impl DummyStep { + fn new(name: &str) -> Self { + Self { + name: name.to_string(), + } + } + + fn boxed(name: &str) -> Box { + Box::new(Self::new(name)) + } + } + + #[async_trait] + impl PipelineStep for DummyStep { + async fn execute( + &self, + _ctx: &mut RequestContext, + _payload: &mut serde_json::Value, + ) -> Result<(), StepError> { + Ok(()) + } + + fn name(&self) -> &str { + &self.name + } + } + + fn step_names(registry: &StepRegistry) -> Vec { + registry + .ordered_steps() + .iter() + .map(|s| s.name().to_string()) + .collect() + } + + #[test] + fn test_core_steps_only() { + let registry = StepRegistry::new(vec![ + DummyStep::boxed("auth"), + DummyStep::boxed("routing"), + DummyStep::boxed("provider"), + ]); + assert_eq!(step_names(®istry), vec!["auth", "routing", "provider"]); + } + + #[test] + fn test_first_and_last() { + let mut registry = StepRegistry::new(vec![DummyStep::boxed("core")]); + registry.register(DummyStep::boxed("first_step"), PipelinePhase::First, 0); + registry.register(DummyStep::boxed("last_step"), PipelinePhase::Last, 0); + assert_eq!( + step_names(®istry), + vec!["first_step", "core", "last_step"] + ); + } + + #[test] + fn test_before_and_after() { + let mut registry = + StepRegistry::new(vec![DummyStep::boxed("auth"), DummyStep::boxed("provider")]); + registry.register( + DummyStep::boxed("pre_auth"), + PipelinePhase::Before("auth".to_string()), + 0, + ); + registry.register( + DummyStep::boxed("post_auth"), + PipelinePhase::After("auth".to_string()), + 0, + ); + assert_eq!( + step_names(®istry), + vec!["pre_auth", "auth", "post_auth", "provider"] + ); + } + + #[test] + fn test_replace() { + let mut registry = + StepRegistry::new(vec![DummyStep::boxed("auth"), DummyStep::boxed("provider")]); + registry.register( + DummyStep::boxed("custom_auth"), + PipelinePhase::Replace("auth".to_string()), + 0, + ); + assert_eq!(step_names(®istry), vec!["custom_auth", "provider"]); + } + + #[test] + fn test_unregister() { + let mut registry = StepRegistry::new(vec![DummyStep::boxed("core")]); + registry.register(DummyStep::boxed("extra"), PipelinePhase::Last, 0); + assert_eq!(step_names(®istry), vec!["core", "extra"]); + + assert!(registry.unregister("extra")); + assert_eq!(step_names(®istry), vec!["core"]); + + // 移除不存在的步骤返回 false + assert!(!registry.unregister("nonexistent")); + } + + #[test] + fn test_priority_ordering() { + let mut registry = StepRegistry::new(vec![DummyStep::boxed("core")]); + registry.register(DummyStep::boxed("low"), PipelinePhase::First, 10); + registry.register(DummyStep::boxed("high"), PipelinePhase::First, 1); + // priority 小的排前面 + assert_eq!(step_names(®istry), vec!["high", "low", "core"]); + } + + #[test] + fn test_complex_pipeline() { + let mut registry = StepRegistry::new(vec![ + DummyStep::boxed("auth"), + DummyStep::boxed("routing"), + DummyStep::boxed("provider"), + ]); + registry.register(DummyStep::boxed("init"), PipelinePhase::First, 0); + registry.register( + DummyStep::boxed("rate_limit"), + PipelinePhase::Before("auth".to_string()), + 0, + ); + registry.register( + DummyStep::boxed("log_auth"), + PipelinePhase::After("auth".to_string()), + 0, + ); + registry.register( + DummyStep::boxed("custom_routing"), + PipelinePhase::Replace("routing".to_string()), + 0, + ); + registry.register(DummyStep::boxed("telemetry"), PipelinePhase::Last, 0); + + assert_eq!( + step_names(®istry), + vec![ + "init", + "rate_limit", + "auth", + "log_auth", + "custom_routing", + "provider", + "telemetry" + ] + ); + } +} diff --git a/src-tauri/crates/services/src/context_memory_service.rs b/src-tauri/crates/services/src/context_memory_service.rs index 06b664f4c..657771a21 100644 --- a/src-tauri/crates/services/src/context_memory_service.rs +++ b/src-tauri/crates/services/src/context_memory_service.rs @@ -3,8 +3,9 @@ //! 基于文件系统的持久化记忆系统,解决 AI Agent 的上下文丢失、目标漂移、错误重复问题 //! 核心理念:Context Window = RAM, Filesystem = Disk +use chrono::TimeZone; use serde::{Deserialize, Serialize}; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::fs; use std::path::PathBuf; use std::sync::{Arc, Mutex}; @@ -79,6 +80,8 @@ pub struct ContextMemoryConfig { pub max_entries_per_session: usize, /// 自动归档天数 pub auto_archive_days: u32, + /// 是否启用自动清理 + pub auto_cleanup_enabled: bool, /// 启用错误跟踪 pub enable_error_tracking: bool, /// 最大错误重试次数 @@ -92,6 +95,7 @@ impl Default for ContextMemoryConfig { memory_dir: home_dir.join(".proxycast").join("memory"), max_entries_per_session: 100, auto_archive_days: 30, + auto_cleanup_enabled: true, enable_error_tracking: true, max_error_retries: 3, } @@ -195,6 +199,9 @@ impl ContextMemoryService { session_id: &str, file_type: MemoryFileType, ) -> Result<(), String> { + let session_dir = self.get_session_memory_dir(session_id); + fs::create_dir_all(&session_dir).map_err(|e| format!("创建会话目录失败: {e}"))?; + let file_path = self.get_memory_file_path(session_id, file_type); let cache = self.memory_cache.lock().map_err(|e| e.to_string())?; @@ -349,40 +356,42 @@ impl ContextMemoryService { return Ok(()); } - let mut error_cache = self.error_cache.lock().map_err(|e| e.to_string())?; - let errors = error_cache - .entry(session_id.to_string()) - .or_insert_with(Vec::new); - - // 查找现有错误 - if let Some(existing_error) = errors - .iter_mut() - .find(|e| e.error_description == error_description) { - existing_error - .attempted_solutions - .push(attempted_solution.to_string()); - existing_error.failure_count += 1; - existing_error.last_failure_at = chrono::Utc::now().timestamp_millis(); + let mut error_cache = self.error_cache.lock().map_err(|e| e.to_string())?; + let errors = error_cache + .entry(session_id.to_string()) + .or_insert_with(Vec::new); - warn!( - "重复错误记录 (第{}次): {} (会话: {})", - existing_error.failure_count, error_description, session_id - ); - } else { - let error_entry = ErrorEntry { - id: uuid::Uuid::new_v4().to_string(), - session_id: session_id.to_string(), - error_description: error_description.to_string(), - attempted_solutions: vec![attempted_solution.to_string()], - failure_count: 1, - last_failure_at: chrono::Utc::now().timestamp_millis(), - resolved: false, - resolution: None, - }; + // 查找现有错误 + if let Some(existing_error) = errors + .iter_mut() + .find(|e| e.error_description == error_description) + { + existing_error + .attempted_solutions + .push(attempted_solution.to_string()); + existing_error.failure_count += 1; + existing_error.last_failure_at = chrono::Utc::now().timestamp_millis(); - errors.push(error_entry); - info!("记录新错误: {} (会话: {})", error_description, session_id); + warn!( + "重复错误记录 (第{}次): {} (会话: {})", + existing_error.failure_count, error_description, session_id + ); + } else { + let error_entry = ErrorEntry { + id: uuid::Uuid::new_v4().to_string(), + session_id: session_id.to_string(), + error_description: error_description.to_string(), + attempted_solutions: vec![attempted_solution.to_string()], + failure_count: 1, + last_failure_at: chrono::Utc::now().timestamp_millis(), + resolved: false, + resolution: None, + }; + + errors.push(error_entry); + info!("记录新错误: {} (会话: {})", error_description, session_id); + } } // 保存到文件 @@ -516,6 +525,34 @@ impl ContextMemoryService { /// 加载会话记忆 fn load_session_memories(&self, session_id: &str) -> Result<(), String> { + let mut loaded_entries = Vec::new(); + + for file_type in [ + MemoryFileType::TaskPlan, + MemoryFileType::Findings, + MemoryFileType::Progress, + ] { + let file_path = self.get_memory_file_path(session_id, file_type); + if !file_path.exists() { + continue; + } + + match fs::read_to_string(&file_path) { + Ok(content) => { + let mut parsed = self.parse_markdown_entries(session_id, file_type, &content); + loaded_entries.append(&mut parsed); + } + Err(err) => { + warn!("读取记忆文件失败: {} - {}", file_path.display(), err); + } + } + } + + if !loaded_entries.is_empty() { + let mut memory_cache = self.memory_cache.lock().map_err(|e| e.to_string())?; + memory_cache.insert(session_id.to_string(), loaded_entries); + } + // 加载错误日志 let error_file = self.get_memory_file_path(session_id, MemoryFileType::ErrorLog); if error_file.exists() { @@ -531,22 +568,217 @@ impl ContextMemoryService { Ok(()) } + fn parse_markdown_entries( + &self, + session_id: &str, + file_type: MemoryFileType, + content: &str, + ) -> Vec { + let mut entries = Vec::new(); + let mut current_title: Option = None; + let mut section_lines: Vec = Vec::new(); + let mut index = 0usize; + + for line in content.lines() { + if let Some(title) = line.strip_prefix("## ") { + if let Some(previous_title) = current_title.take() { + if let Some(entry) = self.build_memory_entry( + session_id, + file_type, + index, + &previous_title, + §ion_lines, + ) { + entries.push(entry); + index += 1; + } + } + + current_title = Some(title.trim().to_string()); + section_lines.clear(); + continue; + } + + if current_title.is_some() { + section_lines.push(line.to_string()); + } + } + + if let Some(previous_title) = current_title { + if let Some(entry) = self.build_memory_entry( + session_id, + file_type, + index, + &previous_title, + §ion_lines, + ) { + entries.push(entry); + } + } + + entries + } + + fn build_memory_entry( + &self, + session_id: &str, + file_type: MemoryFileType, + index: usize, + title: &str, + lines: &[String], + ) -> Option { + let title = title.trim(); + if title.is_empty() { + return None; + } + + let file_type_key = match file_type { + MemoryFileType::TaskPlan => "task_plan", + MemoryFileType::Findings => "findings", + MemoryFileType::Progress => "progress", + MemoryFileType::ErrorLog => "error_log", + }; + + let (priority, tags, parsed_updated_at) = self.parse_entry_metadata(lines); + let now = chrono::Utc::now().timestamp_millis(); + let updated_at = if parsed_updated_at > 0 { + parsed_updated_at + } else { + now + }; + + let content = lines + .iter() + .map(|line| line.trim_end()) + .filter(|line| { + let trimmed = line.trim(); + !trimmed.is_empty() + && !trimmed.starts_with("**优先级**:") + && trimmed != "---" + && trimmed != "----" + }) + .collect::>() + .join("\n") + .trim() + .to_string(); + + Some(MemoryEntry { + id: format!("{session_id}:{file_type_key}:{index}"), + session_id: session_id.to_string(), + file_type, + title: title.to_string(), + content: if content.is_empty() { + "暂无内容".to_string() + } else { + content + }, + tags, + priority, + created_at: updated_at, + updated_at, + archived: false, + }) + } + + fn parse_entry_metadata(&self, lines: &[String]) -> (u8, Vec, i64) { + for line in lines { + let line = line.trim(); + if !line.starts_with("**优先级**:") { + continue; + } + + let priority = line + .split("**优先级**:") + .nth(1) + .and_then(|part| part.split('|').next()) + .and_then(|part| part.trim().parse::().ok()) + .map(|value| value.clamp(1, 5)) + .unwrap_or(3); + + let tags = line + .split("**标签**:") + .nth(1) + .and_then(|part| part.split("| **更新时间**").next()) + .map(|part| { + part.split(',') + .map(|tag| tag.trim().to_string()) + .filter(|tag| !tag.is_empty()) + .collect::>() + }) + .unwrap_or_default(); + + let updated_at = line + .split("**更新时间**:") + .nth(1) + .map(str::trim) + .and_then(Self::parse_datetime_or_timestamp_to_millis) + .unwrap_or(0); + + return (priority, tags, updated_at); + } + + (3, Vec::new(), 0) + } + + fn parse_datetime_or_timestamp_to_millis(value: &str) -> Option { + if let Ok(v) = value.parse::() { + if v > 1_000_000_000_000 { + return Some(v); + } + return Some(v * 1000); + } + + chrono::NaiveDateTime::parse_from_str(value, "%Y-%m-%d %H:%M:%S") + .ok() + .and_then(|naive| { + chrono::Local + .from_local_datetime(&naive) + .single() + .map(|dt| dt.timestamp_millis()) + }) + } + /// 清理过期记忆 pub fn cleanup_expired_memories(&self) -> Result<(), String> { + if !self.config.auto_cleanup_enabled { + debug!("自动清理已关闭,跳过过期记忆清理"); + return Ok(()); + } + + self.cleanup_expired_memories_with_retention_days(self.config.auto_archive_days) + } + + /// 按保留天数清理过期记忆 + pub fn cleanup_expired_memories_with_retention_days( + &self, + retention_days: u32, + ) -> Result<(), String> { let cutoff_time = chrono::Utc::now().timestamp_millis() - - (self.config.auto_archive_days as i64 * 24 * 60 * 60 * 1000); + - (retention_days.max(1) as i64 * 24 * 60 * 60 * 1000); let mut memory_cache = self.memory_cache.lock().map_err(|e| e.to_string())?; let mut archived_count = 0; + let mut dirty_files: HashMap> = HashMap::new(); - for entries in memory_cache.values_mut() { + for (session_id, entries) in memory_cache.iter_mut() { for entry in entries.iter_mut() { if entry.updated_at < cutoff_time && !entry.archived { entry.archived = true; archived_count += 1; + dirty_files + .entry(session_id.clone()) + .or_default() + .insert(entry.file_type); } } } + drop(memory_cache); + + for (session_id, file_types) in dirty_files { + for file_type in file_types { + self.save_memory_to_file(&session_id, file_type)?; + } + } if archived_count > 0 { info!("已归档 {} 个过期记忆条目", archived_count); @@ -613,6 +845,7 @@ mod tests { memory_dir: temp_dir.path().to_path_buf(), max_entries_per_session: 10, auto_archive_days: 1, + auto_cleanup_enabled: true, enable_error_tracking: true, max_error_retries: 3, }; @@ -748,4 +981,69 @@ mod tests { Some(&1) ); } + + #[test] + fn test_reload_markdown_memories_into_cache() { + let (config, _temp_dir) = create_test_config(); + let session_id = "reload-session"; + + let first_service = ContextMemoryService::new(config.clone()).unwrap(); + let entry = MemoryEntry { + id: "reload-entry".to_string(), + session_id: session_id.to_string(), + file_type: MemoryFileType::TaskPlan, + title: "重启恢复测试".to_string(), + content: "验证 markdown 能否在启动时恢复到缓存".to_string(), + tags: vec!["reload".to_string()], + priority: 4, + created_at: chrono::Utc::now().timestamp_millis(), + updated_at: chrono::Utc::now().timestamp_millis(), + archived: false, + }; + first_service.save_memory_entry(&entry).unwrap(); + drop(first_service); + + let second_service = ContextMemoryService::new(config).unwrap(); + let memories = second_service + .get_session_memories(session_id, Some(MemoryFileType::TaskPlan)) + .unwrap(); + + assert_eq!(memories.len(), 1); + assert_eq!(memories[0].title, "重启恢复测试"); + } + + #[test] + fn test_cleanup_persists_to_markdown_file() { + let (config, _temp_dir) = create_test_config(); + let service = ContextMemoryService::new(config.clone()).unwrap(); + let session_id = "cleanup-session"; + + let old_timestamp = chrono::Utc::now().timestamp_millis() - 3 * 24 * 60 * 60 * 1000; + let entry = MemoryEntry { + id: "cleanup-entry".to_string(), + session_id: session_id.to_string(), + file_type: MemoryFileType::TaskPlan, + title: "应被归档的条目".to_string(), + content: "过期内容".to_string(), + tags: vec!["cleanup".to_string()], + priority: 2, + created_at: old_timestamp, + updated_at: old_timestamp, + archived: false, + }; + service.save_memory_entry(&entry).unwrap(); + + service + .cleanup_expired_memories_with_retention_days(1) + .unwrap(); + + let memories = service + .get_session_memories(session_id, Some(MemoryFileType::TaskPlan)) + .unwrap(); + assert!(memories.is_empty()); + + let task_plan_file = config.memory_dir.join(session_id).join("task_plan.md"); + let content = std::fs::read_to_string(task_plan_file).unwrap(); + assert!(!content.contains("应被归档的条目")); + } } diff --git a/src-tauri/crates/skills/src/lib.rs b/src-tauri/crates/skills/src/lib.rs index 4e782b0f3..3ad5b35db 100644 --- a/src-tauri/crates/skills/src/lib.rs +++ b/src-tauri/crates/skills/src/lib.rs @@ -7,6 +7,7 @@ mod execution_callback; mod llm_provider; mod proxycast_llm_provider; mod skill_loader; +mod skill_matcher; // 电商 Skill 模块 pub mod ecommerce_review_reply; @@ -20,5 +21,6 @@ pub use proxycast_llm_provider::ProxyCastLlmProvider; pub use skill_loader::{ find_skill_by_name, get_proxycast_skills_dir, load_skill_from_file, load_skills_from_directory, parse_allowed_tools, parse_boolean, parse_skill_frontmatter, parse_workflow_steps, - LoadedSkillDefinition, SkillFrontmatter, WorkflowStep, + LoadedSkillDefinition, SkillFrontmatter, SkillTriggerConfig, WorkflowStep, }; +pub use skill_matcher::{SkillMatch, SkillMatcher}; diff --git a/src-tauri/crates/skills/src/skill_loader.rs b/src-tauri/crates/skills/src/skill_loader.rs index cfc117e7c..4288d0dbf 100644 --- a/src-tauri/crates/skills/src/skill_loader.rs +++ b/src-tauri/crates/skills/src/skill_loader.rs @@ -6,6 +6,17 @@ use std::path::{Path, PathBuf}; use serde::{Deserialize, Serialize}; +/// Skill 自动触发条件配置 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct SkillTriggerConfig { + /// 触发条件描述列表(自然语言) + #[serde(default)] + pub trigger: Vec, + /// 不触发条件描述列表 + #[serde(default)] + pub do_not_trigger: Vec, +} + /// Workflow 步骤定义 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct WorkflowStep { @@ -61,6 +72,8 @@ pub struct LoadedSkillDefinition { pub allowed_tools: Option>, pub argument_hint: Option, pub when_to_use: Option, + /// 结构化的自动触发条件配置 + pub when_to_use_config: Option, pub model: Option, pub provider: Option, pub disable_model_invocation: bool, @@ -204,6 +217,12 @@ pub fn load_skill_from_file( execution_mode }; + // 尝试将 when_to_use 解析为 JSON 格式的 SkillTriggerConfig + let when_to_use_config = frontmatter + .when_to_use + .as_deref() + .and_then(|v| serde_json::from_str::(v).ok()); + Ok(LoadedSkillDefinition { skill_name: skill_name.to_string(), display_name, @@ -212,6 +231,7 @@ pub fn load_skill_from_file( allowed_tools, argument_hint: frontmatter.argument_hint, when_to_use: frontmatter.when_to_use, + when_to_use_config, model: frontmatter.model, provider: frontmatter.provider, disable_model_invocation, diff --git a/src-tauri/crates/skills/src/skill_matcher.rs b/src-tauri/crates/skills/src/skill_matcher.rs new file mode 100644 index 000000000..2143c8689 --- /dev/null +++ b/src-tauri/crates/skills/src/skill_matcher.rs @@ -0,0 +1,390 @@ +//! Skill 自动匹配器 +//! +//! 基于关键词的简单匹配,不依赖 LLM。 +//! 从 `SkillTriggerConfig` 的 trigger/do_not_trigger 列表提取关键词进行模糊匹配。 + +use crate::skill_loader::LoadedSkillDefinition; + +/// Skill 匹配结果 +#[derive(Debug, Clone)] +pub struct SkillMatch { + pub skill_name: String, + pub confidence: f32, + pub trigger_reason: String, +} + +/// 基于关键词的简单匹配器 +pub struct SkillMatcher { + skills: Vec, +} + +/// 最低置信度阈值 +const CONFIDENCE_THRESHOLD: f32 = 0.6; + +impl SkillMatcher { + pub fn new(skills: Vec) -> Self { + Self { skills } + } + + /// 根据用户输入匹配最合适的 Skill + /// 返回按 confidence 降序排列的匹配结果(仅 >= 0.6) + pub fn match_skills(&self, user_input: &str) -> Vec { + let input_lower = user_input.to_lowercase(); + let mut matches = Vec::new(); + + for skill in &self.skills { + let config = match &skill.when_to_use_config { + Some(c) => c, + None => continue, + }; + + if config.trigger.is_empty() { + continue; + } + + // 先检查排除条件 + if self.check_exclusions(&input_lower, &config.do_not_trigger) { + continue; + } + + // 检查触发条件 + let (matched, confidence, reason) = self.check_triggers(&input_lower, &config.trigger); + + if matched && confidence >= CONFIDENCE_THRESHOLD { + matches.push(SkillMatch { + skill_name: skill.skill_name.clone(), + confidence, + trigger_reason: reason, + }); + } + } + + matches.sort_by(|a, b| { + b.confidence + .partial_cmp(&a.confidence) + .unwrap_or(std::cmp::Ordering::Equal) + }); + matches + } + + /// 检查用户输入是否包含触发关键词 + /// 返回 (是否匹配, 置信度, 匹配原因) + fn check_triggers(&self, input: &str, triggers: &[String]) -> (bool, f32, String) { + if triggers.is_empty() { + return (false, 0.0, String::new()); + } + + let mut matched_triggers = Vec::new(); + + for trigger in triggers { + let keywords = extract_keywords(trigger); + if keywords.is_empty() { + continue; + } + + let matched_count = keywords + .iter() + .filter(|kw| input.contains(kw.as_str())) + .count(); + + if matched_count > 0 { + let ratio = matched_count as f32 / keywords.len() as f32; + if ratio >= 0.5 { + matched_triggers.push((trigger.clone(), ratio)); + } + } + } + + if matched_triggers.is_empty() { + return (false, 0.0, String::new()); + } + + // 置信度 = 匹配的 trigger 条目占比 * 最佳单条匹配率 + let best_ratio = matched_triggers + .iter() + .map(|(_, r)| *r) + .fold(0.0f32, f32::max); + let trigger_coverage = matched_triggers.len() as f32 / triggers.len() as f32; + let confidence = (best_ratio * 0.7 + trigger_coverage * 0.3).min(1.0); + + let reasons: Vec = matched_triggers.iter().map(|(t, _)| t.clone()).collect(); + let reason = format!("匹配触发条件: {}", reasons.join(", ")); + + (true, confidence, reason) + } + + /// 检查是否命中排除条件 + fn check_exclusions(&self, input: &str, exclusions: &[String]) -> bool { + for exclusion in exclusions { + let keywords = extract_keywords(exclusion); + if keywords.is_empty() { + continue; + } + + let matched_count = keywords + .iter() + .filter(|kw| input.contains(kw.as_str())) + .count(); + + // 排除条件中超过一半关键词命中即排除 + if matched_count > 0 && matched_count as f32 / keywords.len() as f32 >= 0.5 { + return true; + } + } + false + } + + /// 生成 skill 描述文本,用于注入 system prompt + pub fn generate_skill_prompt_section(&self) -> String { + if self.skills.is_empty() { + return String::new(); + } + + let mut section = String::from("## 可用 Skills\n\n"); + for skill in &self.skills { + section.push_str(&format!("### /{}\n", skill.skill_name)); + if !skill.description.is_empty() { + section.push_str(&skill.description); + section.push('\n'); + } + if let Some(ref config) = skill.when_to_use_config { + if !config.trigger.is_empty() { + section.push_str("触发条件:"); + section.push_str(&config.trigger.join("、")); + section.push('\n'); + } + if !config.do_not_trigger.is_empty() { + section.push_str("不触发:"); + section.push_str(&config.do_not_trigger.join("、")); + section.push('\n'); + } + } + section.push('\n'); + } + section + } +} + +/// 从自然语言描述中提取关键词(小写) +/// 过滤掉常见停用词,保留有意义的词汇 +fn extract_keywords(text: &str) -> Vec { + // 中英文停用词 + const STOP_WORDS: &[&str] = &[ + // 英文 + "a", "an", "the", "is", "are", "was", "were", "be", "been", "being", "have", "has", "had", + "do", "does", "did", "will", "would", "could", "should", "may", "might", "can", "shall", + "to", "of", "in", "for", "on", "with", "at", "by", "from", "as", "into", "about", "like", + "through", "after", "over", "between", "out", "against", "during", "without", "before", + "under", "around", "among", "and", "but", "or", "nor", "not", "so", "yet", "both", + "either", "neither", "each", "every", "all", "any", "few", "more", "most", "other", "some", + "such", "no", "only", "own", "same", "than", "too", "very", "just", "because", "if", + "when", "where", "how", "what", "which", "who", "whom", "this", "that", "these", "those", + "i", "me", "my", "we", "our", "you", "your", "he", "him", "his", "she", "her", "it", "its", + "they", "them", "their", "user", "want", "wants", "need", "needs", "use", "using", + // 中文 + "的", "了", "在", "是", "我", "有", "和", "就", "不", "人", "都", "一", "一个", "上", "也", + "很", "到", "说", "要", "去", "你", "会", "着", "没有", "看", "好", "自己", "这", "他", + "她", "它", "们", "那", "里", "后", "把", "让", "从", "被", "与", "对", "当", "用", "使用", + "进行", "可以", "需要", "想要", "帮我", "请", "能", "能够", + ]; + + let lower = text.to_lowercase(); + + // 按空格和常见标点分词 + let tokens: Vec = lower + .split(|c: char| c.is_whitespace() || ",.;:!?()[]{}\"'`~@#$%^&*+=|/<>".contains(c)) + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + .collect(); + + // 对于中文文本(没有空格分隔),如果 token 长度 > 4 字符且包含中文, + // 按 2-3 字符切分为子词 + let mut keywords = Vec::new(); + for token in &tokens { + let has_cjk = token.chars().any(|c| is_cjk(c)); + let char_count = token.chars().count(); + + if has_cjk && char_count > 3 { + // 中文长词切分为 bigram + let chars: Vec = token.chars().collect(); + for window in chars.windows(2) { + let bigram: String = window.iter().collect(); + if !STOP_WORDS.contains(&bigram.as_str()) { + keywords.push(bigram); + } + } + } else if !STOP_WORDS.contains(&token.as_str()) && token.len() > 1 { + keywords.push(token.clone()); + } + } + + keywords.sort(); + keywords.dedup(); + keywords +} + +/// 判断字符是否为 CJK 字符 +fn is_cjk(c: char) -> bool { + matches!(c, + '\u{4E00}'..='\u{9FFF}' | // CJK Unified Ideographs + '\u{3400}'..='\u{4DBF}' | // CJK Unified Ideographs Extension A + '\u{F900}'..='\u{FAFF}' // CJK Compatibility Ideographs + ) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::skill_loader::SkillTriggerConfig; + + fn make_skill( + name: &str, + trigger: Vec<&str>, + do_not_trigger: Vec<&str>, + ) -> LoadedSkillDefinition { + LoadedSkillDefinition { + skill_name: name.to_string(), + display_name: name.to_string(), + description: String::new(), + markdown_content: String::new(), + allowed_tools: None, + argument_hint: None, + when_to_use: None, + when_to_use_config: Some(SkillTriggerConfig { + trigger: trigger.into_iter().map(String::from).collect(), + do_not_trigger: do_not_trigger.into_iter().map(String::from).collect(), + }), + model: None, + provider: None, + disable_model_invocation: false, + execution_mode: "prompt".to_string(), + workflow_steps: Vec::new(), + } + } + + #[test] + fn test_basic_trigger_match() { + let skills = vec![make_skill( + "code-review", + vec!["review code", "code review"], + vec![], + )]; + let matcher = SkillMatcher::new(skills); + let results = matcher.match_skills("please review my code"); + assert!(!results.is_empty()); + assert_eq!(results[0].skill_name, "code-review"); + assert!(results[0].confidence >= CONFIDENCE_THRESHOLD); + } + + #[test] + fn test_no_match_below_threshold() { + let skills = vec![make_skill( + "deploy", + vec!["deploy to production server"], + vec![], + )]; + let matcher = SkillMatcher::new(skills); + let results = matcher.match_skills("hello world"); + assert!(results.is_empty()); + } + + #[test] + fn test_exclusion_prevents_match() { + let skills = vec![make_skill( + "translate", + vec!["translate text", "translation"], + vec!["translate code variable names"], + )]; + let matcher = SkillMatcher::new(skills); + let results = matcher.match_skills("translate code variable names to english"); + assert!(results.is_empty()); + } + + #[test] + fn test_multiple_skills_sorted_by_confidence() { + let skills = vec![ + make_skill("git-commit", vec!["commit changes", "git commit"], vec![]), + make_skill( + "code-review", + vec!["review code", "code review", "check code quality"], + vec![], + ), + ]; + let matcher = SkillMatcher::new(skills); + let results = matcher.match_skills("review code quality and commit"); + // code-review 应该有更高的 confidence(匹配了更多 trigger) + assert!(results.len() >= 1); + } + + #[test] + fn test_skill_without_config_is_skipped() { + let mut skill = make_skill("no-config", vec![], vec![]); + skill.when_to_use_config = None; + let matcher = SkillMatcher::new(vec![skill]); + let results = matcher.match_skills("anything"); + assert!(results.is_empty()); + } + + #[test] + fn test_chinese_trigger_match() { + let skills = vec![make_skill( + "ecommerce-reply", + vec!["电商评论回复", "商品评价回复"], + vec![], + )]; + let matcher = SkillMatcher::new(skills); + let results = matcher.match_skills("帮我生成电商评论回复"); + assert!(!results.is_empty()); + assert_eq!(results[0].skill_name, "ecommerce-reply"); + } + + #[test] + fn test_extract_keywords_english() { + let keywords = extract_keywords("review the code quality"); + assert!(keywords.contains(&"review".to_string())); + assert!(keywords.contains(&"code".to_string())); + assert!(keywords.contains(&"quality".to_string())); + // "the" 是停用词,应被过滤 + assert!(!keywords.contains(&"the".to_string())); + } + + #[test] + fn test_extract_keywords_chinese() { + let keywords = extract_keywords("电商评论回复"); + // 应该产生 bigram + assert!(!keywords.is_empty()); + assert!(keywords.contains(&"电商".to_string())); + assert!(keywords.contains(&"评论".to_string())); + } + + #[test] + fn test_empty_triggers() { + let skills = vec![make_skill("empty", vec![], vec![])]; + let matcher = SkillMatcher::new(skills); + let results = matcher.match_skills("anything"); + assert!(results.is_empty()); + } + + #[test] + fn test_generate_skill_prompt_section_empty() { + let matcher = SkillMatcher::new(vec![]); + assert_eq!(matcher.generate_skill_prompt_section(), ""); + } + + #[test] + fn test_generate_skill_prompt_section() { + let skills = vec![make_skill( + "code-review", + vec!["review code", "代码审查"], + vec!["不要自动修复"], + )]; + let matcher = SkillMatcher::new(skills); + let section = matcher.generate_skill_prompt_section(); + assert!(section.contains("## 可用 Skills")); + assert!(section.contains("### /code-review")); + assert!(section.contains("触发条件:")); + assert!(section.contains("review code")); + assert!(section.contains("不触发:")); + assert!(section.contains("不要自动修复")); + } +} diff --git a/src-tauri/src/agent/aster_agent.rs b/src-tauri/src/agent/aster_agent.rs index 53f8c3e9f..a90d80b5a 100644 --- a/src-tauri/src/agent/aster_agent.rs +++ b/src-tauri/src/agent/aster_agent.rs @@ -45,7 +45,9 @@ impl AsterAgentWrapper { let cancel_token = state.create_cancel_token(&session_id).await; let user_message = Message::user().with_text(&message); - let session_config = SessionConfigBuilder::new(&session_id).build(); + let session_config = SessionConfigBuilder::new(&session_id) + .include_context_trace(true) + .build(); let agent_arc = state.get_agent_arc(); let guard = agent_arc.read().await; diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index 612f4370d..dcdcb1aa9 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -26,5 +26,5 @@ pub use credential_bridge::{ pub use heartbeat_service_adapter::HeartbeatServiceAdapter; pub use proxycast_agent::{convert_agent_event, convert_to_tauri_message, TauriAgentEvent}; pub use subagent_scheduler::{ - ProxyCastScheduler, ProxyCastSubAgentExecutor, SubAgentProgressEvent, + ProxyCastScheduler, ProxyCastSubAgentExecutor, SubAgentProgressEvent, SubAgentRole, }; diff --git a/src-tauri/src/agent/subagent_scheduler.rs b/src-tauri/src/agent/subagent_scheduler.rs index 1e76bca65..e0c4ff8d6 100644 --- a/src-tauri/src/agent/subagent_scheduler.rs +++ b/src-tauri/src/agent/subagent_scheduler.rs @@ -14,7 +14,7 @@ use tauri::{AppHandle, Emitter}; use crate::database::DbConnection; pub use proxycast_agent::subagent_scheduler::{ - ProxyCastSubAgentExecutor, SchedulerEventEmitter, SubAgentProgressEvent, + ProxyCastSubAgentExecutor, SchedulerEventEmitter, SubAgentProgressEvent, SubAgentRole, }; /// ProxyCast SubAgent 调度器(Tauri 桥接) @@ -40,6 +40,12 @@ impl ProxyCastScheduler { self } + /// 设置默认角色 + pub fn with_default_role(mut self, role: SubAgentRole) -> Self { + self.inner = self.inner.with_default_role(role); + self + } + /// 初始化调度器 pub async fn init(&self, config: Option) { let event_emitter = self.app_handle.clone().map(|handle| { @@ -64,6 +70,18 @@ impl ProxyCastScheduler { self.inner.execute(tasks, parent_context).await } + /// 使用指定角色执行任务 + pub async fn execute_with_role( + &self, + tasks: Vec, + parent_context: Option<&AgentContext>, + role: SubAgentRole, + ) -> SchedulerResult { + self.inner + .execute_with_role(tasks, parent_context, role) + .await + } + /// 取消执行 pub async fn cancel(&self) { self.inner.cancel().await; diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index 78d6f2914..e3f00fbdb 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -233,7 +233,7 @@ pub fn init_states(config: &Config) -> Result { } // 初始化上下文记忆服务 - let context_memory_config = ContextMemoryConfig::default(); + let context_memory_config = build_context_memory_config(config); let context_memory_service = ContextMemoryService::new(context_memory_config) .map_err(|e| format!("ContextMemoryService 初始化失败: {e}"))?; let context_memory_service_arc = Arc::new(context_memory_service); @@ -291,6 +291,25 @@ pub fn init_states(config: &Config) -> Result { }) } +fn build_context_memory_config(config: &Config) -> ContextMemoryConfig { + let mut context_config = ContextMemoryConfig::default(); + let memory_config = &config.memory; + + if let Some(max_entries) = memory_config.max_entries { + context_config.max_entries_per_session = max_entries.clamp(1, 20_000) as usize; + } + + if let Some(retention_days) = memory_config.retention_days { + context_config.auto_archive_days = retention_days.clamp(1, 3650); + } + + if let Some(auto_cleanup) = memory_config.auto_cleanup { + context_config.auto_cleanup_enabled = auto_cleanup; + } + + context_config +} + /// 初始化插件安装器 fn init_plugin_installer() -> Result { let db_path = database::get_db_path().map_err(|e| format!("获取数据库路径失败: {e}"))?; diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index 0fdddf1f3..d02eae5f5 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -1341,6 +1341,10 @@ pub fn run() { commands::memory_management_cmd::get_conversation_memory_overview, commands::memory_management_cmd::request_conversation_memory_analysis, commands::memory_management_cmd::cleanup_conversation_memory, + commands::memory_management_cmd::memory_get_effective_sources, + commands::memory_management_cmd::memory_get_auto_index, + commands::memory_management_cmd::memory_toggle_auto, + commands::memory_management_cmd::memory_update_auto_note, // Unified Memory commands commands::unified_memory_cmd::unified_memory_list, commands::unified_memory_cmd::unified_memory_get, diff --git a/src-tauri/src/app/state.rs b/src-tauri/src/app/state.rs index d0f530abd..ca7071111 100644 --- a/src-tauri/src/app/state.rs +++ b/src-tauri/src/app/state.rs @@ -104,7 +104,8 @@ pub fn init_service_states() -> ServiceStates { let orchestrator_state = OrchestratorState::new(); // Initialize ContextMemoryService - let context_memory_config = ContextMemoryConfig::default(); + let app_config = proxycast_core::config::load_config().unwrap_or_default(); + let context_memory_config = build_context_memory_config(&app_config); let context_memory_service = ContextMemoryService::new(context_memory_config) .expect("Failed to initialize ContextMemoryService"); let context_memory_service_state = ContextMemoryServiceState(Arc::new(context_memory_service)); @@ -129,6 +130,25 @@ pub fn init_service_states() -> ServiceStates { } } +fn build_context_memory_config(config: &Config) -> ContextMemoryConfig { + let mut context_config = ContextMemoryConfig::default(); + let memory_config = &config.memory; + + if let Some(max_entries) = memory_config.max_entries { + context_config.max_entries_per_session = max_entries.clamp(1, 20_000) as usize; + } + + if let Some(retention_days) = memory_config.retention_days { + context_config.auto_archive_days = retention_days.clamp(1, 3650); + } + + if let Some(auto_cleanup) = memory_config.auto_cleanup { + context_config.auto_cleanup_enabled = auto_cleanup; + } + + context_config +} + /// 初始化插件安装器 fn init_plugin_installer() -> PluginInstallerState { let db_path = database::get_db_path().expect("Failed to get database path for PluginInstaller"); diff --git a/src-tauri/src/commands/agent_cmd.rs b/src-tauri/src/commands/agent_cmd.rs index 77415821b..141c36ef8 100644 --- a/src-tauri/src/commands/agent_cmd.rs +++ b/src-tauri/src/commands/agent_cmd.rs @@ -4,8 +4,10 @@ //! 内部使用 Aster Agent 实现 use crate::agent::{AgentMessage, AgentSession, AsterAgentState}; +use crate::config::GlobalConfigManagerState; use crate::database::dao::agent::AgentDao; use crate::database::DbConnection; +use crate::services::memory_profile_prompt_service::merge_system_prompt_with_memory_profile; use crate::workspace::WorkspaceManager; use crate::AppState; use serde::{Deserialize, Serialize}; @@ -165,6 +167,7 @@ pub struct SkillInfo { pub async fn agent_create_session( agent_state: State<'_, AsterAgentState>, db: State<'_, DbConnection>, + config_manager: State<'_, GlobalConfigManagerState>, provider_type: String, model: Option, system_prompt: Option, @@ -206,8 +209,10 @@ pub async fn agent_create_session( .configure_provider_from_pool(&db, &provider_type, &model_name, &session_id) .await?; - // 构建包含 Skills 的 System Prompt - let final_system_prompt = build_system_prompt_with_skills(system_prompt, skills.as_ref()); + // 构建包含 Skills 的 System Prompt,并附加记忆画像偏好 + let base_system_prompt = build_system_prompt_with_skills(system_prompt, skills.as_ref()); + let final_system_prompt = + merge_system_prompt_with_memory_profile(base_system_prompt, &config_manager.config()); // 保存会话到数据库 let now = chrono::Utc::now().to_rfc3339(); diff --git a/src-tauri/src/commands/aster_agent_cmd.rs b/src-tauri/src/commands/aster_agent_cmd.rs index b5f12bf21..6f2a0c1bf 100644 --- a/src-tauri/src/commands/aster_agent_cmd.rs +++ b/src-tauri/src/commands/aster_agent_cmd.rs @@ -15,6 +15,7 @@ use crate::database::DbConnection; use crate::mcp::{McpManagerState, McpServerConfig}; use crate::services::execution_tracker_service::{ExecutionTracker, RunFinalizeOptions, RunSource}; use crate::services::heartbeat_service::HeartbeatServiceState; +use crate::services::memory_profile_prompt_service::merge_system_prompt_with_memory_profile; use crate::workspace::WorkspaceManager; use aster::agents::extension::{Envs, ExtensionConfig}; use aster::agents::{Agent, AgentEvent}; @@ -1433,16 +1434,17 @@ pub async fn aster_agent_chat_stream( { let session_dir = session.working_dir.unwrap_or_default(); if !session_dir.is_empty() && session_dir != workspace_root { - tracing::warn!( - "[AsterAgent] workspace mismatch: session_id={}, workspace_id={}, session_dir={}, workspace_root={}", - session_id, - workspace_id, + tracing::info!( + "[AsterAgent] workspace 变更,自动更新 session working_dir: {} -> {}", session_dir, workspace_root ); - return Err(format!( - "workspace_mismatch|会话工作目录与 workspace 不匹配: session={session_dir}, workspace={workspace_root}" - )); + db_conn + .execute( + "UPDATE agent_sessions SET working_dir = ?1 WHERE id = ?2", + rusqlite::params![&workspace_root, session_id], + ) + .map_err(|e| format!("更新 session working_dir 失败: {e}"))?; } } } @@ -1520,7 +1522,10 @@ pub async fn aster_agent_chat_stream( } }; - (resolved_prompt, persisted) + let merged_prompt = + merge_system_prompt_with_memory_profile(resolved_prompt, &config_manager.config()); + + (merged_prompt, persisted) }; let requested_strategy = request.execution_strategy.unwrap_or(persisted_strategy); @@ -1641,11 +1646,15 @@ pub async fn aster_agent_chat_stream( let guard = agent_arc.read().await; let agent = guard.as_ref().ok_or("Agent not initialized")?; + let include_context_trace = config_manager.config().memory.enabled; + let build_session_config = || { let mut session_config_builder = SessionConfigBuilder::new(session_id); if let Some(prompt) = system_prompt.clone() { session_config_builder = session_config_builder.system_prompt(prompt); } + session_config_builder = + session_config_builder.include_context_trace(include_context_trace); session_config_builder.build() }; @@ -1921,7 +1930,9 @@ pub async fn aster_agent_submit_elicitation_response( request.user_data, )); - let session_config = SessionConfigBuilder::new(&session_id).build(); + let session_config = SessionConfigBuilder::new(&session_id) + .include_context_trace(true) + .build(); let agent_arc = state.get_agent_arc(); let guard = agent_arc.read().await; diff --git a/src-tauri/src/commands/context_memory.rs b/src-tauri/src/commands/context_memory.rs index b6b0ada0c..c729d7c44 100644 --- a/src-tauri/src/commands/context_memory.rs +++ b/src-tauri/src/commands/context_memory.rs @@ -1,5 +1,6 @@ //! 上下文记忆管理相关的 Tauri 命令 +use crate::config::GlobalConfigManagerState; use proxycast_services::context_memory_service::{ ContextMemoryService, MemoryEntry, MemoryFileType, MemoryStats, }; @@ -159,9 +160,20 @@ pub async fn get_memory_stats( #[tauri::command] pub async fn cleanup_expired_memories( memory_service: State<'_, ContextMemoryServiceState>, + global_config: State<'_, GlobalConfigManagerState>, ) -> Result<(), String> { debug!("清理过期记忆"); - memory_service.0.cleanup_expired_memories()?; + let memory_config = global_config.config().memory; + if matches!(memory_config.auto_cleanup, Some(false)) { + info!("自动清理已关闭,跳过过期记忆清理"); + return Ok(()); + } + + let retention_days = memory_config.retention_days.unwrap_or(30).clamp(1, 3650); + + memory_service + .0 + .cleanup_expired_memories_with_retention_days(retention_days)?; info!("过期记忆清理完成"); Ok(()) } diff --git a/src-tauri/src/commands/ecommerce_review_reply_cmd.rs b/src-tauri/src/commands/ecommerce_review_reply_cmd.rs index 309308a75..810399613 100644 --- a/src-tauri/src/commands/ecommerce_review_reply_cmd.rs +++ b/src-tauri/src/commands/ecommerce_review_reply_cmd.rs @@ -7,6 +7,7 @@ use tauri::State; use crate::agent::AsterAgentState; use crate::commands::skill_exec_cmd::{execute_skill, SkillExecutionResult}; +use crate::config::GlobalConfigManagerState; use crate::database::DbConnection; /// 电商差评回复请求 @@ -45,6 +46,7 @@ pub struct EcommerceReviewReplyRequest { pub async fn execute_ecommerce_review_reply( app_handle: tauri::AppHandle, db: State<'_, DbConnection>, + config_manager: State<'_, GlobalConfigManagerState>, aster_state: State<'_, AsterAgentState>, request: EcommerceReviewReplyRequest, ) -> Result { @@ -72,6 +74,7 @@ pub async fn execute_ecommerce_review_reply( execute_skill( app_handle, db, + config_manager, aster_state, "ecommerce-review-reply".to_string(), user_input, diff --git a/src-tauri/src/commands/memory_management_cmd.rs b/src-tauri/src/commands/memory_management_cmd.rs index 69c3a9a85..2d05965f5 100644 --- a/src-tauri/src/commands/memory_management_cmd.rs +++ b/src-tauri/src/commands/memory_management_cmd.rs @@ -3,7 +3,14 @@ //! 提供对话记忆的统计和管理功能 use crate::commands::context_memory::ContextMemoryServiceState; +use crate::config::GlobalConfigManagerState; use crate::database::DbConnection; +use crate::services::auto_memory_service::{ + get_auto_memory_index, update_auto_memory_note, AutoMemoryIndexResponse, +}; +use crate::services::memory_source_resolver_service::{ + resolve_effective_sources, EffectiveMemorySourcesResponse, +}; use chrono::{Local, NaiveDateTime, TimeZone}; use proxycast_services::context_memory_service::{MemoryEntry, MemoryFileType}; use rusqlite::{params, Connection}; @@ -77,6 +84,12 @@ pub struct MemoryOverviewResponse { pub entries: Vec, } +/// 自动记忆开关响应 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct MemoryAutoToggleResponse { + pub enabled: bool, +} + #[derive(Debug, Clone, Deserialize, Default)] struct ErrorEntryRecord { #[serde(default)] @@ -110,6 +123,7 @@ const CATEGORY_ORDER: [&str; 5] = [ const MAX_SOURCE_MESSAGES: usize = 6000; const MAX_GENERATED_PER_REQUEST: usize = 200; +const MAX_GENERATED_PER_REQUEST_CAP: usize = 2000; const MAX_GENERATED_PER_SESSION: usize = 40; const MIN_MESSAGE_LENGTH: usize = 18; @@ -145,6 +159,7 @@ pub async fn get_conversation_memory_overview( pub async fn request_conversation_memory_analysis( memory_service: State<'_, ContextMemoryServiceState>, db: State<'_, DbConnection>, + global_config: State<'_, GlobalConfigManagerState>, from_timestamp: Option, to_timestamp: Option, ) -> Result { @@ -159,6 +174,23 @@ pub async fn request_conversation_memory_analysis( } } + let memory_config = global_config.config().memory; + if !memory_config.enabled { + info!("[记忆管理] 记忆功能已关闭,跳过分析"); + return Ok(MemoryAnalysisResult { + analyzed_sessions: 0, + analyzed_messages: 0, + generated_entries: 0, + deduplicated_entries: 0, + }); + } + + let max_generated_per_request = memory_config + .max_entries + .unwrap_or(MAX_GENERATED_PER_REQUEST as u32) + .clamp(1, MAX_GENERATED_PER_REQUEST_CAP as u32) + as usize; + let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; let candidates = load_memory_candidates(&conn, from_timestamp, to_timestamp)?; @@ -218,7 +250,7 @@ pub async fn request_conversation_memory_analysis( generated_entries += 1; *counter += 1; - if generated_entries as usize >= MAX_GENERATED_PER_REQUEST { + if generated_entries as usize >= max_generated_per_request { break; } } @@ -237,14 +269,28 @@ pub async fn request_conversation_memory_analysis( #[tauri::command] pub async fn cleanup_conversation_memory( memory_service: State<'_, ContextMemoryServiceState>, + global_config: State<'_, GlobalConfigManagerState>, ) -> Result { info!("[记忆管理] 开始清理过期记忆"); + let memory_config = global_config.config().memory; + if matches!(memory_config.auto_cleanup, Some(false)) { + info!("[记忆管理] 自动清理已关闭,跳过清理"); + return Ok(CleanupMemoryResult { + cleaned_entries: 0, + freed_space: 0, + }); + } + + let retention_days = memory_config.retention_days.unwrap_or(30).clamp(1, 3650); + let memory_dir = resolve_memory_dir(); let before = collect_memory_overview(&memory_dir)?; // 使用 ContextMemoryService 的清理功能 - memory_service.0.cleanup_expired_memories()?; + memory_service + .0 + .cleanup_expired_memories_with_retention_days(retention_days)?; let after = collect_memory_overview(&memory_dir)?; @@ -263,12 +309,93 @@ pub async fn cleanup_conversation_memory( }) } +/// 获取当前会话可见的有效记忆来源(含 AGENTS、规则、自动记忆) +#[tauri::command] +pub async fn memory_get_effective_sources( + global_config: State<'_, GlobalConfigManagerState>, + working_dir: Option, + active_relative_path: Option, +) -> Result { + let config = global_config.config(); + let resolved_working_dir = resolve_working_dir(working_dir)?; + let resolution = resolve_effective_sources( + &config, + &resolved_working_dir, + active_relative_path.as_deref(), + ); + Ok(resolution.response) +} + +/// 获取自动记忆入口索引 +#[tauri::command] +pub async fn memory_get_auto_index( + global_config: State<'_, GlobalConfigManagerState>, + working_dir: Option, +) -> Result { + let config = global_config.config(); + let resolved_working_dir = resolve_working_dir(working_dir)?; + get_auto_memory_index(&config.memory, &resolved_working_dir) +} + +/// 切换自动记忆开关(写入全局配置) +#[tauri::command] +pub async fn memory_toggle_auto( + global_config: State<'_, GlobalConfigManagerState>, + enabled: bool, +) -> Result { + let mut config = global_config.config(); + config.memory.auto.enabled = enabled; + + global_config + .save_config(&config) + .await + .map_err(|e| format!("保存自动记忆开关失败: {e}"))?; + + Ok(MemoryAutoToggleResponse { + enabled: config.memory.auto.enabled, + }) +} + +/// 更新自动记忆笔记(写入 MEMORY.md 或 topic 文件) +#[tauri::command] +pub async fn memory_update_auto_note( + global_config: State<'_, GlobalConfigManagerState>, + working_dir: Option, + note: String, + topic: Option, +) -> Result { + let config = global_config.config(); + let resolved_working_dir = resolve_working_dir(working_dir)?; + update_auto_memory_note( + &config.memory, + &resolved_working_dir, + ¬e, + topic.as_deref(), + ) +} + fn resolve_memory_dir() -> PathBuf { dirs::home_dir() .map(|p| p.join(".proxycast").join("memory")) .unwrap_or_else(|| PathBuf::from(".proxycast/memory")) } +fn resolve_working_dir(working_dir: Option) -> Result { + if let Some(path) = working_dir + .as_deref() + .map(str::trim) + .filter(|p| !p.is_empty()) + { + let candidate = PathBuf::from(path); + let canonical = candidate + .canonicalize() + .map_err(|e| format!("working_dir 无效: {path} ({e})"))?; + return Ok(canonical); + } + + std::env::current_dir().map_err(|e| format!("获取当前工作目录失败: {e}")) +} + fn collect_memory_overview(memory_dir: &Path) -> Result { if !memory_dir.exists() { return Ok(MemoryOverviewResponse { diff --git a/src-tauri/src/commands/persona_cmd.rs b/src-tauri/src/commands/persona_cmd.rs index 1fdd6559b..39e87119f 100644 --- a/src-tauri/src/commands/persona_cmd.rs +++ b/src-tauri/src/commands/persona_cmd.rs @@ -280,6 +280,7 @@ pub struct GeneratedPersona { pub async fn generate_persona( agent_state: State<'_, crate::agent::AsterAgentState>, db: State<'_, DbConnection>, + config_manager: State<'_, crate::config::GlobalConfigManagerState>, prompt: String, ) -> Result { use aster::conversation::message::Message; @@ -354,7 +355,17 @@ pub async fn generate_persona( let cancel_token = agent_state.create_cancel_token(&session_id).await; let user_message = Message::user().with_text(&user_prompt); - let session_config = crate::agent::aster_state::SessionConfigBuilder::new(&session_id).build(); + let mut session_config_builder = + crate::agent::aster_state::SessionConfigBuilder::new(&session_id) + .include_context_trace(true); + if let Some(memory_prompt) = + crate::services::memory_profile_prompt_service::build_memory_profile_prompt( + &config_manager.config(), + ) + { + session_config_builder = session_config_builder.system_prompt(memory_prompt); + } + let session_config = session_config_builder.build(); // 获取 Agent 引用 let agent_arc = agent_state.get_agent_arc(); diff --git a/src-tauri/src/commands/skill_exec_cmd.rs b/src-tauri/src/commands/skill_exec_cmd.rs index 3c4be2b03..c7628728c 100644 --- a/src-tauri/src/commands/skill_exec_cmd.rs +++ b/src-tauri/src/commands/skill_exec_cmd.rs @@ -29,8 +29,10 @@ use crate::commands::skill_error::{ SKILL_ERR_EXECUTE_FAILED, SKILL_ERR_PROVIDER_UNAVAILABLE, SKILL_ERR_SESSION_INIT_FAILED, SKILL_ERR_STREAM_FAILED, }; +use crate::config::GlobalConfigManagerState; use crate::database::DbConnection; use crate::services::execution_tracker_service::{ExecutionTracker, RunFinishDecision, RunSource}; +use crate::services::memory_profile_prompt_service::build_memory_profile_prompt; use crate::skills::TauriExecutionCallback; use proxycast_agent::event_converter::convert_agent_event; use proxycast_skills::{ @@ -154,6 +156,7 @@ pub struct SkillExecutionResult { pub async fn execute_skill( app_handle: tauri::AppHandle, db: State<'_, DbConnection>, + config_manager: State<'_, GlobalConfigManagerState>, aster_state: State<'_, AsterAgentState>, skill_name: String, user_input: String, @@ -165,6 +168,7 @@ pub async fn execute_skill( // 生成执行 ID,并优先复用前端会话 ID(提升 /skill 与主会话上下文一致性) let execution_id = execution_id.unwrap_or_else(|| Uuid::new_v4().to_string()); let session_id = session_id.unwrap_or_else(|| format!("skill-exec-{}", Uuid::new_v4())); + let memory_profile_prompt = build_memory_profile_prompt(&config_manager.config()); let tracker = ExecutionTracker::new(db.inner().clone()); tracker @@ -293,6 +297,7 @@ pub async fn execute_skill( &execution_id, &session_id, &callback, + memory_profile_prompt.as_deref(), ) .await } else { @@ -305,6 +310,7 @@ pub async fn execute_skill( &execution_id, &session_id, &callback, + memory_profile_prompt.as_deref(), ) .await } @@ -352,13 +358,20 @@ async fn execute_skill_prompt( execution_id: &str, session_id: &str, callback: &TauriExecutionCallback, + memory_profile_prompt: Option<&str>, ) -> Result { // 发送步骤开始事件 callback.on_step_start("main", &skill.display_name, 1, 1); // 构建 SessionConfig + let mut combined_prompt = skill.markdown_content.clone(); + if let Some(memory_prompt) = memory_profile_prompt { + combined_prompt = format!("{combined_prompt}\n\n{memory_prompt}"); + } + let session_config = SessionConfigBuilder::new(session_id) - .system_prompt(&skill.markdown_content) + .system_prompt(combined_prompt) + .include_context_trace(true) .build(); let user_message = Message::user().with_text(user_input); @@ -469,6 +482,7 @@ async fn execute_skill_workflow( execution_id: &str, session_id: &str, callback: &TauriExecutionCallback, + memory_profile_prompt: Option<&str>, ) -> Result { let steps = &skill.workflow_steps; let total_steps = steps.len(); @@ -503,9 +517,16 @@ async fn execute_skill_workflow( skill.markdown_content, step.name, step_num, total_steps, step.prompt ); + let step_prompt_with_memory = if let Some(memory_prompt) = memory_profile_prompt { + format!("{step_system_prompt}\n\n{memory_prompt}") + } else { + step_system_prompt + }; + let step_session_id = format!("{}-step-{}", session_id, step.id); let session_config = SessionConfigBuilder::new(&step_session_id) - .system_prompt(&step_system_prompt) + .system_prompt(step_prompt_with_memory) + .include_context_trace(true) .build(); // 用户消息 = 原始输入 + 前序步骤的累积上下文 diff --git a/src-tauri/src/commands/subagent_cmd.rs b/src-tauri/src/commands/subagent_cmd.rs index fc7283123..4bcae19eb 100644 --- a/src-tauri/src/commands/subagent_cmd.rs +++ b/src-tauri/src/commands/subagent_cmd.rs @@ -9,7 +9,7 @@ use tokio::sync::RwLock; use aster::agents::context::AgentContext; use aster::agents::subagent_scheduler::{SchedulerConfig, SchedulerExecutionResult, SubAgentTask}; -use crate::agent::subagent_scheduler::ProxyCastScheduler; +use crate::agent::subagent_scheduler::{ProxyCastScheduler, SubAgentRole}; use crate::database::DbConnection; /// SubAgent 调度器状态 @@ -59,6 +59,7 @@ pub async fn execute_subagent_tasks( state: State<'_, SubAgentSchedulerState>, tasks: Vec, config: Option, + role: Option, ) -> Result { // 确保调度器已初始化 let scheduler_guard = state.scheduler.read().await; @@ -79,11 +80,17 @@ pub async fn execute_subagent_tasks( // 创建父上下文 let parent_context = AgentContext::new(); - // 执行任务 - scheduler - .execute(tasks, Some(&parent_context)) - .await - .map_err(|e| e.to_string()) + // 根据是否指定角色选择执行方式 + match role { + Some(role) => scheduler + .execute_with_role(tasks, Some(&parent_context), role) + .await + .map_err(|e| e.to_string()), + None => scheduler + .execute(tasks, Some(&parent_context)) + .await + .map_err(|e| e.to_string()), + } } /// 取消 SubAgent 任务 diff --git a/src-tauri/src/commands/unified_chat_cmd.rs b/src-tauri/src/commands/unified_chat_cmd.rs index c807e4173..452a49639 100644 --- a/src-tauri/src/commands/unified_chat_cmd.rs +++ b/src-tauri/src/commands/unified_chat_cmd.rs @@ -15,8 +15,10 @@ use crate::agent::aster_state::SessionConfigBuilder; use crate::agent::{AsterAgentState, TauriAgentEvent}; +use crate::config::GlobalConfigManagerState; use crate::database::dao::chat::{ChatDao, ChatMessage, ChatMode, ChatSession}; use crate::database::DbConnection; +use crate::services::memory_profile_prompt_service::merge_system_prompt_with_memory_profile; use aster::conversation::message::Message; use futures::StreamExt; use proxycast_agent::event_converter::convert_agent_event; @@ -104,17 +106,23 @@ impl From for SessionResponse { pub async fn chat_create_session( db: State<'_, DbConnection>, agent_state: State<'_, AsterAgentState>, + config_manager: State<'_, GlobalConfigManagerState>, request: CreateSessionRequest, ) -> Result { let now = chrono::Utc::now().to_rfc3339(); let session_id = uuid::Uuid::new_v4().to_string(); + let merged_system_prompt = merge_system_prompt_with_memory_profile( + request.system_prompt.clone(), + &config_manager.config(), + ); + // 创建会话 let session = ChatSession { id: session_id.clone(), mode: request.mode, title: request.title, - system_prompt: request.system_prompt.clone(), + system_prompt: merged_system_prompt, model: request.model.clone(), provider_type: request.provider_type.clone(), credential_uuid: None, @@ -290,6 +298,7 @@ pub async fn chat_send_message( app: AppHandle, db: State<'_, DbConnection>, agent_state: State<'_, AsterAgentState>, + config_manager: State<'_, GlobalConfigManagerState>, request: SendMessageRequest, ) -> Result<(), String> { let start_time = std::time::Instant::now(); @@ -338,6 +347,11 @@ pub async fn chat_send_message( tracing::debug!("[UnifiedChat] 数据库查询耗时: {:?}", db_elapsed); // 根据模式处理 + let merged_system_prompt = merge_system_prompt_with_memory_profile( + session.system_prompt.clone(), + &config_manager.config(), + ); + let result = match session.mode { ChatMode::Agent | ChatMode::Creator => { // 使用 Aster Agent 处理 @@ -348,7 +362,8 @@ pub async fn chat_send_message( &request.session_id, &request.message, &request.event_name, - session.system_prompt.as_deref(), + merged_system_prompt.as_deref(), + config_manager.config().memory.enabled, ) .await } @@ -361,7 +376,8 @@ pub async fn chat_send_message( &request.session_id, &request.message, &request.event_name, - session.system_prompt.as_deref(), + merged_system_prompt.as_deref(), + config_manager.config().memory.enabled, ) .await } @@ -386,6 +402,7 @@ async fn send_message_with_aster( message: &str, event_name: &str, system_prompt: Option<&str>, + include_context_trace: bool, ) -> Result<(), String> { let start_time = std::time::Instant::now(); @@ -419,7 +436,9 @@ async fn send_message_with_aster( }; let user_message = Message::user().with_text(&final_message); - let session_config = SessionConfigBuilder::new(session_id).build(); + let session_config = SessionConfigBuilder::new(session_id) + .include_context_trace(include_context_trace) + .build(); // 获取 Agent 引用 let agent_arc = agent_state.get_agent_arc(); diff --git a/src-tauri/src/commands/unified_memory_cmd.rs b/src-tauri/src/commands/unified_memory_cmd.rs index 825fb7117..59500cab2 100644 --- a/src-tauri/src/commands/unified_memory_cmd.rs +++ b/src-tauri/src/commands/unified_memory_cmd.rs @@ -2,6 +2,7 @@ //! //! Provides unified memory CRUD operations and analysis pipeline. +use crate::config::GlobalConfigManagerState; use crate::database::DbConnection; use chrono::{Local, TimeZone}; use proxycast_memory::extractor::{self, ExtractionContext}; @@ -17,6 +18,7 @@ const DEFAULT_LIST_LIMIT: usize = 120; const MAX_LIST_LIMIT: usize = 1000; const MAX_SOURCE_MESSAGES: usize = 6000; const MAX_GENERATED_PER_REQUEST: usize = 200; +const MAX_GENERATED_PER_REQUEST_CAP: usize = 2000; const MAX_GENERATED_PER_SESSION: usize = 40; const MIN_MESSAGE_LENGTH: usize = 18; const MAX_LLM_SESSIONS: usize = 20; @@ -429,6 +431,7 @@ pub async fn unified_memory_stats( #[tauri::command] pub async fn unified_memory_analyze( db: State<'_, DbConnection>, + global_config: State<'_, GlobalConfigManagerState>, from_timestamp: Option, to_timestamp: Option, ) -> Result { @@ -443,6 +446,23 @@ pub async fn unified_memory_analyze( } } + let memory_config = global_config.config().memory; + if !memory_config.enabled { + info!("[Unified Memory] 记忆功能已关闭,跳过分析"); + return Ok(MemoryAnalysisResult { + analyzed_sessions: 0, + analyzed_messages: 0, + generated_entries: 0, + deduplicated_entries: 0, + }); + } + + let max_generated_per_request = memory_config + .max_entries + .unwrap_or(MAX_GENERATED_PER_REQUEST as u32) + .clamp(1, MAX_GENERATED_PER_REQUEST_CAP as u32) + as usize; + let candidates = { let conn = db.lock().map_err(|e| format!("数据库锁定失败: {e}"))?; load_memory_candidates(&conn, from_timestamp, to_timestamp)? @@ -470,7 +490,7 @@ pub async fn unified_memory_analyze( let llm_attempted = llm_api_key.is_some(); if let Some(api_key) = llm_api_key { - match build_pending_from_llm(&db, &candidates, &api_key).await { + match build_pending_from_llm(&db, &candidates, &api_key, max_generated_per_request).await { Ok((mut llm_pending, llm_dedup)) => { deduplicated_entries += llm_dedup; pending_memories.append(&mut llm_pending); @@ -482,13 +502,14 @@ pub async fn unified_memory_analyze( } if !llm_attempted || pending_memories.is_empty() { - let (mut fallback_pending, fallback_dedup) = build_pending_from_rules(&db, &candidates)?; + let (mut fallback_pending, fallback_dedup) = + build_pending_from_rules(&db, &candidates, max_generated_per_request)?; deduplicated_entries += fallback_dedup; pending_memories.append(&mut fallback_pending); } - if pending_memories.len() > MAX_GENERATED_PER_REQUEST { - pending_memories.truncate(MAX_GENERATED_PER_REQUEST); + if pending_memories.len() > max_generated_per_request { + pending_memories.truncate(max_generated_per_request); } let generated_entries = { @@ -518,6 +539,7 @@ pub async fn unified_memory_analyze( fn build_pending_from_rules( db: &State<'_, DbConnection>, candidates: &[MemorySourceCandidate], + max_generated_per_request: usize, ) -> Result<(Vec, u32), String> { let mut pending_memories = Vec::new(); let mut deduplicated_entries = 0u32; @@ -572,7 +594,7 @@ fn build_pending_from_rules( pending_memories.push(pending); *counter += 1; - if pending_memories.len() >= MAX_GENERATED_PER_REQUEST { + if pending_memories.len() >= max_generated_per_request { break; } } @@ -584,6 +606,7 @@ async fn build_pending_from_llm( db: &State<'_, DbConnection>, candidates: &[MemorySourceCandidate], api_key: &str, + max_generated_per_request: usize, ) -> Result<(Vec, u32), String> { let mut grouped: HashMap> = HashMap::new(); for candidate in candidates.iter().cloned() { @@ -687,12 +710,12 @@ async fn build_pending_from_llm( existing_mut.push(pending_to_memory(pending.clone())); pending_memories.push(pending); - if pending_memories.len() >= MAX_GENERATED_PER_REQUEST { + if pending_memories.len() >= max_generated_per_request { break; } } - if pending_memories.len() >= MAX_GENERATED_PER_REQUEST { + if pending_memories.len() >= max_generated_per_request { break; } } diff --git a/src-tauri/src/config/tests.rs b/src-tauri/src/config/tests.rs index 97cdea0a3..a970778db 100644 --- a/src-tauri/src/config/tests.rs +++ b/src-tauri/src/config/tests.rs @@ -203,6 +203,7 @@ fn arb_config() -> impl Strategy { hint_router: proxycast_core::config::HintRouterSettings::default(), pairing: proxycast_core::config::PairingSettings::default(), heartbeat: proxycast_core::config::HeartbeatSettings::default(), + channels: proxycast_core::config::ChannelsConfig::default(), }) } @@ -454,6 +455,7 @@ fn arb_valid_config() -> impl Strategy { hint_router: proxycast_core::config::HintRouterSettings::default(), pairing: proxycast_core::config::PairingSettings::default(), heartbeat: proxycast_core::config::HeartbeatSettings::default(), + channels: proxycast_core::config::ChannelsConfig::default(), }) } @@ -516,6 +518,7 @@ fn arb_invalid_config() -> impl Strategy { hint_router: proxycast_core::config::HintRouterSettings::default(), pairing: proxycast_core::config::PairingSettings::default(), heartbeat: proxycast_core::config::HeartbeatSettings::default(), + channels: proxycast_core::config::ChannelsConfig::default(), }; // 根据类型使配置无效 match invalid_type { diff --git a/src-tauri/src/services/auto_memory_service.rs b/src-tauri/src/services/auto_memory_service.rs new file mode 100644 index 000000000..ecf76ae2f --- /dev/null +++ b/src-tauri/src/services/auto_memory_service.rs @@ -0,0 +1,343 @@ +//! 自动记忆服务 +//! +//! 提供自动记忆目录定位、入口索引读取与笔记更新能力。 + +use chrono::Local; +use proxycast_core::config::{MemoryAutoConfig, MemoryConfig}; +use serde::{Deserialize, Serialize}; +use std::fs; +use std::path::{Path, PathBuf}; + +/// 自动记忆索引项 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct AutoMemoryIndexItem { + pub title: String, + pub relative_path: String, + pub exists: bool, + pub summary: Option, +} + +/// 自动记忆索引响应 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct AutoMemoryIndexResponse { + pub enabled: bool, + pub root_dir: String, + pub entrypoint: String, + pub max_loaded_lines: u32, + pub entry_exists: bool, + pub total_lines: u32, + pub preview_lines: Vec, + pub items: Vec, +} + +/// 读取自动记忆索引 +pub fn get_auto_memory_index( + memory_config: &MemoryConfig, + working_dir: &Path, +) -> Result { + let auto = &memory_config.auto; + let root_dir = resolve_auto_memory_root(working_dir, auto); + let entry_name = auto.entrypoint.trim(); + let entry_name = if entry_name.is_empty() { + "MEMORY.md" + } else { + entry_name + }; + let entry_path = root_dir.join(entry_name); + + let mut response = AutoMemoryIndexResponse { + enabled: auto.enabled, + root_dir: root_dir.to_string_lossy().to_string(), + entrypoint: entry_name.to_string(), + max_loaded_lines: auto.max_loaded_lines, + entry_exists: entry_path.is_file(), + total_lines: 0, + preview_lines: Vec::new(), + items: Vec::new(), + }; + + if !entry_path.is_file() { + return Ok(response); + } + + let raw = fs::read_to_string(&entry_path) + .map_err(|e| format!("读取自动记忆入口失败 {}: {e}", entry_path.display()))?; + let lines: Vec = raw.lines().map(|s| s.to_string()).collect(); + response.total_lines = lines.len() as u32; + response.preview_lines = lines + .iter() + .take(auto.max_loaded_lines as usize) + .cloned() + .collect(); + response.items = parse_index_items(&lines, &root_dir); + + Ok(response) +} + +/// 更新自动记忆笔记 +pub fn update_auto_memory_note( + memory_config: &MemoryConfig, + working_dir: &Path, + note: &str, + topic: Option<&str>, +) -> Result { + let trimmed_note = note.trim(); + if trimmed_note.is_empty() { + return Err("note 不能为空".to_string()); + } + + let auto = &memory_config.auto; + let root_dir = resolve_auto_memory_root(working_dir, auto); + fs::create_dir_all(&root_dir) + .map_err(|e| format!("创建自动记忆目录失败 {}: {e}", root_dir.display()))?; + + let entry_name = auto.entrypoint.trim(); + let entry_name = if entry_name.is_empty() { + "MEMORY.md" + } else { + entry_name + }; + let entry_path = root_dir.join(entry_name); + + if let Some(topic_name) = topic.map(str::trim).filter(|v| !v.is_empty()) { + let topic_file = normalize_topic_filename(topic_name); + let topic_path = root_dir.join(&topic_file); + append_topic_note(&topic_path, topic_name, trimmed_note)?; + ensure_entry_link(&entry_path, topic_name, &topic_file)?; + } else { + append_entry_note(&entry_path, trimmed_note)?; + } + + get_auto_memory_index(memory_config, working_dir) +} + +/// 解析自动记忆根目录 +pub fn resolve_auto_memory_root(working_dir: &Path, auto: &MemoryAutoConfig) -> PathBuf { + if let Some(custom_root) = auto.root_dir.as_deref().map(str::trim) { + if !custom_root.is_empty() { + return expand_path(custom_root, Some(working_dir)); + } + } + + let project_anchor = find_git_root(working_dir).unwrap_or_else(|| working_dir.to_path_buf()); + let slug = project_anchor + .to_string_lossy() + .replace(['\\', '/', ':', ' '], "_") + .trim_matches('_') + .to_string(); + let project_slug = if slug.is_empty() { + "default".to_string() + } else { + slug + }; + + dirs::home_dir() + .unwrap_or_else(|| PathBuf::from(".")) + .join(".proxycast") + .join("projects") + .join(project_slug) + .join("memory") +} + +fn parse_index_items(lines: &[String], root_dir: &Path) -> Vec { + let mut items = Vec::new(); + + for line in lines { + let trimmed = line.trim(); + if trimmed.is_empty() { + continue; + } + + // Markdown link: - [title](path) + if let Some((title, relative_path)) = parse_markdown_link(trimmed) { + let path = root_dir.join(&relative_path); + items.push(AutoMemoryIndexItem { + title, + relative_path, + exists: path.is_file(), + summary: None, + }); + continue; + } + + // import 风格:@topic.md + if let Some(import_target) = trimmed.strip_prefix('@') { + let relative_path = import_target.trim().to_string(); + if relative_path.is_empty() { + continue; + } + let path = root_dir.join(&relative_path); + items.push(AutoMemoryIndexItem { + title: relative_path.clone(), + relative_path, + exists: path.is_file(), + summary: None, + }); + } + } + + items +} + +fn parse_markdown_link(line: &str) -> Option<(String, String)> { + let cleaned = line + .trim_start_matches("- ") + .trim_start_matches("* ") + .trim(); + let title_start = cleaned.find('[')?; + let title_end = cleaned[title_start + 1..].find(']')? + title_start + 1; + let path_start = cleaned[title_end + 1..].find('(')? + title_end + 1; + let path_end = cleaned[path_start + 1..].find(')')? + path_start + 1; + + let title = cleaned[title_start + 1..title_end].trim().to_string(); + let path = cleaned[path_start + 1..path_end].trim().to_string(); + if title.is_empty() || path.is_empty() { + return None; + } + Some((title, path)) +} + +fn append_entry_note(entry_path: &Path, note: &str) -> Result<(), String> { + let timestamp = Local::now().format("%Y-%m-%d %H:%M:%S"); + let line = format!("- [{timestamp}] {note}\n"); + let mut existing = if entry_path.is_file() { + fs::read_to_string(entry_path) + .map_err(|e| format!("读取 MEMORY 入口失败 {}: {e}", entry_path.display()))? + } else { + "# Auto Memory Index\n\n".to_string() + }; + if !existing.ends_with('\n') { + existing.push('\n'); + } + existing.push_str(&line); + fs::write(entry_path, existing) + .map_err(|e| format!("写入 MEMORY 入口失败 {}: {e}", entry_path.display())) +} + +fn append_topic_note(topic_path: &Path, topic_name: &str, note: &str) -> Result<(), String> { + let timestamp = Local::now().format("%Y-%m-%d %H:%M:%S"); + let mut content = if topic_path.is_file() { + fs::read_to_string(topic_path) + .map_err(|e| format!("读取主题记忆失败 {}: {e}", topic_path.display()))? + } else { + format!("# {topic_name}\n\n") + }; + if !content.ends_with('\n') { + content.push('\n'); + } + content.push_str(&format!("## {timestamp}\n\n{note}\n\n")); + fs::write(topic_path, content) + .map_err(|e| format!("写入主题记忆失败 {}: {e}", topic_path.display())) +} + +fn ensure_entry_link(entry_path: &Path, topic_name: &str, topic_file: &str) -> Result<(), String> { + let mut content = if entry_path.is_file() { + fs::read_to_string(entry_path) + .map_err(|e| format!("读取 MEMORY 入口失败 {}: {e}", entry_path.display()))? + } else { + "# Auto Memory Index\n\n".to_string() + }; + let marker = format!("({topic_file})"); + if !content.contains(&marker) { + if !content.ends_with('\n') { + content.push('\n'); + } + content.push_str(&format!("- [{topic_name}]({topic_file})\n")); + } + fs::write(entry_path, content) + .map_err(|e| format!("写入 MEMORY 入口失败 {}: {e}", entry_path.display())) +} + +fn normalize_topic_filename(topic: &str) -> String { + let lowered = topic.trim().to_lowercase(); + let mut slug = String::with_capacity(lowered.len() + 3); + for ch in lowered.chars() { + if ch.is_ascii_alphanumeric() { + slug.push(ch); + } else if ch == '-' || ch == '_' || ch == ' ' { + slug.push('-'); + } + } + while slug.contains("--") { + slug = slug.replace("--", "-"); + } + let slug = slug.trim_matches('-'); + if slug.is_empty() { + "notes.md".to_string() + } else if slug.ends_with(".md") { + slug.to_string() + } else { + format!("{slug}.md") + } +} + +fn expand_path(path: &str, working_dir: Option<&Path>) -> PathBuf { + if path.starts_with("~/") { + if let Some(home) = dirs::home_dir() { + return home.join(path.trim_start_matches("~/")); + } + } + + let p = PathBuf::from(path); + if p.is_absolute() { + return p; + } + + if let Some(base) = working_dir { + return base.join(p); + } + p +} + +fn find_git_root(start: &Path) -> Option { + let mut current = if start.is_file() { + start.parent()?.to_path_buf() + } else { + start.to_path_buf() + }; + + loop { + if current.join(".git").exists() { + return Some(current); + } + if !current.pop() { + return None; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + #[test] + fn should_create_entry_when_update_note_without_topic() { + let tmp = TempDir::new().expect("create temp dir"); + let mut cfg = MemoryConfig::default(); + cfg.auto.root_dir = Some(tmp.path().to_string_lossy().to_string()); + cfg.auto.entrypoint = "MEMORY.md".to_string(); + cfg.auto.enabled = true; + + let result = + update_auto_memory_note(&cfg, tmp.path(), "记下这个偏好", None).expect("update note"); + assert!(result.entry_exists); + assert!(!result.preview_lines.is_empty()); + } + + #[test] + fn should_add_topic_and_index_link() { + let tmp = TempDir::new().expect("create temp dir"); + let mut cfg = MemoryConfig::default(); + cfg.auto.root_dir = Some(tmp.path().to_string_lossy().to_string()); + cfg.auto.entrypoint = "MEMORY.md".to_string(); + cfg.auto.enabled = true; + + let result = update_auto_memory_note(&cfg, tmp.path(), "pnpm only", Some("workflow")) + .expect("update topic note"); + assert!(result + .items + .iter() + .any(|item| item.relative_path == "workflow.md")); + } +} diff --git a/src-tauri/src/services/memory_import_parser_service.rs b/src-tauri/src/services/memory_import_parser_service.rs new file mode 100644 index 000000000..9ac357eaa --- /dev/null +++ b/src-tauri/src/services/memory_import_parser_service.rs @@ -0,0 +1,264 @@ +//! 记忆文件 @import 解析服务 +//! +//! 支持从 Markdown 文档中解析以 `@` 开头的导入行,并递归展开。 + +use std::collections::HashSet; +use std::fs; +use std::path::{Path, PathBuf}; + +/// @import 解析选项 +#[derive(Debug, Clone)] +pub struct MemoryImportParseOptions { + /// 是否启用导入解析 + pub follow_imports: bool, + /// 最大递归深度(根文件为 0) + pub max_depth: usize, +} + +impl Default for MemoryImportParseOptions { + fn default() -> Self { + Self { + follow_imports: true, + max_depth: 5, + } + } +} + +/// @import 解析结果 +#[derive(Debug, Clone, Default)] +pub struct MemoryImportParseResult { + /// 展开后的完整内容 + pub content: String, + /// 成功导入的文件列表 + pub imported_files: Vec, + /// 解析过程中的告警 + pub warnings: Vec, +} + +/// 读取并解析记忆文件(支持 @import) +pub fn parse_memory_file( + entry_path: &Path, + options: &MemoryImportParseOptions, +) -> Result { + if !entry_path.exists() { + return Err(format!("文件不存在: {}", entry_path.display())); + } + if !entry_path.is_file() { + return Err(format!("路径不是文件: {}", entry_path.display())); + } + + let mut result = MemoryImportParseResult::default(); + let mut visited = HashSet::new(); + let normalized_entry = normalize_path(entry_path); + visited.insert(normalized_entry.clone()); + + let content = parse_file_recursive( + &normalized_entry, + options, + 0, + &mut visited, + &mut result.imported_files, + &mut result.warnings, + )?; + + result.content = content; + Ok(result) +} + +fn parse_file_recursive( + file_path: &Path, + options: &MemoryImportParseOptions, + depth: usize, + visited: &mut HashSet, + imported_files: &mut Vec, + warnings: &mut Vec, +) -> Result { + let raw = fs::read_to_string(file_path) + .map_err(|e| format!("读取记忆文件失败 {}: {e}", file_path.display()))?; + + if !options.follow_imports { + return Ok(raw); + } + + let mut output = String::new(); + let mut in_code_block = false; + + for line in raw.lines() { + let trimmed = line.trim(); + + if trimmed.starts_with("```") { + in_code_block = !in_code_block; + output.push_str(line); + output.push('\n'); + continue; + } + + if in_code_block || !trimmed.starts_with('@') || trimmed.starts_with("@@") { + output.push_str(line); + output.push('\n'); + continue; + } + + let import_target = trimmed.trim_start_matches('@').trim(); + if import_target.is_empty() { + output.push_str(line); + output.push('\n'); + continue; + } + + if depth >= options.max_depth { + warnings.push(format!( + "导入深度超限({}),已跳过: {} -> {}", + options.max_depth, + file_path.display(), + import_target + )); + output.push_str(&format!( + "\n", + options.max_depth + )); + continue; + } + + let resolved = resolve_import_path(import_target, file_path.parent()); + let Some(import_path) = resolved else { + warnings.push(format!( + "无法解析导入路径: {} -> {}", + file_path.display(), + import_target + )); + output.push_str(line); + output.push('\n'); + continue; + }; + let normalized_import = normalize_path(&import_path); + + if visited.contains(&normalized_import) { + warnings.push(format!( + "检测到循环导入,已跳过: {}", + normalized_import.display() + )); + output.push_str(&format!( + "\n", + normalized_import.display() + )); + continue; + } + + if !normalized_import.exists() || !normalized_import.is_file() { + warnings.push(format!("导入目标不存在: {}", normalized_import.display())); + output.push_str(&format!( + "\n", + normalized_import.display() + )); + continue; + } + + visited.insert(normalized_import.clone()); + imported_files.push(normalized_import.clone()); + let imported_content = parse_file_recursive( + &normalized_import, + options, + depth + 1, + visited, + imported_files, + warnings, + )?; + visited.remove(&normalized_import); + + output.push_str(&format!( + "\n", + normalized_import.display() + )); + output.push_str(imported_content.trim_end()); + output.push('\n'); + output.push_str(&format!( + "\n", + normalized_import.display() + )); + } + + Ok(output) +} + +fn resolve_import_path(import_target: &str, base_dir: Option<&Path>) -> Option { + let normalized = import_target.replace("\\ ", " "); + + if normalized.starts_with('/') { + return Some(PathBuf::from(normalized)); + } + + if normalized.starts_with("~/") { + let home = dirs::home_dir()?; + return Some(home.join(normalized.trim_start_matches("~/"))); + } + + let base = base_dir.unwrap_or_else(|| Path::new(".")); + Some(base.join(normalized)) +} + +fn normalize_path(path: &Path) -> PathBuf { + path.canonicalize().unwrap_or_else(|_| path.to_path_buf()) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + use tempfile::TempDir; + + #[test] + fn should_parse_nested_imports() { + let tmp = TempDir::new().expect("create temp dir"); + let root = tmp.path(); + let main = root.join("main.md"); + let a = root.join("a.md"); + let b = root.join("b.md"); + fs::write(&b, "B-Content").expect("write b"); + fs::write(&a, format!("A-Header\n@{}\nA-Footer", b.display())).expect("write a"); + fs::write(&main, format!("Main\n@{}\nDone", a.display())).expect("write main"); + + let result = parse_memory_file(&main, &MemoryImportParseOptions::default()) + .expect("parse memory file"); + + assert!(result.content.contains("Main")); + assert!(result.content.contains("A-Header")); + assert!(result.content.contains("B-Content")); + assert!(result.imported_files.len() >= 2); + } + + #[test] + fn should_handle_cyclic_imports() { + let tmp = TempDir::new().expect("create temp dir"); + let root = tmp.path(); + let a = root.join("a.md"); + let b = root.join("b.md"); + fs::write(&a, "@./b.md").expect("write a"); + fs::write(&b, "@./a.md").expect("write b"); + + let result = + parse_memory_file(&a, &MemoryImportParseOptions::default()).expect("parse memory file"); + + assert!(!result.warnings.is_empty()); + assert!(result + .warnings + .iter() + .any(|w| w.contains("循环导入") || w.contains("cyclic"))); + } + + #[test] + fn should_stop_at_max_depth() { + let tmp = TempDir::new().expect("create temp dir"); + let root = tmp.path(); + fs::write(root.join("1.md"), "@./2.md").expect("write 1"); + fs::write(root.join("2.md"), "@./3.md").expect("write 2"); + fs::write(root.join("3.md"), "deep").expect("write 3"); + + let options = MemoryImportParseOptions { + follow_imports: true, + max_depth: 1, + }; + let result = parse_memory_file(&root.join("1.md"), &options).expect("parse 1"); + assert!(result.warnings.iter().any(|w| w.contains("导入深度超限"))); + } +} diff --git a/src-tauri/src/services/memory_profile_prompt_service.rs b/src-tauri/src/services/memory_profile_prompt_service.rs new file mode 100644 index 000000000..8837d459a --- /dev/null +++ b/src-tauri/src/services/memory_profile_prompt_service.rs @@ -0,0 +1,173 @@ +//! 记忆画像提示词服务 +//! +//! 将设置页中的记忆画像(学习状态、擅长领域、解释偏好、难题偏好) +//! 转换为可注入到系统提示词中的统一指令片段。 + +use proxycast_core::config::Config; +use std::path::PathBuf; + +use crate::services::memory_source_resolver_service::build_memory_sources_prompt; + +const MEMORY_PROFILE_PROMPT_MARKER: &str = "【用户记忆画像偏好】"; + +fn normalize_text(input: &str) -> Option { + let trimmed = input.trim(); + if trimmed.is_empty() { + None + } else { + Some(trimmed.to_string()) + } +} + +fn normalize_list(items: &[String]) -> Vec { + items + .iter() + .filter_map(|item| normalize_text(item)) + .collect() +} + +/// 构建记忆画像提示词 +/// +/// 仅在以下条件满足时返回: +/// - 记忆功能已启用 +/// - 至少有一项画像字段有值 +pub fn build_memory_profile_prompt(config: &Config) -> Option { + let memory = &config.memory; + if !memory.enabled { + return None; + } + + let profile = memory.profile.as_ref()?; + + let current_status = profile.current_status.as_deref().and_then(normalize_text); + let strengths = normalize_list(&profile.strengths); + let explanation_style = normalize_list(&profile.explanation_style); + let challenge_preference = normalize_list(&profile.challenge_preference); + + let has_profile_data = current_status.is_some() + || !strengths.is_empty() + || !explanation_style.is_empty() + || !challenge_preference.is_empty(); + + if !has_profile_data { + return None; + } + + let mut lines: Vec = vec![ + MEMORY_PROFILE_PROMPT_MARKER.to_string(), + "以下是用户在设置中明确给出的长期偏好,请在回答中持续遵循:".to_string(), + ]; + + if let Some(status) = current_status { + lines.push(format!("- 当前状态:{status}")); + } + if !strengths.is_empty() { + lines.push(format!("- 擅长领域:{}", strengths.join("、"))); + } + if !explanation_style.is_empty() { + lines.push(format!("- 偏好解释方式:{}", explanation_style.join("、"))); + } + if !challenge_preference.is_empty() { + lines.push(format!( + "- 遇到难题时偏好:{}", + challenge_preference.join("、") + )); + } + + lines.push("执行要求:".to_string()); + lines.push("1. 优先按上述偏好组织回答结构、例子与解释顺序。".to_string()); + lines.push("2. 在保证正确性的前提下,控制解释粒度并匹配用户理解路径。".to_string()); + lines.push("3. 不要显式提及你看到了该画像配置。".to_string()); + + // 记忆来源补充(AGENTS、规则、自动记忆等) + let working_dir = std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")); + if let Some(source_prompt) = build_memory_sources_prompt(config, &working_dir, None, 4000) { + lines.push(String::new()); + lines.push(source_prompt); + } + + Some(lines.join("\n")) +} + +/// 合并基础系统提示词与记忆画像提示词 +/// +/// - 已包含画像标记时不会重复追加 +/// - 任一方为空时返回另一方 +pub fn merge_system_prompt_with_memory_profile( + base_prompt: Option, + config: &Config, +) -> Option { + let memory_prompt = build_memory_profile_prompt(config); + + match (base_prompt, memory_prompt) { + (Some(base), Some(memory)) => { + if base.contains(MEMORY_PROFILE_PROMPT_MARKER) { + Some(base) + } else if base.trim().is_empty() { + Some(memory) + } else { + Some(format!("{base}\n\n{memory}")) + } + } + (Some(base), None) => Some(base), + (None, Some(memory)) => Some(memory), + (None, None) => None, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use proxycast_core::config::Config; + + #[test] + fn memory_disabled_should_not_build_prompt() { + let mut config = Config::default(); + config.memory.enabled = false; + config.memory.profile = Some(Default::default()); + + let result = build_memory_profile_prompt(&config); + assert!(result.is_none()); + } + + #[test] + fn empty_profile_should_not_build_prompt() { + let mut config = Config::default(); + config.memory.enabled = true; + config.memory.profile = Some(Default::default()); + + let result = build_memory_profile_prompt(&config); + assert!(result.is_none()); + } + + #[test] + fn should_build_prompt_when_profile_has_data() { + let mut config = Config::default(); + config.memory.enabled = true; + let mut profile = config.memory.profile.clone().unwrap_or_default(); + profile.current_status = Some("研究生".to_string()); + profile.strengths = vec!["数学/逻辑推理".to_string()]; + profile.explanation_style = vec!["先举例,后讲理论".to_string()]; + profile.challenge_preference = vec!["一步一步地分解".to_string()]; + config.memory.profile = Some(profile); + + let result = build_memory_profile_prompt(&config); + assert!(result.is_some()); + let text = result.unwrap_or_default(); + assert!(text.contains("研究生")); + assert!(text.contains("先举例,后讲理论")); + } + + #[test] + fn should_not_duplicate_when_base_contains_marker() { + let mut config = Config::default(); + config.memory.enabled = true; + let mut profile = config.memory.profile.clone().unwrap_or_default(); + profile.current_status = Some("本科生".to_string()); + config.memory.profile = Some(profile); + + let base = Some("前置内容\n\n【用户记忆画像偏好】\n已有内容".to_string()); + let merged = merge_system_prompt_with_memory_profile(base.clone(), &config); + assert_eq!(merged, base); + } +} diff --git a/src-tauri/src/services/memory_rules_loader_service.rs b/src-tauri/src/services/memory_rules_loader_service.rs new file mode 100644 index 000000000..71e8dd780 --- /dev/null +++ b/src-tauri/src/services/memory_rules_loader_service.rs @@ -0,0 +1,243 @@ +//! 记忆规则加载服务 +//! +//! 负责从 `.agents/rules/**/*.md` 加载规则,并支持基于 frontmatter `paths` 的条件匹配。 + +use glob::Pattern; +use std::fs; +use std::path::{Path, PathBuf}; + +/// 规则文档 +#[derive(Debug, Clone)] +pub struct LoadedMemoryRule { + /// 规则文件路径 + pub path: PathBuf, + /// 标题(优先第一个一级标题,否则文件名) + pub title: String, + /// 规则正文(已去除 frontmatter) + pub content: String, + /// frontmatter 中的 paths 条件 + pub path_patterns: Vec, + /// 是否命中当前 active_path + pub matched: bool, +} + +/// 从规则目录递归加载规则 +/// +/// - `rule_dir`: 规则目录(通常是 `.agents/rules`) +/// - `active_path`: 当前正在处理的相对路径(用于 paths 匹配) +pub fn load_rules(rule_dir: &Path, active_path: Option<&str>) -> Vec { + let mut rule_files = Vec::new(); + collect_markdown_files(rule_dir, &mut rule_files); + rule_files.sort(); + + rule_files + .into_iter() + .filter_map(|path| parse_rule_file(&path, active_path)) + .collect() +} + +fn collect_markdown_files(dir: &Path, output: &mut Vec) { + let entries = match fs::read_dir(dir) { + Ok(entries) => entries, + Err(_) => return, + }; + + for entry in entries.flatten() { + let path = entry.path(); + if path.is_dir() { + collect_markdown_files(&path, output); + continue; + } + + if !path.is_file() { + continue; + } + + if path + .extension() + .and_then(|ext| ext.to_str()) + .map(|ext| ext.eq_ignore_ascii_case("md")) + .unwrap_or(false) + { + output.push(path); + } + } +} + +fn parse_rule_file(path: &Path, active_path: Option<&str>) -> Option { + let raw = fs::read_to_string(path).ok()?; + let (path_patterns, content) = strip_frontmatter_and_extract_paths(&raw); + let title = extract_title(path, &content); + + let normalized_active = active_path.map(normalize_glob_path); + let matched = if path_patterns.is_empty() { + true + } else if let Some(active) = normalized_active.as_deref() { + matches_any_pattern(active, &path_patterns) + } else { + false + }; + + if !matched { + return None; + } + + let trimmed_content = content.trim().to_string(); + if trimmed_content.is_empty() { + return None; + } + + Some(LoadedMemoryRule { + path: path.to_path_buf(), + title, + content: trimmed_content, + path_patterns, + matched: true, + }) +} + +fn strip_frontmatter_and_extract_paths(raw: &str) -> (Vec, String) { + if !raw.starts_with("---\n") && !raw.starts_with("---\r\n") { + return (Vec::new(), raw.to_string()); + } + + let mut lines = raw.lines(); + let Some(first) = lines.next() else { + return (Vec::new(), String::new()); + }; + if first.trim() != "---" { + return (Vec::new(), raw.to_string()); + } + + let mut frontmatter_lines = Vec::new(); + let mut body_lines = Vec::new(); + let mut in_frontmatter = true; + + for line in lines { + if in_frontmatter && line.trim() == "---" { + in_frontmatter = false; + continue; + } + + if in_frontmatter { + frontmatter_lines.push(line.to_string()); + } else { + body_lines.push(line.to_string()); + } + } + + if in_frontmatter { + // 未闭合 frontmatter,按普通 markdown 处理 + return (Vec::new(), raw.to_string()); + } + + let patterns = extract_paths_from_frontmatter(&frontmatter_lines); + (patterns, body_lines.join("\n")) +} + +fn extract_paths_from_frontmatter(lines: &[String]) -> Vec { + let mut patterns = Vec::new(); + let mut in_paths_block = false; + + for line in lines { + let trimmed = line.trim(); + if trimmed.is_empty() { + continue; + } + + if !in_paths_block { + if trimmed == "paths:" { + in_paths_block = true; + } + continue; + } + + if trimmed.starts_with('-') { + let value = trimmed.trim_start_matches('-').trim(); + if !value.is_empty() { + patterns.push(value.trim_matches('"').trim_matches('\'').to_string()); + } + continue; + } + + // paths 块结束 + break; + } + + patterns +} + +fn extract_title(path: &Path, content: &str) -> String { + for line in content.lines() { + let trimmed = line.trim(); + if let Some(title) = trimmed.strip_prefix("# ") { + let title = title.trim(); + if !title.is_empty() { + return title.to_string(); + } + } + } + + path.file_stem() + .and_then(|v| v.to_str()) + .unwrap_or("rule") + .to_string() +} + +fn normalize_glob_path(path: &str) -> String { + path.replace('\\', "/") +} + +fn matches_any_pattern(active_path: &str, patterns: &[String]) -> bool { + patterns.iter().any(|pattern| { + Pattern::new(pattern) + .map(|p| p.matches(active_path)) + .unwrap_or(false) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + use tempfile::TempDir; + + #[test] + fn should_load_unconditional_rules() { + let tmp = TempDir::new().expect("create temp dir"); + let rules_dir = tmp.path().join(".agents/rules"); + fs::create_dir_all(&rules_dir).expect("create rules dir"); + fs::write(rules_dir.join("general.md"), "# 通用规则\n- 保持简洁").expect("write rule"); + + let rules = load_rules(&rules_dir, None); + assert_eq!(rules.len(), 1); + assert!(rules[0].matched); + assert!(rules[0].content.contains("保持简洁")); + } + + #[test] + fn should_match_conditional_rule_by_paths() { + let tmp = TempDir::new().expect("create temp dir"); + let rules_dir = tmp.path().join(".agents/rules"); + fs::create_dir_all(&rules_dir).expect("create rules dir"); + fs::write( + rules_dir.join("api.md"), + r#"--- +paths: + - "src/api/**/*.ts" +--- +# API 规则 +- 必须做输入校验 +"#, + ) + .expect("write rule"); + + let matched = load_rules(&rules_dir, Some("src/api/user/index.ts")); + assert_eq!(matched.len(), 1); + assert!(matched[0].matched); + assert!(matched[0].content.contains("输入校验")); + + let not_matched = load_rules(&rules_dir, Some("src/ui/index.tsx")); + assert_eq!(not_matched.len(), 0); + } +} diff --git a/src-tauri/src/services/memory_source_resolver_service.rs b/src-tauri/src/services/memory_source_resolver_service.rs new file mode 100644 index 000000000..a7a0c5203 --- /dev/null +++ b/src-tauri/src/services/memory_source_resolver_service.rs @@ -0,0 +1,664 @@ +//! 记忆来源解析服务 +//! +//! 将配置中的记忆来源(AGENTS、规则、自动记忆等)统一解析为可观察结果与可注入提示词片段。 + +use crate::services::auto_memory_service::{get_auto_memory_index, resolve_auto_memory_root}; +use crate::services::memory_import_parser_service::{parse_memory_file, MemoryImportParseOptions}; +use crate::services::memory_rules_loader_service::load_rules; +use proxycast_core::config::{Config, MemoryConfig}; +use serde::{Deserialize, Serialize}; +use std::collections::HashSet; +use std::path::{Path, PathBuf}; + +/// 单个来源解析结果 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct EffectiveMemorySource { + /// 来源类型:managed_policy/project/user/local/rule/auto_memory/additional + pub kind: String, + /// 来源路径 + pub path: String, + /// 文件或目录是否存在 + pub exists: bool, + /// 是否被实际加载 + pub loaded: bool, + /// 内容行数(目录类来源为 0) + pub line_count: u32, + /// 导入展开后额外包含的文件数 + pub import_count: u32, + /// 告警信息 + pub warnings: Vec, + /// 预览(最多 300 字) + pub preview: Option, +} + +/// 来源解析总览 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct EffectiveMemorySourcesResponse { + pub working_dir: String, + pub total_sources: u32, + pub loaded_sources: u32, + pub follow_imports: bool, + pub import_max_depth: u8, + pub sources: Vec, +} + +/// 内部解析结果(包含可注入片段) +#[derive(Debug, Clone)] +pub struct MemorySourceResolution { + pub response: EffectiveMemorySourcesResponse, + pub prompt_segments: Vec, +} + +/// 解析有效记忆来源 +pub fn resolve_effective_sources( + config: &Config, + working_dir: &Path, + active_relative_path: Option<&str>, +) -> MemorySourceResolution { + let memory = &config.memory; + let options = MemoryImportParseOptions { + follow_imports: memory.resolve.follow_imports, + max_depth: memory.resolve.import_max_depth as usize, + }; + + let mut sources = Vec::new(); + let mut prompt_segments = Vec::new(); + let mut seen = HashSet::new(); + + // 1. managed policy + let managed_policy_path = memory + .sources + .managed_policy_path + .as_deref() + .map(|v| expand_path(v, Some(working_dir))) + .unwrap_or_else(default_managed_policy_path); + resolve_file_source( + "managed_policy", + &managed_policy_path, + true, + &options, + &mut seen, + &mut sources, + &mut prompt_segments, + ); + + // 2. user memory + let user_memory_path = memory + .sources + .user_memory_path + .as_deref() + .map(|v| expand_path(v, Some(working_dir))) + .unwrap_or_else(default_user_memory_path); + resolve_file_source( + "user_memory", + &user_memory_path, + true, + &options, + &mut seen, + &mut sources, + &mut prompt_segments, + ); + + // 3. project hierarchy memory + rules + let ancestors = collect_ancestor_dirs(working_dir); + for ancestor in &ancestors { + for rel in &memory.sources.project_memory_paths { + if rel.trim().is_empty() { + continue; + } + let candidate = ancestor.join(rel); + resolve_file_source( + "project_memory", + &candidate, + false, + &options, + &mut seen, + &mut sources, + &mut prompt_segments, + ); + } + + if let Some(project_local_rel) = memory + .sources + .project_local_memory_path + .as_deref() + .map(str::trim) + .filter(|v| !v.is_empty()) + { + let candidate = ancestor.join(project_local_rel); + resolve_file_source( + "project_local", + &candidate, + false, + &options, + &mut seen, + &mut sources, + &mut prompt_segments, + ); + } + + for rel in &memory.sources.project_rule_dirs { + if rel.trim().is_empty() { + continue; + } + let rule_dir = ancestor.join(rel); + resolve_rule_sources( + &rule_dir, + active_relative_path, + false, + &mut seen, + &mut sources, + &mut prompt_segments, + ); + } + } + + // 4. additional directories + if memory.resolve.load_additional_dirs_memory { + for additional in &memory.resolve.additional_dirs { + let additional_dir = expand_path(additional, Some(working_dir)); + for rel in &memory.sources.project_memory_paths { + if rel.trim().is_empty() { + continue; + } + let candidate = additional_dir.join(rel); + resolve_file_source( + "additional_memory", + &candidate, + false, + &options, + &mut seen, + &mut sources, + &mut prompt_segments, + ); + } + for rel in &memory.sources.project_rule_dirs { + if rel.trim().is_empty() { + continue; + } + let rule_dir = additional_dir.join(rel); + resolve_rule_sources( + &rule_dir, + active_relative_path, + false, + &mut seen, + &mut sources, + &mut prompt_segments, + ); + } + } + } + + // 5. auto memory + resolve_auto_memory_source( + memory, + working_dir, + &mut sources, + &mut prompt_segments, + &mut seen, + ); + + let loaded_sources = sources.iter().filter(|s| s.loaded).count() as u32; + let response = EffectiveMemorySourcesResponse { + working_dir: working_dir.to_string_lossy().to_string(), + total_sources: sources.len() as u32, + loaded_sources, + follow_imports: options.follow_imports, + import_max_depth: options.max_depth as u8, + sources, + }; + + MemorySourceResolution { + response, + prompt_segments, + } +} + +/// 构建可注入到 system prompt 的记忆来源片段 +pub fn build_memory_sources_prompt( + config: &Config, + working_dir: &Path, + active_relative_path: Option<&str>, + max_chars: usize, +) -> Option { + let resolution = resolve_effective_sources(config, working_dir, active_relative_path); + if resolution.prompt_segments.is_empty() { + return None; + } + + let mut output = String::from("【记忆来源补充指令】\n"); + output.push_str("以下内容来自配置化记忆来源,请优先遵循:\n"); + + let mut used = 0usize; + for segment in resolution.prompt_segments { + if segment.trim().is_empty() { + continue; + } + if used >= max_chars { + break; + } + let remaining = max_chars.saturating_sub(used); + let clipped = clip_text(&segment, remaining); + if clipped.trim().is_empty() { + continue; + } + output.push('\n'); + output.push_str(&clipped); + output.push('\n'); + used += clipped.chars().count(); + } + + if used == 0 { + None + } else { + Some(output.trim().to_string()) + } +} + +fn resolve_file_source( + kind: &str, + file_path: &Path, + include_missing: bool, + options: &MemoryImportParseOptions, + seen: &mut HashSet, + output: &mut Vec, + prompt_segments: &mut Vec, +) { + let normalized = normalize_path(file_path); + if !seen.insert(normalized.clone()) { + return; + } + + if !normalized.exists() || !normalized.is_file() { + if !include_missing { + return; + } + output.push(EffectiveMemorySource { + kind: kind.to_string(), + path: normalized.to_string_lossy().to_string(), + exists: false, + loaded: false, + line_count: 0, + import_count: 0, + warnings: Vec::new(), + preview: None, + }); + return; + } + + match parse_memory_file(&normalized, options) { + Ok(parsed) => { + let content = parsed.content.trim().to_string(); + let preview = if content.is_empty() { + None + } else { + Some(clip_text(&content, 300)) + }; + + let loaded = !content.is_empty(); + let line_count = if loaded { + content.lines().count() as u32 + } else { + 0 + }; + + output.push(EffectiveMemorySource { + kind: kind.to_string(), + path: normalized.to_string_lossy().to_string(), + exists: true, + loaded, + line_count, + import_count: parsed.imported_files.len() as u32, + warnings: parsed.warnings.clone(), + preview, + }); + + if loaded { + prompt_segments.push(format!( + "### {} ({})\n{}", + kind, + normalized.display(), + content + )); + } + } + Err(err) => { + output.push(EffectiveMemorySource { + kind: kind.to_string(), + path: normalized.to_string_lossy().to_string(), + exists: true, + loaded: false, + line_count: 0, + import_count: 0, + warnings: vec![err], + preview: None, + }); + } + } +} + +fn resolve_rule_sources( + rule_dir: &Path, + active_relative_path: Option<&str>, + include_missing: bool, + seen: &mut HashSet, + output: &mut Vec, + prompt_segments: &mut Vec, +) { + let normalized = normalize_path(rule_dir); + let dir_key = normalized.join("__rules_dir__"); + if !seen.insert(dir_key) { + return; + } + + if !normalized.exists() || !normalized.is_dir() { + if !include_missing { + return; + } + output.push(EffectiveMemorySource { + kind: "project_rules".to_string(), + path: normalized.to_string_lossy().to_string(), + exists: false, + loaded: false, + line_count: 0, + import_count: 0, + warnings: Vec::new(), + preview: None, + }); + return; + } + + let rules = load_rules(&normalized, active_relative_path); + if rules.is_empty() { + if !include_missing { + return; + } + output.push(EffectiveMemorySource { + kind: "project_rules".to_string(), + path: normalized.to_string_lossy().to_string(), + exists: true, + loaded: false, + line_count: 0, + import_count: 0, + warnings: vec!["规则目录存在,但未发现可用规则".to_string()], + preview: None, + }); + return; + } + + for rule in rules { + let normalized_rule = normalize_path(&rule.path); + if !seen.insert(normalized_rule.clone()) { + continue; + } + + let loaded = rule.matched && !rule.content.trim().is_empty(); + let mut warnings = Vec::new(); + if !rule.matched && !rule.path_patterns.is_empty() { + warnings.push(format!( + "规则 paths 未命中: {}", + rule.path_patterns.join(", ") + )); + } + output.push(EffectiveMemorySource { + kind: "project_rule".to_string(), + path: normalized_rule.to_string_lossy().to_string(), + exists: true, + loaded, + line_count: if loaded { + rule.content.lines().count() as u32 + } else { + 0 + }, + import_count: 0, + warnings, + preview: if loaded { + Some(clip_text(&rule.content, 300)) + } else { + None + }, + }); + + if loaded { + prompt_segments.push(format!( + "### 规则: {} ({})\n{}", + rule.title, + normalized_rule.display(), + rule.content + )); + } + } +} + +fn resolve_auto_memory_source( + memory_config: &MemoryConfig, + working_dir: &Path, + output: &mut Vec, + prompt_segments: &mut Vec, + seen: &mut HashSet, +) { + let auto_root = resolve_auto_memory_root(working_dir, &memory_config.auto); + let entry_name = memory_config.auto.entrypoint.trim(); + let entry_name = if entry_name.is_empty() { + "MEMORY.md" + } else { + entry_name + }; + let entry_path = normalize_path(&auto_root.join(entry_name)); + if !seen.insert(entry_path.clone()) { + return; + } + + let index = get_auto_memory_index(memory_config, working_dir); + match index { + Ok(idx) => { + let loaded = idx.entry_exists && !idx.preview_lines.is_empty(); + output.push(EffectiveMemorySource { + kind: "auto_memory".to_string(), + path: entry_path.to_string_lossy().to_string(), + exists: idx.entry_exists, + loaded, + line_count: idx.total_lines, + import_count: idx.items.len() as u32, + warnings: if !memory_config.auto.enabled { + vec!["自动记忆已关闭".to_string()] + } else { + Vec::new() + }, + preview: if loaded { + Some(clip_text(&idx.preview_lines.join("\n"), 300)) + } else { + None + }, + }); + + if loaded { + prompt_segments.push(format!( + "### auto_memory ({})\n{}", + entry_path.display(), + idx.preview_lines.join("\n") + )); + } + } + Err(err) => { + output.push(EffectiveMemorySource { + kind: "auto_memory".to_string(), + path: entry_path.to_string_lossy().to_string(), + exists: entry_path.exists(), + loaded: false, + line_count: 0, + import_count: 0, + warnings: vec![err], + preview: None, + }); + } + } +} + +fn collect_ancestor_dirs(start: &Path) -> Vec { + let mut dirs = Vec::new(); + let mut current = if start.is_file() { + start + .parent() + .unwrap_or_else(|| Path::new(".")) + .to_path_buf() + } else { + start.to_path_buf() + }; + let project_root = find_git_root(¤t); + let home_dir = dirs::home_dir(); + let mut depth = 0usize; + + loop { + dirs.push(current.clone()); + if let Some(root) = project_root.as_ref() { + if ¤t == root { + break; + } + } + if let Some(home) = home_dir.as_ref() { + if ¤t == home { + break; + } + } + // 兜底保护,避免跨层级扫描过深导致来源列表爆炸 + if depth >= 12 { + break; + } + if !current.pop() { + break; + } + depth += 1; + } + + dirs +} + +fn expand_path(path: &str, working_dir: Option<&Path>) -> PathBuf { + let trimmed = path.trim(); + if trimmed.starts_with("~/") { + if let Some(home) = dirs::home_dir() { + return home.join(trimmed.trim_start_matches("~/")); + } + } + + let p = PathBuf::from(trimmed); + if p.is_absolute() { + return p; + } + + if let Some(base) = working_dir { + return base.join(p); + } + p +} + +fn default_user_memory_path() -> PathBuf { + dirs::home_dir() + .unwrap_or_else(|| PathBuf::from(".")) + .join(".proxycast") + .join("AGENTS.md") +} + +fn default_managed_policy_path() -> PathBuf { + #[cfg(target_os = "macos")] + { + return PathBuf::from("/Library/Application Support/ProxyCast/AGENTS.md"); + } + #[cfg(target_os = "linux")] + { + return PathBuf::from("/etc/proxycast/AGENTS.md"); + } + #[cfg(target_os = "windows")] + { + return PathBuf::from("C:/Program Files/ProxyCast/AGENTS.md"); + } + #[allow(unreachable_code)] + PathBuf::from("/etc/proxycast/AGENTS.md") +} + +fn normalize_path(path: &Path) -> PathBuf { + path.canonicalize().unwrap_or_else(|_| path.to_path_buf()) +} + +fn find_git_root(start: &Path) -> Option { + let mut current = if start.is_file() { + start.parent()?.to_path_buf() + } else { + start.to_path_buf() + }; + + loop { + if current.join(".git").exists() { + return Some(current); + } + if !current.pop() { + return None; + } + } +} + +fn clip_text(text: &str, max_chars: usize) -> String { + if max_chars == 0 { + return String::new(); + } + let mut chars = text.chars(); + let clipped: String = chars.by_ref().take(max_chars).collect(); + if chars.next().is_some() { + format!("{clipped}...") + } else { + clipped + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::fs; + use tempfile::TempDir; + + #[test] + fn should_resolve_project_memory_and_rules() { + let tmp = TempDir::new().expect("create temp dir"); + let root = tmp.path(); + fs::create_dir_all(root.join(".agents/rules")).expect("create rules"); + fs::write(root.join("AGENTS.md"), "# 项目记忆\n- use rust").expect("write agents"); + fs::write(root.join(".agents/rules/general.md"), "# 规则\n- KISS").expect("write rule"); + + let mut cfg = Config::default(); + cfg.memory.enabled = true; + cfg.memory.sources.project_memory_paths = vec!["AGENTS.md".to_string()]; + cfg.memory.sources.project_rule_dirs = vec![".agents/rules".to_string()]; + cfg.memory.resolve.follow_imports = true; + cfg.memory.resolve.import_max_depth = 3; + + let resolved = resolve_effective_sources(&cfg, root, Some("src/main.rs")); + assert!(resolved.response.total_sources > 0); + assert!(resolved.response.loaded_sources > 0); + assert!(!resolved.prompt_segments.is_empty()); + } + + #[test] + fn should_support_additional_dirs_when_enabled() { + let tmp = TempDir::new().expect("create temp dir"); + let root = tmp.path().join("main"); + let ext = tmp.path().join("extra"); + fs::create_dir_all(&root).expect("create main"); + fs::create_dir_all(&ext).expect("create extra"); + fs::write(ext.join("AGENTS.md"), "extra memory").expect("write extra agents"); + + let mut cfg = Config::default(); + cfg.memory.enabled = true; + cfg.memory.sources.project_memory_paths = vec!["AGENTS.md".to_string()]; + cfg.memory.resolve.load_additional_dirs_memory = true; + cfg.memory.resolve.additional_dirs = vec![ext.to_string_lossy().to_string()]; + + let resolved = resolve_effective_sources(&cfg, &root, None); + let has_additional_loaded = resolved + .response + .sources + .iter() + .any(|s| s.kind == "additional_memory" && s.loaded); + assert!(has_additional_loaded); + } +} diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index c90b0ddd2..b125f138d 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -4,10 +4,15 @@ //! 本模块保留 Tauri 相关服务。 // 保留在主 crate 的 Tauri 相关服务 +pub mod auto_memory_service; pub mod conversation_statistics_service; pub mod execution_tracker_service; pub mod file_browser_service; pub mod heartbeat_service; +pub mod memory_import_parser_service; +pub mod memory_profile_prompt_service; +pub mod memory_rules_loader_service; +pub mod memory_source_resolver_service; pub mod sysinfo_service; pub mod update_check_service; pub mod update_window; diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 24971c85b..d0ccbcce7 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.72.0", + "version": "0.73.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/components/agent/chat/components/MessageList.tsx b/src/components/agent/chat/components/MessageList.tsx index 8201d1918..9330ff166 100644 --- a/src/components/agent/chat/components/MessageList.tsx +++ b/src/components/agent/chat/components/MessageList.tsx @@ -273,6 +273,28 @@ export const MessageList: React.FC = ({ )} + {msg.role === "assistant" && + !msg.isThinking && + msg.contextTrace && + msg.contextTrace.length > 0 && ( +
+ + 上下文轨迹 ({msg.contextTrace.length}) + +
+ {msg.contextTrace.map((step, index) => ( +
+ + {step.stage} + + : + {step.detail} +
+ ))} +
+
+ )} + {editingId !== msg.id && ( + + ) : ( + <> + {!card.available && ( + + )} + + + )} + + )} + + ))} + + +
diff --git a/src/components/memory/UnifiedMemoryPage.tsx b/src/components/memory/UnifiedMemoryPage.tsx index dfffb9fb6..65f215fe3 100644 --- a/src/components/memory/UnifiedMemoryPage.tsx +++ b/src/components/memory/UnifiedMemoryPage.tsx @@ -225,7 +225,7 @@ export default function UnifiedMemoryPage() {
  • 点击"刷新记忆列表"加载所有记忆
  • 点击"创建新记忆"添加测试数据
  • -
  • 点击"删除"按钮软删除记忆(数据不会真正删除)
  • +
  • 点击"删除"按钮会永久删除记忆(数据不可恢复)
  • 所有操作会在控制台输出详细日志
diff --git a/src/components/memory/UnifiedMemoryTest.tsx b/src/components/memory/UnifiedMemoryTest.tsx index 50f1437a5..e120c22e9 100644 --- a/src/components/memory/UnifiedMemoryTest.tsx +++ b/src/components/memory/UnifiedMemoryTest.tsx @@ -130,7 +130,7 @@ export default function UnifiedMemoryTest() {
  • 创建记忆会自动生成 ID
  • 创建成功后,复制 ID 用于其他操作
  • 所有操作都会在控制台输出详细结果
  • -
  • 删除是软删除,数据不会真正删除
  • +
  • 删除是永久删除,数据会被真正移除
  • diff --git a/src/components/memory/memoryLayerMetrics.test.ts b/src/components/memory/memoryLayerMetrics.test.ts new file mode 100644 index 000000000..b76338f74 --- /dev/null +++ b/src/components/memory/memoryLayerMetrics.test.ts @@ -0,0 +1,152 @@ +import { describe, expect, it } from "vitest"; +import { buildLayerMetrics } from "./memoryLayerMetrics"; + +describe("buildLayerMetrics", () => { + it("仅第一层有数据时应返回 1/3 可用", () => { + const result = buildLayerMetrics({ + unifiedTotalEntries: 3, + contextTotalEntries: 0, + projectId: null, + projectMemory: null, + }); + + const unifiedCard = result.cards.find((card) => card.key === "unified"); + const contextCard = result.cards.find((card) => card.key === "context"); + const projectCard = result.cards.find((card) => card.key === "project"); + + expect(unifiedCard?.available).toBe(true); + expect(contextCard?.available).toBe(false); + expect(projectCard?.available).toBe(false); + expect(result.readyLayers).toBe(1); + }); + + it("仅第二层有数据时应返回 1/3 可用", () => { + const result = buildLayerMetrics({ + unifiedTotalEntries: 0, + contextTotalEntries: 6, + projectId: null, + projectMemory: null, + }); + + const unifiedCard = result.cards.find((card) => card.key === "unified"); + const contextCard = result.cards.find((card) => card.key === "context"); + const projectCard = result.cards.find((card) => card.key === "project"); + + expect(unifiedCard?.available).toBe(false); + expect(contextCard?.available).toBe(true); + expect(projectCard?.available).toBe(false); + expect(result.readyLayers).toBe(1); + }); + + it("三层都有数据时应返回 3/3 可用", () => { + const result = buildLayerMetrics({ + unifiedTotalEntries: 12, + contextTotalEntries: 5, + projectId: "project-1", + projectMemory: { + characters: [ + { + id: "c1", + project_id: "project-1", + name: "主角", + aliases: [], + relationships: [], + is_main: true, + order: 0, + created_at: "2026-01-01T00:00:00Z", + updated_at: "2026-01-01T00:00:00Z", + }, + ], + world_building: { + project_id: "project-1", + description: "未来都市", + updated_at: "2026-01-01T00:00:00Z", + }, + style_guide: { + project_id: "project-1", + style: "克制叙事", + forbidden_words: [], + preferred_words: [], + updated_at: "2026-01-01T00:00:00Z", + }, + outline: [ + { + id: "o1", + project_id: "project-1", + title: "第一章", + order: 0, + expanded: true, + created_at: "2026-01-01T00:00:00Z", + updated_at: "2026-01-01T00:00:00Z", + }, + ], + }, + }); + + expect(result.totalLayers).toBe(3); + expect(result.readyLayers).toBe(3); + expect(result.cards[2]?.value).toBe(4); + expect(result.cards[2]?.available).toBe(true); + }); + + it("第三层部分维度已完善时也应判定为可用", () => { + const result = buildLayerMetrics({ + unifiedTotalEntries: 0, + contextTotalEntries: 0, + projectId: "project-1", + projectMemory: { + characters: [ + { + id: "c1", + project_id: "project-1", + name: "主角", + aliases: [], + relationships: [], + is_main: true, + order: 0, + created_at: "2026-01-01T00:00:00Z", + updated_at: "2026-01-01T00:00:00Z", + }, + ], + outline: [], + }, + }); + + const projectCard = result.cards.find((card) => card.key === "project"); + expect(projectCard?.available).toBe(true); + expect(projectCard?.value).toBe(1); + expect(result.readyLayers).toBe(1); + }); + + it("未选择项目时第三层应不可用并给出说明", () => { + const result = buildLayerMetrics({ + unifiedTotalEntries: 4, + contextTotalEntries: 2, + projectId: null, + projectMemory: null, + }); + + const projectCard = result.cards.find((card) => card.key === "project"); + expect(projectCard?.available).toBe(false); + expect(projectCard?.description).toContain("未选择项目"); + expect(result.readyLayers).toBe(2); + }); + + it("已选项目但无项目记忆内容时第三层仍不可用", () => { + const result = buildLayerMetrics({ + unifiedTotalEntries: 0, + contextTotalEntries: 1, + projectId: "project-2", + projectMemory: { + characters: [], + outline: [], + }, + }); + + const projectCard = result.cards.find((card) => card.key === "project"); + expect(projectCard?.value).toBe(0); + expect(projectCard?.available).toBe(false); + expect(projectCard?.description).toContain("还未填写"); + expect(result.readyLayers).toBe(1); + }); +}); diff --git a/src/components/memory/memoryLayerMetrics.ts b/src/components/memory/memoryLayerMetrics.ts new file mode 100644 index 000000000..651133ff9 --- /dev/null +++ b/src/components/memory/memoryLayerMetrics.ts @@ -0,0 +1,92 @@ +import type { ProjectMemory } from "@/lib/api/memory"; + +export interface LayerMetricsInput { + unifiedTotalEntries: number; + contextTotalEntries: number; + projectId: string | null; + projectMemory: ProjectMemory | null; +} + +export interface LayerCard { + key: "unified" | "context" | "project"; + title: string; + value: number; + unit: string; + available: boolean; + description: string; +} + +export interface LayerMetricsResult { + cards: LayerCard[]; + readyLayers: number; + totalLayers: number; +} + +function hasWorldBuilding(memory: ProjectMemory | null): boolean { + return !!memory?.world_building?.description?.trim(); +} + +function hasStyleGuide(memory: ProjectMemory | null): boolean { + return !!memory?.style_guide?.style?.trim(); +} + +function projectCoverageCount(memory: ProjectMemory | null): number { + if (!memory) { + return 0; + } + + let covered = 0; + if (memory.characters.length > 0) covered += 1; + if (hasWorldBuilding(memory)) covered += 1; + if (hasStyleGuide(memory)) covered += 1; + if (memory.outline.length > 0) covered += 1; + return covered; +} + +export function buildLayerMetrics(input: LayerMetricsInput): LayerMetricsResult { + const projectCoverage = projectCoverageCount(input.projectMemory); + const hasProjectSelection = !!input.projectId; + + const cards: LayerCard[] = [ + { + key: "unified", + title: "第一层:统一记忆", + value: input.unifiedTotalEntries, + unit: "条", + available: input.unifiedTotalEntries > 0, + description: + input.unifiedTotalEntries > 0 + ? "从历史对话沉淀出的结构化记忆。" + : "暂无沉淀结果,可点击“请求记忆分析”。", + }, + { + key: "context", + title: "第二层:上下文记忆", + value: input.contextTotalEntries, + unit: "条", + available: input.contextTotalEntries > 0, + description: + input.contextTotalEntries > 0 + ? "工作流文件记忆(计划/发现/进度)已生效。" + : "当前会话尚未形成文件记忆。", + }, + { + key: "project", + title: "第三层:项目记忆", + value: projectCoverage, + unit: "/4 维", + available: projectCoverage > 0, + description: !hasProjectSelection + ? "未选择项目,无法加载角色/世界观/风格/大纲。" + : projectCoverage > 0 + ? "项目级长期记忆已参与。" + : "项目已选择,但还未填写项目记忆内容。", + }, + ]; + + return { + readyLayers: cards.filter((card) => card.available).length, + totalLayers: cards.length, + cards, + }; +} diff --git a/src/components/settings-v2/_layout/index.tsx b/src/components/settings-v2/_layout/index.tsx index ec33c27e9..d3a21b270 100644 --- a/src/components/settings-v2/_layout/index.tsx +++ b/src/components/settings-v2/_layout/index.tsx @@ -16,6 +16,7 @@ import { CanvasBreadcrumbHeader } from "@/components/content-creator/canvas/shar // 外观设置 import { AppearanceSettings } from '../general/appearance'; import { ChatAppearanceSettings } from '../general/chat-appearance'; +import { MemorySettings } from "../general/memory"; // 网络代理 import { ProxySettings } from "../system/proxy"; // 安全与性能 @@ -23,8 +24,6 @@ import { SecurityPerformanceSettings } from "../system/security-performance"; // 心跳引擎 import { HeartbeatSettings } from "../system/heartbeat"; import { ExecutionTrackerSettings } from "../system/execution-tracker"; -// 外部工具 -import { ExternalToolsSettings } from "../system/external-tools"; // 实验功能 import { ExperimentalSettings } from "../system/experimental"; // 开发者 @@ -155,6 +154,14 @@ function renderSettingsContent(tab: SettingsTabs): ReactNode { ); + case SettingsTabs.Memory: + return ( + <> + + + + ); + // 智能体组 case SettingsTabs.Providers: return ( @@ -252,14 +259,6 @@ function renderSettingsContent(tab: SettingsTabs): ReactNode { ); - case SettingsTabs.ExternalTools: - return ( - <> - - - - ); - case SettingsTabs.Experimental: return ( <> diff --git a/src/components/settings-v2/general/appearance/index.tsx b/src/components/settings-v2/general/appearance/index.tsx index be0a1f63b..b3e47576d 100644 --- a/src/components/settings-v2/general/appearance/index.tsx +++ b/src/components/settings-v2/general/appearance/index.tsx @@ -6,7 +6,6 @@ import { useState, useEffect, useCallback } from "react"; import styled from "styled-components"; import { Moon, Sun, Monitor, Volume2, RotateCcw } from "lucide-react"; -import { cn } from "@/lib/utils"; import { getConfig, saveConfig, Config } from "@/hooks/useTauri"; import { useOnboardingState } from "@/components/onboarding"; import { diff --git a/src/components/settings-v2/general/memory/index.test.tsx b/src/components/settings-v2/general/memory/index.test.tsx new file mode 100644 index 000000000..3e503c7ba --- /dev/null +++ b/src/components/settings-v2/general/memory/index.test.tsx @@ -0,0 +1,242 @@ +import { act } from "react"; +import { createRoot, type Root } from "react-dom/client"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +const { + mockGetConfig, + mockSaveConfig, + mockGetMemoryOverview, + mockGetMemoryEffectiveSources, + mockGetMemoryAutoIndex, + mockToggleMemoryAuto, + mockUpdateMemoryAutoNote, + mockGetUnifiedMemoryStats, + mockGetProjectMemory, +} = vi.hoisted(() => ({ + mockGetConfig: vi.fn(), + mockSaveConfig: vi.fn(), + mockGetMemoryOverview: vi.fn(), + mockGetMemoryEffectiveSources: vi.fn(), + mockGetMemoryAutoIndex: vi.fn(), + mockToggleMemoryAuto: vi.fn(), + mockUpdateMemoryAutoNote: vi.fn(), + mockGetUnifiedMemoryStats: vi.fn(), + mockGetProjectMemory: vi.fn(), +})); + +vi.mock("@/hooks/useTauri", () => ({ + getConfig: mockGetConfig, + saveConfig: mockSaveConfig, + getMemoryOverview: mockGetMemoryOverview, + getMemoryEffectiveSources: mockGetMemoryEffectiveSources, + getMemoryAutoIndex: mockGetMemoryAutoIndex, + toggleMemoryAuto: mockToggleMemoryAuto, + updateMemoryAutoNote: mockUpdateMemoryAutoNote, +})); + +vi.mock("@/lib/api/unifiedMemory", () => ({ + getUnifiedMemoryStats: mockGetUnifiedMemoryStats, +})); + +vi.mock("@/lib/api/memory", () => ({ + getProjectMemory: mockGetProjectMemory, +})); + +vi.mock("@/lib/resourceProjectSelection", () => ({ + getStoredResourceProjectId: vi.fn(() => null), + onResourceProjectChange: vi.fn(() => () => {}), +})); + +vi.mock("@/components/memory/memoryLayerMetrics", () => ({ + buildLayerMetrics: vi.fn(() => ({ + cards: [ + { + key: "unified", + title: "第一层", + value: 1, + unit: "条", + available: true, + description: "ok", + }, + { + key: "context", + title: "第二层", + value: 0, + unit: "条", + available: false, + description: "wait", + }, + { + key: "project", + title: "第三层", + value: 0, + unit: "/4 维", + available: false, + description: "wait", + }, + ], + readyLayers: 1, + totalLayers: 3, + })), +})); + +import { MemorySettings } from "."; + +interface Mounted { + container: HTMLDivElement; + root: Root; +} + +const mounted: Mounted[] = []; + +function renderComponent(): HTMLDivElement { + const container = document.createElement("div"); + document.body.appendChild(container); + const root = createRoot(container); + act(() => { + root.render(); + }); + mounted.push({ container, root }); + return container; +} + +function findButton(container: HTMLElement, text: string): HTMLButtonElement { + const buttons = Array.from(container.querySelectorAll("button")); + const matched = buttons.find((button) => button.textContent?.includes(text)); + if (!matched) { + throw new Error(`未找到按钮: ${text}`); + } + return matched as HTMLButtonElement; +} + +async function flushEffects() { + await act(async () => { + await Promise.resolve(); + }); +} + +beforeEach(() => { + ( + globalThis as typeof globalThis & { + IS_REACT_ACT_ENVIRONMENT?: boolean; + } + ).IS_REACT_ACT_ENVIRONMENT = true; + + vi.clearAllMocks(); + + mockGetConfig.mockResolvedValue({ + memory: { + enabled: true, + max_entries: 1000, + retention_days: 30, + auto_cleanup: true, + profile: { + strengths: [], + explanation_style: [], + challenge_preference: [], + }, + auto: { + enabled: true, + entrypoint: "MEMORY.md", + max_loaded_lines: 200, + }, + resolve: { + additional_dirs: [], + follow_imports: true, + import_max_depth: 5, + load_additional_dirs_memory: false, + }, + sources: { + project_memory_paths: ["AGENTS.md"], + project_rule_dirs: [".agents/rules"], + user_memory_path: "~/.proxycast/AGENTS.md", + }, + }, + }); + + mockGetUnifiedMemoryStats.mockResolvedValue({ total_entries: 1 }); + mockGetMemoryOverview.mockResolvedValue({ + stats: { total_entries: 0, storage_used: 0, memory_count: 0 }, + categories: [], + entries: [], + }); + mockGetProjectMemory.mockResolvedValue(null); + mockGetMemoryEffectiveSources.mockResolvedValue({ + working_dir: "/tmp", + total_sources: 2, + loaded_sources: 1, + follow_imports: true, + import_max_depth: 5, + sources: [], + }); + mockGetMemoryAutoIndex.mockResolvedValue({ + enabled: true, + root_dir: "/tmp/memory", + entrypoint: "MEMORY.md", + max_loaded_lines: 200, + entry_exists: false, + total_lines: 0, + preview_lines: [], + items: [], + }); + mockToggleMemoryAuto.mockResolvedValue({ enabled: false }); + mockUpdateMemoryAutoNote.mockResolvedValue({ + enabled: true, + root_dir: "/tmp/memory", + entrypoint: "MEMORY.md", + max_loaded_lines: 200, + entry_exists: true, + total_lines: 1, + preview_lines: ["- test"], + items: [], + }); +}); + +afterEach(() => { + while (mounted.length > 0) { + const target = mounted.pop(); + if (!target) break; + act(() => { + target.root.unmount(); + }); + target.container.remove(); + } + vi.clearAllTimers(); +}); + +describe("MemorySettings", () => { + it("初始化时应加载来源与自动记忆索引", async () => { + renderComponent(); + await flushEffects(); + await flushEffects(); + + expect(mockGetMemoryEffectiveSources).toHaveBeenCalledTimes(1); + expect(mockGetMemoryAutoIndex).toHaveBeenCalledTimes(1); + }); + + it("点击立即关闭应调用 toggleMemoryAuto", async () => { + const container = renderComponent(); + await flushEffects(); + await flushEffects(); + + await act(async () => { + findButton(container, "立即关闭").click(); + }); + + expect(mockToggleMemoryAuto).toHaveBeenCalledWith(false); + }); + + it("未填写内容时写入自动记忆应阻止调用", async () => { + const container = renderComponent(); + await flushEffects(); + await flushEffects(); + + await act(async () => { + findButton(container, "写入自动记忆").click(); + }); + await flushEffects(); + + expect(mockUpdateMemoryAutoNote).not.toHaveBeenCalled(); + expect(container.textContent).toContain("请先输入要保存的自动记忆内容"); + }); +}); diff --git a/src/components/settings-v2/general/memory/index.tsx b/src/components/settings-v2/general/memory/index.tsx new file mode 100644 index 000000000..806e1a745 --- /dev/null +++ b/src/components/settings-v2/general/memory/index.tsx @@ -0,0 +1,922 @@ +import { useCallback, useEffect, useMemo, useState } from "react"; +import { Brain, Loader2, RefreshCw } from "lucide-react"; +import { cn } from "@/lib/utils"; +import { + getConfig, + getMemoryAutoIndex, + getMemoryEffectiveSources, + getMemoryOverview as getContextMemoryOverview, + saveConfig, + toggleMemoryAuto, + updateMemoryAutoNote, + type AutoMemoryIndexResponse, + type Config, + type EffectiveMemorySourcesResponse, + type MemoryAutoConfig, + type MemoryConfig, + type MemoryProfileConfig, + type MemoryResolveConfig, + type MemorySourcesConfig, +} from "@/hooks/useTauri"; +import { getUnifiedMemoryStats } from "@/lib/api/unifiedMemory"; +import { getProjectMemory } from "@/lib/api/memory"; +import { + getStoredResourceProjectId, + onResourceProjectChange, +} from "@/lib/resourceProjectSelection"; +import { + buildLayerMetrics, + type LayerMetricsResult, +} from "@/components/memory/memoryLayerMetrics"; + +const STATUS_OPTIONS = [ + "高中生", + "大学生/本科生", + "研究生", + "自学者/专业人士", + "其他", +]; + +const STRENGTH_OPTIONS = [ + "数学/逻辑推理", + "计算机科学/编程", + "自然科学(物理学、化学、生物学)", + "写作/阅读/人文", + "商业/经济学", + "没有——我还在探索中。", +]; + +const EXPLANATION_STYLE_OPTIONS = [ + "将晦涩难懂的概念变得直观易懂", + "先举例,后讲理论", + "概念结构与全局观", + "类比和隐喻", + "考试导向型讲解", + "我没有偏好——随机应变", +]; + +const CHALLENGE_OPTIONS = [ + "照本宣科——把所有细节都直接告诉我(我能应付)", + "一步一步地分解", + "先从简单的例子或类比入手", + "先解释重点和难点在哪里", + "多种解释/角度", +]; + +function normalizeProfile(profile?: MemoryProfileConfig): MemoryProfileConfig { + return { + current_status: profile?.current_status || undefined, + strengths: profile?.strengths || [], + explanation_style: profile?.explanation_style || [], + challenge_preference: profile?.challenge_preference || [], + }; +} + +function normalizeSources(sources?: MemorySourcesConfig): MemorySourcesConfig { + return { + managed_policy_path: sources?.managed_policy_path ?? undefined, + project_memory_paths: + sources?.project_memory_paths?.length && + sources.project_memory_paths.filter((item) => item.trim().length > 0) + ? sources.project_memory_paths + : ["AGENTS.md", ".agents/AGENTS.md"], + project_rule_dirs: + sources?.project_rule_dirs?.length && + sources.project_rule_dirs.filter((item) => item.trim().length > 0) + ? sources.project_rule_dirs + : [".agents/rules"], + user_memory_path: sources?.user_memory_path ?? "~/.proxycast/AGENTS.md", + project_local_memory_path: + sources?.project_local_memory_path ?? "AGENTS.local.md", + }; +} + +function normalizeAuto(auto?: MemoryAutoConfig): MemoryAutoConfig { + return { + enabled: auto?.enabled ?? true, + entrypoint: auto?.entrypoint || "MEMORY.md", + max_loaded_lines: auto?.max_loaded_lines ?? 200, + root_dir: auto?.root_dir ?? undefined, + }; +} + +function normalizeResolve(resolve?: MemoryResolveConfig): MemoryResolveConfig { + return { + additional_dirs: resolve?.additional_dirs || [], + follow_imports: resolve?.follow_imports ?? true, + import_max_depth: resolve?.import_max_depth ?? 5, + load_additional_dirs_memory: resolve?.load_additional_dirs_memory ?? false, + }; +} + +function normalizeMemoryConfig(memory?: MemoryConfig): MemoryConfig { + return { + enabled: memory?.enabled ?? true, + max_entries: memory?.max_entries ?? 1000, + retention_days: memory?.retention_days ?? 30, + auto_cleanup: memory?.auto_cleanup ?? true, + profile: normalizeProfile(memory?.profile), + sources: normalizeSources(memory?.sources), + auto: normalizeAuto(memory?.auto), + resolve: normalizeResolve(memory?.resolve), + }; +} + +function parseLines(input: string): string[] { + return input + .split("\n") + .map((line) => line.trim()) + .filter((line) => line.length > 0); +} + +interface MultiSelectSectionProps { + title: string; + subtitle?: string; + options: string[]; + value: string[]; + onToggle: (value: string) => void; +} + +function MultiSelectSection({ + title, + subtitle, + options, + value, + onToggle, +}: MultiSelectSectionProps) { + return ( +
    +
    +

    {title}

    + {subtitle &&

    {subtitle}

    } +
    + +
    + {options.map((option) => { + const selected = value.includes(option); + return ( + + ); + })} +
    +
    + ); +} + +export function MemorySettings() { + const [config, setConfig] = useState(null); + const [draft, setDraft] = useState(() => + normalizeMemoryConfig(), + ); + const [snapshot, setSnapshot] = useState(() => + normalizeMemoryConfig(), + ); + const [loading, setLoading] = useState(true); + const [saving, setSaving] = useState(false); + const [loadingLayerMetrics, setLoadingLayerMetrics] = useState(false); + const [loadingSourceState, setLoadingSourceState] = useState(false); + const [savingAutoNote, setSavingAutoNote] = useState(false); + const [projectId, setProjectId] = useState(() => + getStoredResourceProjectId({ includeLegacy: true }), + ); + const [layerMetrics, setLayerMetrics] = useState( + null, + ); + const [effectiveSources, setEffectiveSources] = + useState(null); + const [autoIndex, setAutoIndex] = useState( + null, + ); + const [autoTopic, setAutoTopic] = useState(""); + const [autoNote, setAutoNote] = useState(""); + const [message, setMessage] = useState(null); + + const loadLayerMetrics = useCallback( + async (targetProjectId?: string | null) => { + const currentProjectId = targetProjectId ?? projectId; + setLoadingLayerMetrics(true); + try { + const [unifiedStats, contextOverview, projectMemory] = await Promise.all([ + getUnifiedMemoryStats(), + getContextMemoryOverview(200).catch(() => null), + currentProjectId + ? getProjectMemory(currentProjectId).catch(() => null) + : Promise.resolve(null), + ]); + + setLayerMetrics( + buildLayerMetrics({ + unifiedTotalEntries: unifiedStats.total_entries, + contextTotalEntries: contextOverview?.stats.total_entries ?? 0, + projectId: currentProjectId ?? null, + projectMemory, + }), + ); + } catch (error) { + console.error("加载三层记忆状态失败:", error); + } finally { + setLoadingLayerMetrics(false); + } + }, + [projectId], + ); + + const loadSourceState = useCallback(async () => { + setLoadingSourceState(true); + try { + const [sources, index] = await Promise.all([ + getMemoryEffectiveSources().catch(() => null), + getMemoryAutoIndex().catch(() => null), + ]); + setEffectiveSources(sources); + setAutoIndex(index); + } finally { + setLoadingSourceState(false); + } + }, []); + + useEffect(() => { + const load = async () => { + setLoading(true); + try { + const nextConfig = await getConfig(); + const nextMemory = normalizeMemoryConfig(nextConfig.memory); + setConfig(nextConfig); + setDraft(nextMemory); + setSnapshot(nextMemory); + } catch (error) { + console.error("加载记忆设置失败:", error); + } finally { + setLoading(false); + } + }; + + load(); + }, []); + + useEffect(() => { + loadLayerMetrics(); + loadSourceState(); + }, [loadLayerMetrics, loadSourceState]); + + useEffect(() => { + return onResourceProjectChange((detail) => { + setProjectId(detail.projectId); + loadLayerMetrics(detail.projectId); + }); + }, [loadLayerMetrics]); + + const dirty = useMemo( + () => JSON.stringify(draft) !== JSON.stringify(snapshot), + [draft, snapshot], + ); + + const toggleMulti = ( + key: "strengths" | "explanation_style" | "challenge_preference", + option: string, + ) => { + setDraft((prev) => { + const profile = normalizeProfile(prev.profile); + const current = profile[key] || []; + const exists = current.includes(option); + return { + ...prev, + profile: { + ...profile, + [key]: exists + ? current.filter((item) => item !== option) + : [...current, option], + }, + }; + }); + }; + + const setStatus = (value: string) => { + setDraft((prev) => ({ + ...prev, + profile: { + ...normalizeProfile(prev.profile), + current_status: value, + }, + })); + }; + + const handleCancel = () => { + setDraft(snapshot); + setMessage("已恢复为上次保存内容"); + setTimeout(() => setMessage(null), 2500); + }; + + const handleSave = async () => { + if (!config) return; + setSaving(true); + try { + const updatedConfig: Config = { + ...config, + memory: draft, + }; + await saveConfig(updatedConfig); + setConfig(updatedConfig); + setSnapshot(draft); + setMessage("记忆设置已保存"); + setTimeout(() => setMessage(null), 2500); + await loadSourceState(); + } catch (error) { + console.error("保存记忆设置失败:", error); + setMessage("保存失败,请稍后重试"); + setTimeout(() => setMessage(null), 2500); + } finally { + setSaving(false); + } + }; + + const handleToggleAutoImmediately = async () => { + const current = normalizeAuto(draft.auto).enabled ?? true; + const next = !current; + try { + const result = await toggleMemoryAuto(next); + setDraft((prev) => ({ + ...prev, + auto: { + ...normalizeAuto(prev.auto), + enabled: result.enabled, + }, + })); + setSnapshot((prev) => ({ + ...prev, + auto: { + ...normalizeAuto(prev.auto), + enabled: result.enabled, + }, + })); + setMessage(result.enabled ? "自动记忆已开启" : "自动记忆已关闭"); + setTimeout(() => setMessage(null), 2500); + await loadSourceState(); + } catch (error) { + console.error("切换自动记忆失败:", error); + setMessage("切换自动记忆失败"); + setTimeout(() => setMessage(null), 2500); + } + }; + + const handleUpdateAutoNote = async () => { + const note = autoNote.trim(); + if (!note) { + setMessage("请先输入要保存的自动记忆内容"); + setTimeout(() => setMessage(null), 2500); + return; + } + + setSavingAutoNote(true); + try { + const index = await updateMemoryAutoNote(note, autoTopic.trim() || undefined); + setAutoIndex(index); + setAutoNote(""); + setMessage("已写入自动记忆"); + setTimeout(() => setMessage(null), 2500); + } catch (error) { + console.error("写入自动记忆失败:", error); + setMessage("写入自动记忆失败"); + setTimeout(() => setMessage(null), 2500); + } finally { + setSavingAutoNote(false); + } + }; + + if (loading) { + return ( +
    + + 正在加载记忆设置... +
    + ); + } + + const profile = normalizeProfile(draft.profile); + const sourcesConfig = normalizeSources(draft.sources); + const autoConfig = normalizeAuto(draft.auto); + const resolveConfig = normalizeResolve(draft.resolve); + + return ( +
    +
    +
    +
    + +
    +

    记忆

    +

    + 启用对话记忆功能,以便更好地理解上下文 +

    +
    +
    + +
    + + +
    +
    + +
    +
    启用记忆
    + + setDraft((prev) => ({ ...prev, enabled: event.target.checked })) + } + className="h-4 w-4 rounded border-gray-300" + /> +
    +
    + +
    +

    以下哪个选项最能形容你现在的状态?

    +
    + {STATUS_OPTIONS.map((option) => { + const selected = profile.current_status === option; + return ( + + ); + })} +
    +
    + + toggleMulti("strengths", option)} + /> + + toggleMulti("explanation_style", option)} + /> + + toggleMulti("challenge_preference", option)} + /> + +
    +
    +

    三层记忆可用性

    + +
    + + {layerMetrics ? ( + <> +
    + 已可用 {layerMetrics.readyLayers}/{layerMetrics.totalLayers} 层 +
    +
    + {layerMetrics.cards.map((card) => ( +
    +
    + {card.title} + + {card.available ? "已生效" : "待完善"} + +
    +

    + {card.description} +

    +
    + ))} +
    +

    + 第三层(项目记忆)的补全操作在「记忆」页面进行(支持一键初始化)。 +

    + + ) : ( +

    正在加载三层状态...

    + )} +
    + +
    +
    +

    记忆来源策略

    + +
    + +
    + + + + + + + +
    + +