feat: release v0.73.0 with full pending changes

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
coso
2026-02-27 16:37:53 +08:00
co-authored by Claude Opus 4.6
parent da0a4e0a6b
commit 0686737331
73 changed files with 7844 additions and 462 deletions
+29 -35
View File
@@ -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 安全、技能匹配等模块
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.72.0",
"version": "0.73.0",
"type": "module",
"repository": {
"type": "git",
+16 -15
View File
@@ -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",
+2 -2
View File
@@ -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"
+1
View File
@@ -22,6 +22,7 @@ chrono.workspace = true
dirs.workspace = true
uuid.workspace = true
thiserror.workspace = true
regex.workspace = true
[dev-dependencies]
tempfile.workspace = true
@@ -102,6 +102,7 @@ pub struct SessionConfigBuilder {
id: String,
max_turns: Option<u32>,
system_prompt: Option<String>,
include_context_trace: Option<bool>,
}
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,
}
}
}
+41 -4
View File
@@ -264,6 +264,10 @@ pub enum TauriAgentEvent {
#[serde(rename = "model_change")]
ModelChange { model: String, mode: String },
/// 上下文准备轨迹
#[serde(rename = "context_trace")]
ContextTrace { steps: Vec<TauriContextTraceStep> },
/// 完成(单次响应完成)
#[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<TauriAgentEvent> {
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!({
+593
View File
@@ -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<String>,
/// 匹配特定模式的内容
pub content_pattern: Option<String>,
}
/// 单个 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<String>,
}
/// Hook 执行上下文
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HookContext {
pub tool_name: Option<String>,
pub content: Option<String>,
pub metadata: HashMap<String, String>,
}
/// Hook 配置文件结构(旧格式)
#[derive(Debug, Deserialize)]
struct HookConfig {
hooks: Vec<HookDefinition>,
}
/// 新格式:按事件分组
#[derive(Debug, Deserialize)]
#[serde(untagged)]
enum HookConfigFormat {
/// 旧格式:{ "hooks": [...] }
Legacy(HookConfig),
/// 新格式:{ "hooks": { "BeforeToolCall": [...], ... } }
Grouped(GroupedHookConfig),
}
#[derive(Debug, Deserialize)]
struct GroupedHookConfig {
hooks: HashMap<HookEvent, Vec<GroupedHookEntry>>,
}
#[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<HookDefinition>,
}
impl HookManager {
pub fn new() -> Self {
Self { hooks: Vec::new() }
}
/// 从配置文件加载 hooks
pub fn load_from_config(config_path: &Path) -> Result<Self, Box<dyn std::error::Error>> {
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<HookResult> {
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::<serde_json::Value>(&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);
}
}
+6
View File
@@ -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};
+85 -3
View File
@@ -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<PathBuf>,
/// Skill 描述(注入到 system prompt)
skill_prompt: Option<String>,
}
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<Path>) -> 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("项目指令"));
}
}
@@ -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<PathBuf> {
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<InstructionLayer> {
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::<Vec<_>>()
.join("\n\n")
}
/// 从 path 向上查找包含 .git 的目录作为项目根
fn find_project_root(path: &Path) -> Option<PathBuf> {
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<InstructionLayer> {
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<PathBuf>) -> 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<PathBuf> {
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<InstructionLayer>,
cached_at: Instant,
}
// InstructionLayer 没有实现 Clone,手动实现缓存的 clone
impl CachedInstruction {
fn clone_layers(&self) -> Vec<InstructionLayer> {
self.layers.clone()
}
}
static CACHE: std::sync::LazyLock<RwLock<std::collections::HashMap<PathBuf, CachedInstruction>>> =
std::sync::LazyLock::new(|| RwLock::new(std::collections::HashMap::new()));
/// 带缓存的指令发现(TTL 默认 60 秒)
pub fn discover_instructions_cached(working_dir: &Path, ttl: Duration) -> Vec<InstructionLayer> {
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();
}
}
+5
View File
@@ -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::*;
@@ -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<String>,
pub is_readonly: bool,
pub reason: Option<String>,
}
/// 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<String> {
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);
}
}
@@ -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<dyn Fn(&serde_json::Value) + Send + Sync>;
// ---------------------------------------------------------------------------
// 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<String> = 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<AsterProviderConfig> {
let model = task.model.as_deref().unwrap_or(&self.default_model);
@@ -96,7 +215,7 @@ impl SubAgentExecutor for ProxyCastSubAgentExecutor {
context: &AgentContext,
) -> SchedulerResult<SubAgentResult> {
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<RwLock<Option<SubAgentScheduler<ProxyCastSubAgentExecutor>>>>,
/// 数据库连接
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<SchedulerConfig>) {
self.init_with_event_emitter(config, None).await;
@@ -174,7 +316,7 @@ impl ProxyCastScheduler {
config: Option<SchedulerConfig>,
event_emitter: Option<SchedulerEventEmitter>,
) {
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<SubAgentTask>,
parent_context: Option<&AgentContext>,
) -> SchedulerResult<SchedulerExecutionResult> {
self.execute_with_role(tasks, parent_context, self.default_role)
.await
}
/// 使用指定角色执行任务
///
/// 角色的工具限制会应用到每个任务上(与任务自身的 allowed_tools 取交集)。
pub async fn execute_with_role(
&self,
tasks: Vec<SubAgentTask>,
parent_context: Option<&AgentContext>,
role: SubAgentRole,
) -> SchedulerResult<SchedulerExecutionResult> {
let scheduler = self.scheduler.read().await;
let scheduler = scheduler
.as_ref()
.ok_or_else(|| SchedulerError::ContextError("调度器未初始化".to_string()))?;
// 应用角色工具限制
let tasks: Vec<SubAgentTask> = 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<String>,
/// SubAgent 角色
pub role: Option<String>,
}
impl From<SchedulerProgress> for SubAgentProgressEvent {
@@ -241,6 +410,143 @@ impl From<SchedulerProgress> 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()));
}
}
@@ -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<String, ToolPermissionMeta>,
auto_approve_level: ToolRiskLevel,
/// 会话内用户已允许的工具(tool_name → 允许次数)
session_allowed: HashMap<String, usize>,
/// 会话内用户已拒绝的工具
session_denied: HashSet<String>,
/// 动态权限检查器
dynamic_checker: Option<Box<dyn DynamicPermissionCheck>>,
}
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<dyn DynamicPermissionCheck>) {
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<bool> {
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<ToolPermissionMeta> {
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
);
}
}
+4 -2
View File
@@ -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,
+137
View File
@@ -1818,6 +1818,131 @@ pub struct ChatAppearanceConfig {
pub append_selected_text_to_recommendation: Option<bool>,
}
/// 记忆管理配置
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
pub struct MemoryProfileConfig {
/// 当前学习/工作状态(单选)
#[serde(default)]
pub current_status: Option<String>,
/// 擅长领域(多选)
#[serde(default)]
pub strengths: Vec<String>,
/// 偏好的解释风格(多选)
#[serde(default)]
pub explanation_style: Vec<String>,
/// 遇到难题时的偏好(多选)
#[serde(default)]
pub challenge_preference: Vec<String>,
}
/// 记忆来源配置
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct MemorySourcesConfig {
/// 组织级策略文件(可选)
#[serde(default, skip_serializing_if = "Option::is_none")]
pub managed_policy_path: Option<String>,
/// 项目级记忆文件相对路径列表(会按目录层级向上查找)
#[serde(default)]
pub project_memory_paths: Vec<String>,
/// 项目规则目录相对路径列表(会按目录层级向上查找)
#[serde(default)]
pub project_rule_dirs: Vec<String>,
/// 用户级记忆文件(可选)
#[serde(default, skip_serializing_if = "Option::is_none")]
pub user_memory_path: Option<String>,
/// 项目本地私有记忆文件(可选)
#[serde(default, skip_serializing_if = "Option::is_none")]
pub project_local_memory_path: Option<String>,
}
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/<project>/memory)
#[serde(default, skip_serializing_if = "Option::is_none")]
pub root_dir: Option<String>,
}
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<String>,
/// 是否跟随 @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<bool>,
/// 记忆偏好画像
#[serde(default)]
pub profile: Option<MemoryProfileConfig>,
/// 记忆来源配置
#[serde(default)]
pub sources: MemorySourcesConfig,
/// 自动记忆配置
#[serde(default)]
pub auto: MemoryAutoConfig,
/// 记忆解析行为配置
#[serde(default)]
pub resolve: MemoryResolveConfig,
}
/// 语音服务配置
@@ -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<usize>,
}
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<usize> {
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::<Vec<_>>()
.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<SummaryRequest> {
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<String> {
// 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<String> {
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<serde_json::Value> {
(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));
}
}
@@ -6,6 +6,7 @@ mod auth;
mod injection;
mod plugin;
mod provider;
pub mod registry;
mod routing;
mod telemetry;
mod traits;
@@ -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<dyn PipelineStep>,
phase: PipelinePhase,
priority: i32,
}
/// Pipeline 步骤注册表
pub struct StepRegistry {
core_steps: Vec<Box<dyn PipelineStep>>,
dynamic_steps: Vec<RegisteredStep>,
}
impl StepRegistry {
pub fn new(core_steps: Vec<Box<dyn PipelineStep>>) -> Self {
Self {
core_steps,
dynamic_steps: Vec::new(),
}
}
/// 运行时注册自定义步骤
pub fn register(&mut self, step: Box<dyn PipelineStep>, 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<dyn PipelineStep> {
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<String> {
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(&registry), 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(&registry),
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(&registry),
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(&registry), 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(&registry), vec!["core", "extra"]);
assert!(registry.unregister("extra"));
assert_eq!(step_names(&registry), 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(&registry), 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(&registry),
vec![
"init",
"rate_limit",
"auth",
"log_auth",
"custom_routing",
"provider",
"telemetry"
]
);
}
}
@@ -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<MemoryEntry> {
let mut entries = Vec::new();
let mut current_title: Option<String> = None;
let mut section_lines: Vec<String> = 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,
&section_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,
&section_lines,
) {
entries.push(entry);
}
}
entries
}
fn build_memory_entry(
&self,
session_id: &str,
file_type: MemoryFileType,
index: usize,
title: &str,
lines: &[String],
) -> Option<MemoryEntry> {
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::<Vec<_>>()
.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<String>, 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::<u8>().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::<Vec<_>>()
})
.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<i64> {
if let Ok(v) = value.parse::<i64>() {
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<String, HashSet<MemoryFileType>> = 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("应被归档的条目"));
}
}
+3 -1
View File
@@ -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};
@@ -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<String>,
/// 不触发条件描述列表
#[serde(default)]
pub do_not_trigger: Vec<String>,
}
/// Workflow 步骤定义
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WorkflowStep {
@@ -61,6 +72,8 @@ pub struct LoadedSkillDefinition {
pub allowed_tools: Option<Vec<String>>,
pub argument_hint: Option<String>,
pub when_to_use: Option<String>,
/// 结构化的自动触发条件配置
pub when_to_use_config: Option<SkillTriggerConfig>,
pub model: Option<String>,
pub provider: Option<String>,
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::<SkillTriggerConfig>(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,
@@ -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<LoadedSkillDefinition>,
}
/// 最低置信度阈值
const CONFIDENCE_THRESHOLD: f32 = 0.6;
impl SkillMatcher {
pub fn new(skills: Vec<LoadedSkillDefinition>) -> Self {
Self { skills }
}
/// 根据用户输入匹配最合适的 Skill
/// 返回按 confidence 降序排列的匹配结果(仅 >= 0.6)
pub fn match_skills(&self, user_input: &str) -> Vec<SkillMatch> {
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<String> = 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<String> {
// 中英文停用词
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<String> = 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<char> = 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("不要自动修复"));
}
}
+3 -1
View File
@@ -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;
+1 -1
View File
@@ -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,
};
+19 -1
View File
@@ -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<SchedulerConfig>) {
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<SubAgentTask>,
parent_context: Option<&AgentContext>,
role: SubAgentRole,
) -> SchedulerResult<SchedulerExecutionResult> {
self.inner
.execute_with_role(tasks, parent_context, role)
.await
}
/// 取消执行
pub async fn cancel(&self) {
self.inner.cancel().await;
+20 -1
View File
@@ -233,7 +233,7 @@ pub fn init_states(config: &Config) -> Result<AppStates, String> {
}
// 初始化上下文记忆服务
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<AppStates, String> {
})
}
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<PluginInstallerState, String> {
let db_path = database::get_db_path().map_err(|e| format!("获取数据库路径失败: {e}"))?;
+4
View File
@@ -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,
+21 -1
View File
@@ -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");
+7 -2
View File
@@ -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<String>,
system_prompt: Option<String>,
@@ -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();
+20 -9
View File
@@ -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;
+13 -1
View File
@@ -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(())
}
@@ -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<SkillExecutionResult, String> {
@@ -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,
+129 -2
View File
@@ -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<MemoryEntryPreview>,
}
/// 自动记忆开关响应
#[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<i64>,
to_timestamp: Option<i64>,
) -> Result<MemoryAnalysisResult, String> {
@@ -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<CleanupMemoryResult, String> {
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<String>,
active_relative_path: Option<String>,
) -> Result<EffectiveMemorySourcesResponse, String> {
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<String>,
) -> Result<AutoMemoryIndexResponse, String> {
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<MemoryAutoToggleResponse, String> {
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<String>,
note: String,
topic: Option<String>,
) -> Result<AutoMemoryIndexResponse, String> {
let config = global_config.config();
let resolved_working_dir = resolve_working_dir(working_dir)?;
update_auto_memory_note(
&config.memory,
&resolved_working_dir,
&note,
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<String>) -> Result<PathBuf, String> {
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<MemoryOverviewResponse, String> {
if !memory_dir.exists() {
return Ok(MemoryOverviewResponse {
+12 -1
View File
@@ -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<GeneratedPersona, String> {
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();
+23 -2
View File
@@ -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<SkillExecutionResult, String> {
// 发送步骤开始事件
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<SkillExecutionResult, String> {
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();
// 用户消息 = 原始输入 + 前序步骤的累积上下文
+13 -6
View File
@@ -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<SubAgentTask>,
config: Option<SchedulerConfig>,
role: Option<SubAgentRole>,
) -> Result<SchedulerExecutionResult, String> {
// 确保调度器已初始化
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 任务
+23 -4
View File
@@ -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<ChatSession> for SessionResponse {
pub async fn chat_create_session(
db: State<'_, DbConnection>,
agent_state: State<'_, AsterAgentState>,
config_manager: State<'_, GlobalConfigManagerState>,
request: CreateSessionRequest,
) -> Result<SessionResponse, String> {
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();
+30 -7
View File
@@ -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<i64>,
to_timestamp: Option<i64>,
) -> Result<MemoryAnalysisResult, String> {
@@ -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<PendingMemory>, 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<PendingMemory>, u32), String> {
let mut grouped: HashMap<String, Vec<MemorySourceCandidate>> = 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;
}
}
+3
View File
@@ -203,6 +203,7 @@ fn arb_config() -> impl Strategy<Value = Config> {
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<Value = Config> {
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<Value = Config> {
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 {
@@ -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<String>,
}
/// 自动记忆索引响应
#[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<String>,
pub items: Vec<AutoMemoryIndexItem>,
}
/// 读取自动记忆索引
pub fn get_auto_memory_index(
memory_config: &MemoryConfig,
working_dir: &Path,
) -> Result<AutoMemoryIndexResponse, String> {
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<String> = 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<AutoMemoryIndexResponse, String> {
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<AutoMemoryIndexItem> {
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<PathBuf> {
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"));
}
}
@@ -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<PathBuf>,
/// 解析过程中的告警
pub warnings: Vec<String>,
}
/// 读取并解析记忆文件(支持 @import)
pub fn parse_memory_file(
entry_path: &Path,
options: &MemoryImportParseOptions,
) -> Result<MemoryImportParseResult, String> {
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<PathBuf>,
imported_files: &mut Vec<PathBuf>,
warnings: &mut Vec<String>,
) -> Result<String, String> {
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!(
"<!-- import skipped: max depth reached ({}) -->\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!(
"<!-- import skipped: cyclic {} -->\n",
normalized_import.display()
));
continue;
}
if !normalized_import.exists() || !normalized_import.is_file() {
warnings.push(format!("导入目标不存在: {}", normalized_import.display()));
output.push_str(&format!(
"<!-- import missing: {} -->\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!(
"<!-- import begin: {} -->\n",
normalized_import.display()
));
output.push_str(imported_content.trim_end());
output.push('\n');
output.push_str(&format!(
"<!-- import end: {} -->\n",
normalized_import.display()
));
}
Ok(output)
}
fn resolve_import_path(import_target: &str, base_dir: Option<&Path>) -> Option<PathBuf> {
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("导入深度超限")));
}
}
@@ -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<String> {
let trimmed = input.trim();
if trimmed.is_empty() {
None
} else {
Some(trimmed.to_string())
}
}
fn normalize_list(items: &[String]) -> Vec<String> {
items
.iter()
.filter_map(|item| normalize_text(item))
.collect()
}
/// 构建记忆画像提示词
///
/// 仅在以下条件满足时返回:
/// - 记忆功能已启用
/// - 至少有一项画像字段有值
pub fn build_memory_profile_prompt(config: &Config) -> Option<String> {
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<String> = 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<String>,
config: &Config,
) -> Option<String> {
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);
}
}
@@ -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<String>,
/// 是否命中当前 active_path
pub matched: bool,
}
/// 从规则目录递归加载规则
///
/// - `rule_dir`: 规则目录(通常是 `.agents/rules`)
/// - `active_path`: 当前正在处理的相对路径(用于 paths 匹配)
pub fn load_rules(rule_dir: &Path, active_path: Option<&str>) -> Vec<LoadedMemoryRule> {
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<PathBuf>) {
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<LoadedMemoryRule> {
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>, 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<String> {
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);
}
}
@@ -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<String>,
/// 预览(最多 300 字)
pub preview: Option<String>,
}
/// 来源解析总览
#[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<EffectiveMemorySource>,
}
/// 内部解析结果(包含可注入片段)
#[derive(Debug, Clone)]
pub struct MemorySourceResolution {
pub response: EffectiveMemorySourcesResponse,
pub prompt_segments: Vec<String>,
}
/// 解析有效记忆来源
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<String> {
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<PathBuf>,
output: &mut Vec<EffectiveMemorySource>,
prompt_segments: &mut Vec<String>,
) {
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<PathBuf>,
output: &mut Vec<EffectiveMemorySource>,
prompt_segments: &mut Vec<String>,
) {
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<EffectiveMemorySource>,
prompt_segments: &mut Vec<String>,
seen: &mut HashSet<PathBuf>,
) {
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<PathBuf> {
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(&current);
let home_dir = dirs::home_dir();
let mut depth = 0usize;
loop {
dirs.push(current.clone());
if let Some(root) = project_root.as_ref() {
if &current == root {
break;
}
}
if let Some(home) = home_dir.as_ref() {
if &current == 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<PathBuf> {
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);
}
}
+5
View File
@@ -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;
+1 -1
View File
@@ -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",
@@ -273,6 +273,28 @@ export const MessageList: React.FC<MessageListProps> = ({
<TokenUsageDisplay usage={msg.usage} />
)}
{msg.role === "assistant" &&
!msg.isThinking &&
msg.contextTrace &&
msg.contextTrace.length > 0 && (
<details className="mt-3 rounded border border-border/60 bg-muted/20">
<summary className="cursor-pointer px-3 py-2 text-xs text-muted-foreground hover:text-foreground">
上下文轨迹 ({msg.contextTrace.length})
</summary>
<div className="border-t border-border/60 px-3 py-2 space-y-1.5">
{msg.contextTrace.map((step, index) => (
<div key={`${step.stage}-${index}`} className="text-xs">
<span className="font-medium text-foreground/90">
{step.stage}
</span>
<span className="text-muted-foreground">: </span>
<span className="text-muted-foreground">{step.detail}</span>
</div>
))}
</div>
</details>
)}
{editingId !== msg.id && (
<MessageActions className="message-actions">
<Button
@@ -647,6 +647,62 @@ describe("useAsterAgentChat action_required 渲染链路", () => {
harness.unmount();
}
});
it("收到 context_trace 事件后应写入当前 assistant 消息", async () => {
const workspaceId = "ws-context-trace";
seedSession(workspaceId, "session-context-trace");
const harness = mountHook(workspaceId);
let streamHandler:
| ((event: { payload: unknown }) => void)
| null = null;
mockSafeListen.mockImplementationOnce(async (_eventName, handler) => {
streamHandler = handler as (event: { payload: unknown }) => void;
return () => {
streamHandler = null;
};
});
try {
await flushEffects();
await act(async () => {
await harness
.getValue()
.sendMessage("检查轨迹", [], false, false, false, "react");
});
act(() => {
streamHandler?.({
payload: {
type: "context_trace",
steps: [
{
stage: "memory_injection",
detail: "query_len=8,injected=2",
},
{
stage: "memory_injection",
detail: "query_len=8,injected=2",
},
],
},
});
});
const assistantMessage = [...harness.getValue().messages]
.reverse()
.find((msg) => msg.role === "assistant");
expect(assistantMessage?.contextTrace).toBeDefined();
expect(assistantMessage?.contextTrace?.length).toBe(1);
expect(assistantMessage?.contextTrace?.[0]?.stage).toBe(
"memory_injection",
);
} finally {
harness.unmount();
}
});
});
describe("useAsterAgentChat 偏好持久化", () => {
@@ -23,6 +23,7 @@ import {
submitAsterElicitationResponse,
parseStreamEvent,
type StreamEvent,
type ContextTraceStep,
type AsterSessionInfo,
type AsterExecutionStrategy,
type ToolResultImage,
@@ -571,12 +572,25 @@ const mergeAdjacentAssistantMessages = (messages: Message[]): Message[] => {
toolCallMap.set(toolCall.id, toolCall);
}
const toolCalls = Array.from(toolCallMap.values());
const contextTrace = (() => {
const seen = new Set<string>();
const mergedSteps: ContextTraceStep[] = [];
for (const step of [...(previous.contextTrace || []), ...(current.contextTrace || [])]) {
const key = `${step.stage}::${step.detail}`;
if (!seen.has(key)) {
seen.add(key);
mergedSteps.push(step);
}
}
return mergedSteps;
})();
merged[merged.length - 1] = {
...previous,
content,
contentParts: contentParts.length > 0 ? contentParts : undefined,
toolCalls: toolCalls.length > 0 ? toolCalls : undefined,
contextTrace: contextTrace.length > 0 ? contextTrace : undefined,
timestamp: current.timestamp,
isThinking: false,
thinkingContent: undefined,
@@ -1654,6 +1668,38 @@ export function useAsterAgentChat(options: UseAsterAgentChatOptions) {
break;
}
case "context_trace":
if (!Array.isArray(data.steps) || data.steps.length === 0) {
break;
}
setMessages((prev) =>
prev.map((msg) => {
if (msg.id !== assistantMsgId) return msg;
const seen = new Set(
(msg.contextTrace || []).map(
(step) => `${step.stage}::${step.detail}`,
),
);
const nextSteps = [...(msg.contextTrace || [])];
for (const step of data.steps) {
const key = `${step.stage}::${step.detail}`;
if (!seen.has(key)) {
seen.add(key);
nextSteps.push(step);
}
}
return {
...msg,
contextTrace: nextSteps,
};
}),
);
break;
case "final_done":
setMessages((prev) =>
prev.map((msg) =>
+3
View File
@@ -1,4 +1,5 @@
import type { ToolCallState, TokenUsage } from "@/lib/api/agent";
import type { ContextTraceStep } from "@/lib/api/agent";
import { safeInvoke } from "@/lib/dev-bridge";
export interface MessageImage {
@@ -99,6 +100,8 @@ export interface Message {
* 否则回退到 content + toolCalls 渲染方式
*/
contentParts?: ContentPart[];
/** 上下文准备轨迹(可选) */
contextTrace?: ContextTraceStep[];
}
export interface ChatSession {
+234 -2
View File
@@ -25,6 +25,7 @@ import {
MessagesSquare,
RefreshCw,
Search,
Settings2,
Signature,
Trash2,
type LucideIcon,
@@ -32,13 +33,23 @@ import {
import { cn } from "@/lib/utils";
import { buildHomeAgentParams } from "@/lib/workspace/navigation";
import type { Page, PageParams } from "@/types/page";
import { SettingsTabs } from "@/types/settings";
import { CanvasBreadcrumbHeader } from "@/components/content-creator/canvas/shared/CanvasBreadcrumbHeader";
import {
getConfig,
getMemoryOverview as getContextMemoryOverview,
saveConfig,
type Config,
type MemoryConfig as TauriMemoryConfig,
} from "@/hooks/useTauri";
import {
createCharacter,
createOutlineNode,
getProjectMemory,
type ProjectMemory,
updateStyleGuide,
updateWorldBuilding,
} from "@/lib/api/memory";
import {
analyzeUnifiedMemories,
deleteUnifiedMemory,
@@ -49,6 +60,11 @@ import {
type UnifiedMemoryAnalysisResult,
type UnifiedMemoryStatsResponse,
} from "@/lib/api/unifiedMemory";
import {
getStoredResourceProjectId,
onResourceProjectChange,
} from "@/lib/resourceProjectSelection";
import { buildLayerMetrics } from "./memoryLayerMetrics";
type CategoryType = MemoryCategory;
type CategoryFilter = "all" | CategoryType;
type MemorySection = "home" | CategoryType;
@@ -85,6 +101,10 @@ interface MemoryOverviewResponse {
entries: MemoryEntryPreview[];
}
interface ContextLayerStats {
total_entries: number;
}
const CATEGORY_META: Record<
CategoryType,
{ label: string; description: string; icon: LucideIcon }
@@ -580,10 +600,18 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
);
const [overview, setOverview] = useState<MemoryOverviewResponse | null>(null);
const [contextLayerStats, setContextLayerStats] =
useState<ContextLayerStats | null>(null);
const [projectId, setProjectId] = useState<string | null>(() =>
getStoredResourceProjectId({ includeLegacy: true }),
);
const [projectMemory, setProjectMemory] = useState<ProjectMemory | null>(null);
const [loading, setLoading] = useState(true);
const [refreshing, setRefreshing] = useState(false);
const [saving, setSaving] = useState(false);
const [analyzing, setAnalyzing] = useState(false);
const [initializingProjectMemory, setInitializingProjectMemory] =
useState(false);
const [deletingEntryId, setDeletingEntryId] = useState<string | null>(null);
const [activeSection, setActiveSection] = useState<MemorySection>("home");
@@ -641,6 +669,16 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
const entries = useMemo(() => overview?.entries ?? [], [overview]);
const hasMemoryData = stats.total_entries > 0;
const layerMetrics = useMemo(
() =>
buildLayerMetrics({
unifiedTotalEntries: stats.total_entries,
contextTotalEntries: contextLayerStats?.total_entries ?? 0,
projectId,
projectMemory,
}),
[contextLayerStats?.total_entries, projectId, projectMemory, stats.total_entries],
);
const activeCategoryFilter: CategoryFilter =
activeSection === "home" ? "all" : activeSection;
@@ -694,7 +732,8 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
}, []);
const loadOverview = useCallback(async () => {
const [statsResult, memories] = await Promise.all([
const [statsResult, memories, contextOverviewResult, projectMemoryResult] =
await Promise.all([
getUnifiedMemoryStats(),
listUnifiedMemories({
archived: false,
@@ -702,6 +741,16 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
order: "desc",
limit: 1000,
}),
getContextMemoryOverview(200).catch((error) => {
console.warn("加载上下文记忆总览失败:", error);
return null;
}),
projectId
? getProjectMemory(projectId).catch((error) => {
console.warn("加载项目记忆失败:", error);
return null;
})
: Promise.resolve(null),
]);
const normalizedStats: MemoryOverviewResponse = {
@@ -715,7 +764,13 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
};
setOverview(normalizedStats);
}, []);
setContextLayerStats(
contextOverviewResult
? { total_entries: contextOverviewResult.stats.total_entries }
: null,
);
setProjectMemory(projectMemoryResult);
}, [projectId]);
const loadAll = useCallback(async () => {
setLoading(true);
@@ -733,6 +788,12 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
loadAll();
}, [loadAll]);
useEffect(() => {
return onResourceProjectChange((detail) => {
setProjectId(detail.projectId);
});
}, []);
useEffect(() => {
const handleKeyDown = (event: KeyboardEvent) => {
if (event.metaKey || event.ctrlKey || event.altKey) {
@@ -782,6 +843,90 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
}
}, [loadOverview, showMessage]);
const handleBootstrapProjectMemory = useCallback(async () => {
if (!projectId) {
showMessage("error", "请先在资源页或项目页选择一个项目");
return;
}
if (initializingProjectMemory) {
return;
}
const hasCharacters = (projectMemory?.characters.length ?? 0) > 0;
const hasWorldBuilding = !!projectMemory?.world_building?.description?.trim();
const hasStyleGuide = !!projectMemory?.style_guide?.style?.trim();
const hasOutline = (projectMemory?.outline.length ?? 0) > 0;
if (hasCharacters && hasWorldBuilding && hasStyleGuide && hasOutline) {
showMessage("success", "第三层项目记忆已完善");
return;
}
setInitializingProjectMemory(true);
try {
const tasks: Promise<unknown>[] = [];
if (!hasCharacters) {
tasks.push(
createCharacter({
project_id: projectId,
name: "默认主角",
description: "待补充角色设定",
is_main: true,
}),
);
}
if (!hasWorldBuilding) {
tasks.push(
updateWorldBuilding(projectId, {
description: "待补充世界观背景与规则",
}),
);
}
if (!hasStyleGuide) {
tasks.push(
updateStyleGuide(projectId, {
style: "待补充写作风格与语气",
}),
);
}
if (!hasOutline) {
tasks.push(
createOutlineNode({
project_id: projectId,
title: "第一章",
content: "待补充章节内容",
}),
);
}
if (tasks.length > 0) {
await Promise.all(tasks);
}
await loadOverview();
showMessage("success", "已初始化第三层项目记忆,请按需继续完善");
} catch (error) {
console.error("初始化项目记忆失败:", error);
showMessage("error", "初始化项目记忆失败,请稍后重试");
} finally {
setInitializingProjectMemory(false);
}
}, [
initializingProjectMemory,
loadOverview,
projectId,
projectMemory?.characters.length,
projectMemory?.outline.length,
projectMemory?.style_guide?.style,
projectMemory?.world_building?.description,
showMessage,
]);
const handleAnalyze = useCallback(async () => {
if (
analysisFromDate &&
@@ -988,6 +1133,16 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
</div>
<div className="flex items-center gap-2">
<button
onClick={() =>
onNavigate?.("settings", { tab: SettingsTabs.Memory })
}
className="inline-flex items-center gap-1.5 rounded border px-3 py-1.5 text-sm hover:bg-muted"
>
<Settings2 className="h-3.5 w-3.5" />
记忆设置
</button>
<button
onClick={handleRefresh}
disabled={refreshing || analyzing || loading}
@@ -1055,6 +1210,83 @@ export function MemoryPage({ onNavigate }: MemoryPageProps) {
</div>
</div>
<div className="rounded-xl border p-4 bg-card">
<div className="flex flex-wrap items-center justify-between gap-2 mb-3">
<div className="text-sm font-medium">三层记忆可用性</div>
<div className="text-xs text-muted-foreground">
已可用 {layerMetrics.readyLayers}/{layerMetrics.totalLayers} 层
</div>
</div>
<div className="grid grid-cols-1 md:grid-cols-3 gap-3">
{layerMetrics.cards.map((card) => (
<div
key={card.key}
className="rounded-lg border bg-background/60 p-3"
>
<div className="flex items-center justify-between gap-2">
<div className="text-xs text-muted-foreground">
{card.title}
</div>
<span
className={cn(
"text-[10px] px-2 py-0.5 rounded-full border",
card.available
? "text-green-700 border-green-200 bg-green-50 dark:text-green-400 dark:border-green-800 dark:bg-green-900/20"
: "text-muted-foreground border-muted",
)}
>
{card.available ? "已生效" : "待完善"}
</span>
</div>
<div className="text-xl font-semibold text-primary mt-1">
{card.value}
<span className="text-sm text-muted-foreground ml-1">
{card.unit}
</span>
</div>
<div className="text-xs text-muted-foreground mt-1 leading-relaxed">
{card.description}
</div>
{card.key === "project" && (
<div className="mt-2 flex flex-wrap gap-2">
{!projectId ? (
<button
onClick={() => onNavigate?.("projects")}
className="rounded border px-2 py-1 text-[11px] hover:bg-muted"
>
去选择项目
</button>
) : (
<>
{!card.available && (
<button
onClick={handleBootstrapProjectMemory}
disabled={initializingProjectMemory}
className="rounded border px-2 py-1 text-[11px] hover:bg-muted disabled:opacity-60"
>
{initializingProjectMemory
? "初始化中..."
: "一键初始化"}
</button>
)}
<button
onClick={() =>
onNavigate?.("project-detail", { projectId })
}
className="rounded border px-2 py-1 text-[11px] hover:bg-muted"
>
前往项目记忆
</button>
</>
)}
</div>
)}
</div>
))}
</div>
</div>
<div className="rounded-lg border p-4 space-y-3">
<div className="flex items-center gap-2 text-sm font-medium">
<Database className="h-4 w-4 text-muted-foreground" />
+1 -1
View File
@@ -225,7 +225,7 @@ export default function UnifiedMemoryPage() {
<ul style={{ margin: 0, paddingLeft: "20px", fontSize: "14px", color: "#333" }}>
<li>点击"刷新记忆列表"加载所有记忆</li>
<li>点击"创建新记忆"添加测试数据</li>
<li>点击"删除"按钮软删除记忆(数据不会真正删除)</li>
<li>点击"删除"按钮会永久删除记忆(数据不可恢复)</li>
<li>所有操作会在控制台输出详细日志</li>
</ul>
</div>
+1 -1
View File
@@ -130,7 +130,7 @@ export default function UnifiedMemoryTest() {
<li>创建记忆会自动生成 ID</li>
<li>创建成功后,复制 ID 用于其他操作</li>
<li>所有操作都会在控制台输出详细结果</li>
<li>删除是软删除,数据不会真正删除</li>
<li>删除是永久删除,数据会被真正移除</li>
</ul>
</div>
</div>
@@ -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);
});
});
@@ -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,
};
}
+9 -10
View File
@@ -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 (
<>
<SettingHeader title="记忆" />
<MemorySettings />
</>
);
// 智能体组
case SettingsTabs.Providers:
return (
@@ -252,14 +259,6 @@ function renderSettingsContent(tab: SettingsTabs): ReactNode {
</>
);
case SettingsTabs.ExternalTools:
return (
<>
<SettingHeader title="外部工具" />
<ExternalToolsSettings />
</>
);
case SettingsTabs.Experimental:
return (
<>
@@ -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 {
@@ -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(<MemorySettings />);
});
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("请先输入要保存的自动记忆内容");
});
});
@@ -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 (
<div className="rounded-lg border p-4 space-y-3">
<div>
<h3 className="text-sm font-medium">{title}</h3>
{subtitle && <p className="text-xs text-muted-foreground">{subtitle}</p>}
</div>
<div className="flex flex-wrap gap-2">
{options.map((option) => {
const selected = value.includes(option);
return (
<button
key={option}
type="button"
onClick={() => onToggle(option)}
className={cn(
"rounded-md border px-3 py-1.5 text-xs transition-colors",
selected
? "border-primary bg-primary/10 text-primary"
: "hover:bg-muted",
)}
>
{option}
</button>
);
})}
</div>
</div>
);
}
export function MemorySettings() {
const [config, setConfig] = useState<Config | null>(null);
const [draft, setDraft] = useState<MemoryConfig>(() =>
normalizeMemoryConfig(),
);
const [snapshot, setSnapshot] = useState<MemoryConfig>(() =>
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<string | null>(() =>
getStoredResourceProjectId({ includeLegacy: true }),
);
const [layerMetrics, setLayerMetrics] = useState<LayerMetricsResult | null>(
null,
);
const [effectiveSources, setEffectiveSources] =
useState<EffectiveMemorySourcesResponse | null>(null);
const [autoIndex, setAutoIndex] = useState<AutoMemoryIndexResponse | null>(
null,
);
const [autoTopic, setAutoTopic] = useState("");
const [autoNote, setAutoNote] = useState("");
const [message, setMessage] = useState<string | null>(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 (
<div className="flex items-center justify-center py-12 text-muted-foreground">
<Loader2 className="mr-2 h-5 w-5 animate-spin" />
正在加载记忆设置...
</div>
);
}
const profile = normalizeProfile(draft.profile);
const sourcesConfig = normalizeSources(draft.sources);
const autoConfig = normalizeAuto(draft.auto);
const resolveConfig = normalizeResolve(draft.resolve);
return (
<div className="space-y-4 max-w-4xl">
<div className="rounded-lg border p-4">
<div className="flex items-start justify-between gap-4">
<div className="flex items-start gap-2">
<Brain className="h-4 w-4 text-muted-foreground mt-0.5" />
<div>
<h3 className="text-sm font-medium">记忆</h3>
<p className="text-xs text-muted-foreground mt-1">
启用对话记忆功能,以便更好地理解上下文
</p>
</div>
</div>
<div className="flex items-center gap-2">
<button
type="button"
onClick={handleCancel}
disabled={!dirty || saving}
className="rounded border px-3 py-1.5 text-xs hover:bg-muted disabled:opacity-60"
>
取消
</button>
<button
type="button"
onClick={handleSave}
disabled={!dirty || saving}
className="rounded bg-primary px-3 py-1.5 text-xs text-primary-foreground hover:opacity-90 disabled:opacity-60"
>
{saving ? "保存中..." : "保存"}
</button>
</div>
</div>
<div className="mt-4 flex items-center justify-between rounded-md border p-3">
<div className="text-sm">启用记忆</div>
<input
type="checkbox"
checked={draft.enabled}
onChange={(event) =>
setDraft((prev) => ({ ...prev, enabled: event.target.checked }))
}
className="h-4 w-4 rounded border-gray-300"
/>
</div>
</div>
<div className="rounded-lg border p-4 space-y-3">
<h3 className="text-sm font-medium">以下哪个选项最能形容你现在的状态?</h3>
<div className="flex flex-wrap gap-2">
{STATUS_OPTIONS.map((option) => {
const selected = profile.current_status === option;
return (
<button
key={option}
type="button"
onClick={() => setStatus(option)}
className={cn(
"rounded-md border px-3 py-1.5 text-xs transition-colors",
selected
? "border-primary bg-primary/10 text-primary"
: "hover:bg-muted",
)}
>
{option}
</button>
);
})}
</div>
</div>
<MultiSelectSection
title="你觉得自己有哪些方面比较擅长?"
subtitle="(可多选)"
options={STRENGTH_OPTIONS}
value={profile.strengths || []}
onToggle={(option) => toggleMulti("strengths", option)}
/>
<MultiSelectSection
title="我解释事情时通常更喜欢:"
subtitle="(可多选)"
options={EXPLANATION_STYLE_OPTIONS}
value={profile.explanation_style || []}
onToggle={(option) => toggleMulti("explanation_style", option)}
/>
<MultiSelectSection
title="当你遇到难题/概念时,你更倾向于:"
subtitle="(可多选)"
options={CHALLENGE_OPTIONS}
value={profile.challenge_preference || []}
onToggle={(option) => toggleMulti("challenge_preference", option)}
/>
<div className="rounded-lg border p-4 space-y-3 bg-muted/20">
<div className="flex items-center justify-between gap-2">
<h3 className="text-sm font-medium">三层记忆可用性</h3>
<button
type="button"
onClick={() => loadLayerMetrics()}
disabled={loadingLayerMetrics}
className="inline-flex items-center gap-1 rounded border px-2 py-1 text-[11px] hover:bg-muted disabled:opacity-60"
>
<RefreshCw className={cn("h-3 w-3", loadingLayerMetrics && "animate-spin")} />
刷新
</button>
</div>
{layerMetrics ? (
<>
<div className="text-xs text-muted-foreground">
已可用 {layerMetrics.readyLayers}/{layerMetrics.totalLayers} 层
</div>
<div className="space-y-2">
{layerMetrics.cards.map((card) => (
<div
key={card.key}
className="rounded border bg-background/60 px-3 py-2"
>
<div className="flex items-center justify-between gap-2">
<span className="text-xs font-medium">{card.title}</span>
<span
className={cn(
"rounded-full border px-2 py-0.5 text-[10px]",
card.available
? "text-green-700 border-green-200 bg-green-50"
: "text-muted-foreground border-muted",
)}
>
{card.available ? "已生效" : "待完善"}
</span>
</div>
<p className="mt-1 text-[11px] text-muted-foreground leading-relaxed">
{card.description}
</p>
</div>
))}
</div>
<p className="text-xs text-muted-foreground">
第三层(项目记忆)的补全操作在「记忆」页面进行(支持一键初始化)。
</p>
</>
) : (
<p className="text-xs text-muted-foreground">正在加载三层状态...</p>
)}
</div>
<div className="rounded-lg border p-4 space-y-4">
<div className="flex items-center justify-between gap-2">
<h3 className="text-sm font-medium">记忆来源策略</h3>
<button
type="button"
onClick={() => loadSourceState()}
disabled={loadingSourceState}
className="inline-flex items-center gap-1 rounded border px-2 py-1 text-[11px] hover:bg-muted disabled:opacity-60"
>
<RefreshCw className={cn("h-3 w-3", loadingSourceState && "animate-spin")} />
刷新来源
</button>
</div>
<div className="grid grid-cols-1 md:grid-cols-2 gap-3">
<label className="space-y-1">
<span className="text-xs text-muted-foreground">组织策略文件</span>
<input
type="text"
value={sourcesConfig.managed_policy_path || ""}
onChange={(event) =>
setDraft((prev) => ({
...prev,
sources: {
...normalizeSources(prev.sources),
managed_policy_path: event.target.value || undefined,
},
}))
}
className="w-full rounded border bg-background px-2 py-1.5 text-xs"
placeholder="例如 /Library/Application Support/ProxyCast/AGENTS.md"
/>
</label>
<label className="space-y-1">
<span className="text-xs text-muted-foreground">用户记忆文件</span>
<input
type="text"
value={sourcesConfig.user_memory_path || ""}
onChange={(event) =>
setDraft((prev) => ({
...prev,
sources: {
...normalizeSources(prev.sources),
user_memory_path: event.target.value || undefined,
},
}))
}
className="w-full rounded border bg-background px-2 py-1.5 text-xs"
placeholder="例如 ~/.proxycast/AGENTS.md"
/>
</label>
<label className="space-y-1">
<span className="text-xs text-muted-foreground">项目本地私有文件</span>
<input
type="text"
value={sourcesConfig.project_local_memory_path || ""}
onChange={(event) =>
setDraft((prev) => ({
...prev,
sources: {
...normalizeSources(prev.sources),
project_local_memory_path: event.target.value || undefined,
},
}))
}
className="w-full rounded border bg-background px-2 py-1.5 text-xs"
placeholder="例如 AGENTS.local.md"
/>
</label>
<label className="space-y-1">
<span className="text-xs text-muted-foreground">最大导入深度</span>
<input
type="number"
min={1}
max={20}
value={resolveConfig.import_max_depth ?? 5}
onChange={(event) => {
const value = Number(event.target.value);
setDraft((prev) => ({
...prev,
resolve: {
...normalizeResolve(prev.resolve),
import_max_depth: Number.isFinite(value)
? Math.max(1, Math.min(20, value))
: 5,
},
}));
}}
className="w-full rounded border bg-background px-2 py-1.5 text-xs"
/>
</label>
</div>
<label className="space-y-1 block">
<span className="text-xs text-muted-foreground">
项目记忆文件(每行一个相对路径)
</span>
<textarea
value={(sourcesConfig.project_memory_paths || []).join("\n")}
onChange={(event) =>
setDraft((prev) => ({
...prev,
sources: {
...normalizeSources(prev.sources),
project_memory_paths: parseLines(event.target.value),
},
}))
}
className="w-full rounded border bg-background px-2 py-1.5 text-xs min-h-20"
/>
</label>
<label className="space-y-1 block">
<span className="text-xs text-muted-foreground">
项目规则目录(每行一个相对路径)
</span>
<textarea
value={(sourcesConfig.project_rule_dirs || []).join("\n")}
onChange={(event) =>
setDraft((prev) => ({
...prev,
sources: {
...normalizeSources(prev.sources),
project_rule_dirs: parseLines(event.target.value),
},
}))
}
className="w-full rounded border bg-background px-2 py-1.5 text-xs min-h-16"
/>
</label>
<label className="space-y-1 block">
<span className="text-xs text-muted-foreground">
额外目录(每行一个绝对路径,可添加 aster-rust 等外部仓库)
</span>
<textarea
value={(resolveConfig.additional_dirs || []).join("\n")}
onChange={(event) =>
setDraft((prev) => ({
...prev,
resolve: {
...normalizeResolve(prev.resolve),
additional_dirs: parseLines(event.target.value),
},
}))
}
className="w-full rounded border bg-background px-2 py-1.5 text-xs min-h-16"
placeholder="例如 /Users/coso/Documents/dev/ai/astercloud/aster-rust"
/>
</label>
<div className="grid grid-cols-1 md:grid-cols-2 gap-3">
<label className="flex items-center justify-between rounded border px-3 py-2 text-xs">
<span>跟随 @import</span>
<input
type="checkbox"
checked={resolveConfig.follow_imports ?? true}
onChange={(event) =>
setDraft((prev) => ({
...prev,
resolve: {
...normalizeResolve(prev.resolve),
follow_imports: event.target.checked,
},
}))
}
className="h-4 w-4 rounded border-gray-300"
/>
</label>
<label className="flex items-center justify-between rounded border px-3 py-2 text-xs">
<span>加载额外目录记忆</span>
<input
type="checkbox"
checked={resolveConfig.load_additional_dirs_memory ?? false}
onChange={(event) =>
setDraft((prev) => ({
...prev,
resolve: {
...normalizeResolve(prev.resolve),
load_additional_dirs_memory: event.target.checked,
},
}))
}
className="h-4 w-4 rounded border-gray-300"
/>
</label>
</div>
</div>
<div className="rounded-lg border p-4 space-y-4">
<div className="flex items-center justify-between gap-2">
<h3 className="text-sm font-medium">自动记忆(Auto Memory)</h3>
<button
type="button"
onClick={handleToggleAutoImmediately}
className="rounded border px-2 py-1 text-[11px] hover:bg-muted"
>
{autoConfig.enabled ? "立即关闭" : "立即开启"}
</button>
</div>
<div className="grid grid-cols-1 md:grid-cols-3 gap-3">
<label className="space-y-1">
<span className="text-xs text-muted-foreground">入口文件</span>
<input
type="text"
value={autoConfig.entrypoint || "MEMORY.md"}
onChange={(event) =>
setDraft((prev) => ({
...prev,
auto: {
...normalizeAuto(prev.auto),
entrypoint: event.target.value,
},
}))
}
className="w-full rounded border bg-background px-2 py-1.5 text-xs"
/>
</label>
<label className="space-y-1">
<span className="text-xs text-muted-foreground">加载行数上限</span>
<input
type="number"
min={20}
max={1000}
value={autoConfig.max_loaded_lines ?? 200}
onChange={(event) => {
const value = Number(event.target.value);
setDraft((prev) => ({
...prev,
auto: {
...normalizeAuto(prev.auto),
max_loaded_lines: Number.isFinite(value)
? Math.max(20, Math.min(1000, value))
: 200,
},
}));
}}
className="w-full rounded border bg-background px-2 py-1.5 text-xs"
/>
</label>
<label className="space-y-1">
<span className="text-xs text-muted-foreground">自动记忆根目录</span>
<input
type="text"
value={autoConfig.root_dir || ""}
onChange={(event) =>
setDraft((prev) => ({
...prev,
auto: {
...normalizeAuto(prev.auto),
root_dir: event.target.value || undefined,
},
}))
}
className="w-full rounded border bg-background px-2 py-1.5 text-xs"
placeholder="默认自动推导,可留空"
/>
</label>
</div>
<div className="rounded border p-3 space-y-2">
<div className="text-xs text-muted-foreground">写入自动记忆</div>
<input
type="text"
value={autoTopic}
onChange={(event) => setAutoTopic(event.target.value)}
className="w-full rounded border bg-background px-2 py-1.5 text-xs"
placeholder="可选:topic,例如 workflow"
/>
<textarea
value={autoNote}
onChange={(event) => setAutoNote(event.target.value)}
className="w-full rounded border bg-background px-2 py-1.5 text-xs min-h-20"
placeholder="输入要写入自动记忆的内容"
/>
<button
type="button"
onClick={handleUpdateAutoNote}
disabled={savingAutoNote}
className="rounded bg-primary px-3 py-1.5 text-xs text-primary-foreground hover:opacity-90 disabled:opacity-60"
>
{savingAutoNote ? "写入中..." : "写入自动记忆"}
</button>
</div>
<div className="rounded border p-3 space-y-2">
<div className="text-xs text-muted-foreground">
当前索引:{autoIndex?.entry_exists ? "已存在" : "未初始化"}
{autoIndex ? `,${autoIndex.total_lines} 行` : ""}
</div>
{autoIndex?.preview_lines?.length ? (
<pre className="text-[11px] leading-relaxed whitespace-pre-wrap break-words bg-muted/30 rounded p-2 max-h-44 overflow-auto">
{autoIndex.preview_lines.join("\n")}
</pre>
) : (
<p className="text-xs text-muted-foreground">暂无自动记忆入口内容</p>
)}
</div>
</div>
<div className="rounded-lg border p-4 space-y-3 bg-muted/20">
<div className="flex items-center justify-between gap-2">
<h3 className="text-sm font-medium">记忆来源命中详情</h3>
<div className="text-xs text-muted-foreground">
{effectiveSources
? `命中 ${effectiveSources.loaded_sources}/${effectiveSources.total_sources}`
: "--"}
</div>
</div>
{effectiveSources ? (
<div className="space-y-2">
{effectiveSources.sources.map((source) => (
<div key={`${source.kind}-${source.path}`} className="rounded border bg-background/60 px-3 py-2">
<div className="flex items-center justify-between gap-2">
<span className="text-xs font-medium">{source.kind}</span>
<span
className={cn(
"rounded-full border px-2 py-0.5 text-[10px]",
source.loaded
? "text-green-700 border-green-200 bg-green-50"
: "text-muted-foreground border-muted",
)}
>
{source.loaded ? "已加载" : source.exists ? "存在未命中" : "未发现"}
</span>
</div>
<p className="mt-1 text-[11px] text-muted-foreground break-all">{source.path}</p>
{source.preview && (
<p className="mt-1 text-[11px] leading-relaxed text-muted-foreground line-clamp-2">
{source.preview}
</p>
)}
{source.warnings?.length > 0 && (
<p className="mt-1 text-[11px] text-amber-600">
{source.warnings.join(";")}
</p>
)}
</div>
))}
</div>
) : (
<p className="text-xs text-muted-foreground">正在加载来源命中结果...</p>
)}
</div>
{message && (
<div className="rounded-md border bg-muted/30 px-3 py-2 text-xs text-muted-foreground">
{message}
</div>
)}
</div>
);
}
export default MemorySettings;
@@ -25,7 +25,7 @@ import {
ShieldCheck,
HeartPulse,
Activity,
Wrench,
FlaskConical,
Code,
Info,
@@ -101,6 +101,11 @@ export function useSettingsCategory(): CategoryGroup[] {
label: t("settings.tab.hotkeys", "快捷键"),
icon: Keyboard,
},
{
key: SettingsTabs.Memory,
label: t("settings.tab.memory", "记忆"),
icon: Brain,
},
],
});
@@ -177,11 +182,6 @@ export function useSettingsCategory(): CategoryGroup[] {
label: t("settings.tab.executionTracker", "执行轨迹"),
icon: Activity,
},
{
key: SettingsTabs.ExternalTools,
label: t("settings.tab.externalTools", "外部工具"),
icon: Wrench,
},
{
key: SettingsTabs.Experimental,
label: t("settings.tab.experimental", "实验功能"),
@@ -1,3 +1,4 @@
/* global process */
import { readdirSync, readFileSync, statSync } from "node:fs";
import { join, relative } from "node:path";
import { describe, expect, it } from "vitest";
@@ -6,7 +6,7 @@
import React, { useState, useEffect, useCallback } from "react";
import { useTranslation } from "react-i18next";
import { Plus, Pencil, Trash2, Power } from "lucide-react";
import { Plus, Pencil, Trash2 } from "lucide-react";
import { toast } from "sonner";
import { Button } from "@/components/ui/button";
import { Switch } from "@/components/ui/switch";
@@ -57,14 +57,10 @@ export function AIChannelsList() {
// 添加渠道
const handleAdd = useCallback(
async (config: AIChannelConfig) => {
try {
await aiChannelsApi.createChannel(config);
toast.success("AI 渠道已创建");
await loadChannels();
setShowAddModal(false);
} catch (e) {
throw e;
}
await aiChannelsApi.createChannel(config);
toast.success("AI 渠道已创建");
await loadChannels();
setShowAddModal(false);
},
[loadChannels],
);
@@ -72,14 +68,10 @@ export function AIChannelsList() {
// 更新渠道
const handleUpdate = useCallback(
async (id: string, config: AIChannelConfig) => {
try {
await aiChannelsApi.updateChannel(id, config);
toast.success("AI 渠道已更新");
await loadChannels();
setEditingChannel(null);
} catch (e) {
throw e;
}
await aiChannelsApi.updateChannel(id, config);
toast.success("AI 渠道已更新");
await loadChannels();
setEditingChannel(null);
},
[loadChannels],
);
@@ -18,7 +18,7 @@ export interface ConnectionTestButtonProps {
export function ConnectionTestButton({
channelId,
channelName,
channelName: _channelName,
}: ConnectionTestButtonProps) {
const { t } = useTranslation();
const [testing, setTesting] = useState(false);
@@ -17,7 +17,6 @@ import {
SelectTrigger,
SelectValue,
} from "@/components/ui/select";
import { Switch } from "@/components/ui/switch";
import {
type NotificationChannel,
type NotificationChannelConfig,
@@ -6,7 +6,7 @@
import React, { useState, useEffect, useCallback } from "react";
import { useTranslation } from "react-i18next";
import { Plus, Pencil, Trash2, Send } from "lucide-react";
import { Plus, Pencil, Trash2 } from "lucide-react";
import { toast } from "sonner";
import { Button } from "@/components/ui/button";
import { Switch } from "@/components/ui/switch";
@@ -57,14 +57,10 @@ export function NotificationChannelsList() {
// 添加渠道
const handleAdd = useCallback(
async (config: NotificationChannelConfig) => {
try {
await notificationChannelsApi.createChannel(config);
toast.success("通知渠道已创建");
await loadChannels();
setShowAddModal(false);
} catch (e) {
throw e;
}
await notificationChannelsApi.createChannel(config);
toast.success("通知渠道已创建");
await loadChannels();
setShowAddModal(false);
},
[loadChannels],
);
@@ -72,14 +68,10 @@ export function NotificationChannelsList() {
// 更新渠道
const handleUpdate = useCallback(
async (id: string, config: NotificationChannelConfig) => {
try {
await notificationChannelsApi.updateChannel(id, config);
toast.success("通知渠道已更新");
await loadChannels();
setEditingChannel(null);
} catch (e) {
throw e;
}
await notificationChannelsApi.updateChannel(id, config);
toast.success("通知渠道已更新");
await loadChannels();
setEditingChannel(null);
},
[loadChannels],
);
@@ -31,7 +31,7 @@ const DEFAULT_TEST_MESSAGES = {
export function SendTestMessageButton({
channelId,
channelName,
channelName: _channelName,
channelType,
}: SendTestMessageButtonProps) {
const { t } = useTranslation();
@@ -1,240 +0,0 @@
/**
* 外部工具设置组件
*
* 管理 Codex CLI 等外部命令行工具的状态和配置
* 这些工具有自己的认证系统,不通过 ProxyCast 凭证池管理
*
* @module components/settings-v2/system/external-tools
*/
import { useState, useEffect, useCallback } from "react";
import { Button } from "@/components/ui/button";
import {
Card,
CardContent,
CardDescription,
CardHeader,
CardTitle,
} from "@/components/ui/card";
import {
RefreshCw,
CheckCircle,
XCircle,
ExternalLink,
Terminal,
Copy,
AlertCircle,
} from "lucide-react";
import {
checkCodexCliStatus,
getCodexLoginCommand,
getCodexLogoutCommand,
type CodexCliStatus,
} from "@/lib/api/externalTools";
import { toast } from "sonner";
export function ExternalToolsSettings() {
const [codexStatus, setCodexStatus] = useState<CodexCliStatus | null>(null);
const [loading, setLoading] = useState(true);
// 加载 Codex CLI 状态
const loadStatus = useCallback(async () => {
setLoading(true);
try {
const status = await checkCodexCliStatus();
setCodexStatus(status);
} catch (err) {
console.error("[ExternalTools] 加载状态失败:", err);
setCodexStatus({
installed: false,
logged_in: false,
error: String(err),
});
} finally {
setLoading(false);
}
}, []);
useEffect(() => {
loadStatus();
}, [loadStatus]);
// 复制命令到剪贴板
const copyCommand = async (command: string) => {
await navigator.clipboard.writeText(command);
toast.success("命令已复制到剪贴板,请在终端中执行");
};
// 处理登录
const handleLogin = async () => {
const cmd = await getCodexLoginCommand();
await copyCommand(cmd);
};
// 处理登出
const handleLogout = async () => {
const cmd = await getCodexLogoutCommand();
await copyCommand(cmd);
};
return (
<div className="space-y-6 max-w-2xl">
{/* Codex CLI */}
<Card>
<CardHeader>
<CardTitle className="flex items-center justify-between">
<div className="flex items-center gap-2">
<Terminal className="w-5 h-5" />
<span>Codex CLI</span>
</div>
<Button
variant="ghost"
size="sm"
onClick={loadStatus}
disabled={loading}
>
<RefreshCw
className={`w-4 h-4 ${loading ? "animate-spin" : ""}`}
/>
</Button>
</CardTitle>
<CardDescription>
OpenAI Codex 命令行工具,用于 Agent 模式的代码生成和工具调用
</CardDescription>
</CardHeader>
<CardContent className="space-y-4">
{/* 状态显示 */}
{codexStatus && (
<div className="space-y-3">
{/* 安装状态 */}
<div className="flex items-center justify-between p-3 bg-muted/50 rounded-md">
<div className="flex items-center gap-2">
{codexStatus.installed ? (
<CheckCircle className="w-4 h-4 text-green-500" />
) : (
<XCircle className="w-4 h-4 text-red-500" />
)}
<span className="text-sm">
{codexStatus.installed ? "已安装" : "未安装"}
</span>
{codexStatus.version && (
<code className="text-xs bg-muted px-2 py-0.5 rounded">
{codexStatus.version}
</code>
)}
</div>
{!codexStatus.installed && (
<Button
variant="outline"
size="sm"
onClick={() => copyCommand("npm i -g @openai/codex")}
>
<Copy className="w-3 h-3 mr-1" />
复制安装命令
</Button>
)}
</div>
{/* 登录状态 */}
{codexStatus.installed && (
<div className="flex items-center justify-between p-3 bg-muted/50 rounded-md">
<div className="flex items-center gap-2">
{codexStatus.logged_in ? (
<CheckCircle className="w-4 h-4 text-green-500" />
) : (
<AlertCircle className="w-4 h-4 text-yellow-500" />
)}
<span className="text-sm">
{codexStatus.logged_in ? "已登录" : "未登录"}
</span>
{codexStatus.auth_type && (
<code className="text-xs bg-muted px-2 py-0.5 rounded">
{codexStatus.auth_type === "api_key"
? "API Key"
: codexStatus.auth_type === "oauth"
? "OAuth"
: codexStatus.auth_type}
</code>
)}
{codexStatus.api_key_prefix && (
<code className="text-xs bg-muted px-2 py-0.5 rounded text-muted-foreground">
{codexStatus.api_key_prefix}
</code>
)}
</div>
<div className="flex gap-2">
{codexStatus.logged_in ? (
<Button
variant="outline"
size="sm"
onClick={handleLogout}
>
登出
</Button>
) : (
<Button variant="default" size="sm" onClick={handleLogin}>
登录
</Button>
)}
</div>
</div>
)}
{/* 错误信息 */}
{codexStatus.error && (
<div className="flex items-start gap-2 p-3 text-sm text-red-500 bg-red-500/10 rounded-md">
<AlertCircle className="w-4 h-4 mt-0.5 flex-shrink-0" />
<span className="whitespace-pre-wrap">
{codexStatus.error}
</span>
</div>
)}
</div>
)}
{/* 说明 */}
<div className="p-4 bg-muted/30 rounded-md space-y-2">
<h4 className="text-sm font-medium">关于 Codex CLI</h4>
<p className="text-xs text-muted-foreground">
Codex CLI 是 OpenAI 提供的命令行工具,支持 Agent
模式进行代码生成和工具调用。 它使用自己的认证系统(通过{" "}
<code>codex login</code>), 与 ProxyCast 凭证池中的 API Key
是独立的。
</p>
<div className="flex gap-2 mt-2">
<Button
variant="ghost"
size="sm"
className="h-auto p-0 text-xs text-primary hover:underline"
onClick={() =>
window.open("https://github.com/openai/codex", "_blank")
}
>
<ExternalLink className="w-3 h-3 mr-1" />
GitHub 文档
</Button>
</div>
</div>
</CardContent>
</Card>
{/* 说明卡片 */}
<Card>
<CardHeader>
<CardTitle className="text-base">CLI 工具 vs API 凭证</CardTitle>
</CardHeader>
<CardContent className="text-sm text-muted-foreground space-y-2">
<p>
<strong>CLI 工具</strong>(如 Codex CLI)有自己的认证系统,
通过命令行登录后可以在 Agent 模式中使用。
</p>
<p>
<strong>API 凭证</strong>(在凭证池中管理)用于 ProxyCast
代理服务器,将请求转发到各个 AI 服务。
</p>
<p>两者是独立的,可以同时使用不同的账号。</p>
</CardContent>
</Card>
</div>
);
}
+150
View File
@@ -163,6 +163,64 @@ export interface ChatAppearanceConfig {
append_selected_text_to_recommendation?: boolean;
}
/**
* 记忆管理系统配置
*/
export interface MemoryProfileConfig {
/** 当前学习/工作状态(单选) */
current_status?: string;
/** 擅长领域(多选) */
strengths?: string[];
/** 偏好的解释风格(多选) */
explanation_style?: string[];
/** 遇到难题时的偏好(多选) */
challenge_preference?: string[];
}
/**
* 记忆来源配置
*/
export interface MemorySourcesConfig {
/** 组织级策略文件路径 */
managed_policy_path?: string | null;
/** 项目记忆文件(按目录层级向上发现) */
project_memory_paths?: string[];
/** 项目规则目录(按目录层级向上发现) */
project_rule_dirs?: string[];
/** 用户级记忆文件路径 */
user_memory_path?: string | null;
/** 项目本地记忆文件路径 */
project_local_memory_path?: string | null;
}
/**
* 自动记忆配置
*/
export interface MemoryAutoConfig {
/** 是否启用自动记忆 */
enabled?: boolean;
/** 入口文件名 */
entrypoint?: string;
/** 启动时加载入口的最大行数 */
max_loaded_lines?: number;
/** 自动记忆根目录 */
root_dir?: string | null;
}
/**
* 记忆解析行为配置
*/
export interface MemoryResolveConfig {
/** 额外目录(例如外部 workspace) */
additional_dirs?: string[];
/** 是否跟随 @import */
follow_imports?: boolean;
/** 最大导入深度 */
import_max_depth?: number;
/** 是否加载额外目录中的记忆来源 */
load_additional_dirs_memory?: boolean;
}
/**
* 记忆管理系统配置
*/
@@ -175,6 +233,14 @@ export interface MemoryConfig {
retention_days?: number;
/** 自动清理过期记忆 */
auto_cleanup?: boolean;
/** 记忆偏好画像 */
profile?: MemoryProfileConfig;
/** 记忆来源配置 */
sources?: MemorySourcesConfig;
/** 自动记忆配置 */
auto?: MemoryAutoConfig;
/** 记忆解析行为配置 */
resolve?: MemoryResolveConfig;
}
/**
@@ -861,6 +927,48 @@ export interface MemoryAnalysisResult {
deduplicated_entries: number;
}
export interface EffectiveMemorySource {
kind: string;
path: string;
exists: boolean;
loaded: boolean;
line_count: number;
import_count: number;
warnings: string[];
preview?: string | null;
}
export interface EffectiveMemorySourcesResponse {
working_dir: string;
total_sources: number;
loaded_sources: number;
follow_imports: boolean;
import_max_depth: number;
sources: EffectiveMemorySource[];
}
export interface AutoMemoryIndexItem {
title: string;
relative_path: string;
exists: boolean;
summary?: string | null;
}
export interface AutoMemoryIndexResponse {
enabled: boolean;
root_dir: string;
entrypoint: string;
max_loaded_lines: number;
entry_exists: boolean;
total_lines: number;
preview_lines: string[];
items: AutoMemoryIndexItem[];
}
export interface MemoryAutoToggleResponse {
enabled: boolean;
}
/**
* 获取记忆统计信息
*/
@@ -897,6 +1005,48 @@ export async function cleanupMemory(): Promise<CleanupMemoryResult> {
return safeInvoke("cleanup_conversation_memory");
}
/**
* 获取记忆来源解析结果
*/
export async function getMemoryEffectiveSources(
workingDir?: string,
activeRelativePath?: string,
): Promise<EffectiveMemorySourcesResponse> {
return safeInvoke("memory_get_effective_sources", {
workingDir,
activeRelativePath,
});
}
/**
* 获取自动记忆入口索引
*/
export async function getMemoryAutoIndex(
workingDir?: string,
): Promise<AutoMemoryIndexResponse> {
return safeInvoke("memory_get_auto_index", { workingDir });
}
/**
* 切换自动记忆开关
*/
export async function toggleMemoryAuto(
enabled: boolean,
): Promise<MemoryAutoToggleResponse> {
return safeInvoke("memory_toggle_auto", { enabled });
}
/**
* 写入自动记忆笔记
*/
export async function updateMemoryAutoNote(
note: string,
topic?: string,
workingDir?: string,
): Promise<AutoMemoryIndexResponse> {
return safeInvoke("memory_update_auto_note", { note, topic, workingDir });
}
// ============ 语音测试 API ============
export interface TtsTestResult {
+18
View File
@@ -57,6 +57,7 @@ export type StreamEvent =
| StreamEventToolStart
| StreamEventToolEnd
| StreamEventActionRequired
| StreamEventContextTrace
| StreamEventDone
| StreamEventFinalDone
| StreamEventWarning
@@ -136,6 +137,16 @@ export interface StreamEventActionRequired {
requested_schema?: Record<string, unknown>;
}
export interface ContextTraceStep {
stage: string;
detail: string;
}
export interface StreamEventContextTrace {
type: "context_trace";
steps: ContextTraceStep[];
}
/**
* 完成事件(单次 API 响应完成,工具循环可能继续)
* Requirements: 9.5 - THE Frontend SHALL display token usage statistics after each Agent response
@@ -298,6 +309,13 @@ export function parseStreamEvent(data: unknown): StreamEvent | null {
type: "done",
usage: event.usage as TokenUsage | undefined,
};
case "context_trace":
return {
type: "context_trace",
steps: Array.isArray(event.steps)
? (event.steps as ContextTraceStep[])
: [],
};
case "final_done":
return {
type: "final_done",
+14 -10
View File
@@ -84,7 +84,7 @@ export interface UnifiedMemory {
/** 更新时间(毫秒时间戳) */
updated_at: number;
/** 是否已归档(软删除) */
/** 是否已归档 */
archived: boolean;
}
@@ -367,10 +367,12 @@ export async function semanticSearch(
console.log('[语义搜索] Query:', query, 'Category:', category, 'MinSimilarity:', minSimilarity);
const result = await invoke<UnifiedMemory[]>("unified_memory_semantic_search", {
query,
category: category?.toString(),
min_similarity: minSimilarity,
limit,
options: {
query,
category,
min_similarity: minSimilarity,
limit,
},
});
console.log('[语义搜索] Results:', result);
@@ -397,11 +399,13 @@ export async function hybridSearch(
console.log('[混合搜索] Query:', query, 'Category:', category, 'SemanticWeight:', semanticWeight, 'MinSimilarity:', minSimilarity);
const result = await invoke<UnifiedMemory[]>("unified_memory_hybrid_search", {
query,
category: category?.toString(),
semantic_weight: semanticWeight,
min_similarity: minSimilarity,
limit,
options: {
query,
category,
semantic_weight: semanticWeight,
min_similarity: minSimilarity,
limit,
},
});
console.log('[混合搜索] Results:', result);
+4 -2
View File
@@ -26,6 +26,7 @@ export enum SettingsTabs {
Appearance = "appearance",
ChatAppearance = "chat-appearance",
Hotkeys = "hotkeys",
Memory = "memory",
// 智能体
Providers = "providers",
@@ -42,7 +43,7 @@ export enum SettingsTabs {
SecurityPerformance = "security-performance",
Heartbeat = "heartbeat",
ExecutionTracker = "execution-tracker",
ExternalTools = "external-tools",
Experimental = "experimental",
Developer = "developer",
About = "about",
@@ -75,6 +76,7 @@ export const SETTINGS_GROUPS: Record<SettingsGroupKey, SettingsTabs[]> = {
SettingsTabs.Appearance,
SettingsTabs.ChatAppearance,
SettingsTabs.Hotkeys,
SettingsTabs.Memory,
],
[SettingsGroupKey.Agent]: [
SettingsTabs.Providers,
@@ -91,7 +93,7 @@ export const SETTINGS_GROUPS: Record<SettingsGroupKey, SettingsTabs[]> = {
SettingsTabs.SecurityPerformance,
SettingsTabs.Heartbeat,
SettingsTabs.ExecutionTracker,
SettingsTabs.ExternalTools,
SettingsTabs.Experimental,
SettingsTabs.Developer,
SettingsTabs.About,