mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
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:
+29
-35
@@ -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
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"private": true,
|
||||
"version": "0.72.0",
|
||||
"version": "0.73.0",
|
||||
"type": "module",
|
||||
"repository": {
|
||||
"type": "git",
|
||||
|
||||
Generated
+16
-15
@@ -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",
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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!({
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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};
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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(®istry), vec!["auth", "routing", "provider"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_first_and_last() {
|
||||
let mut registry = StepRegistry::new(vec![DummyStep::boxed("core")]);
|
||||
registry.register(DummyStep::boxed("first_step"), PipelinePhase::First, 0);
|
||||
registry.register(DummyStep::boxed("last_step"), PipelinePhase::Last, 0);
|
||||
assert_eq!(
|
||||
step_names(®istry),
|
||||
vec!["first_step", "core", "last_step"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_before_and_after() {
|
||||
let mut registry =
|
||||
StepRegistry::new(vec![DummyStep::boxed("auth"), DummyStep::boxed("provider")]);
|
||||
registry.register(
|
||||
DummyStep::boxed("pre_auth"),
|
||||
PipelinePhase::Before("auth".to_string()),
|
||||
0,
|
||||
);
|
||||
registry.register(
|
||||
DummyStep::boxed("post_auth"),
|
||||
PipelinePhase::After("auth".to_string()),
|
||||
0,
|
||||
);
|
||||
assert_eq!(
|
||||
step_names(®istry),
|
||||
vec!["pre_auth", "auth", "post_auth", "provider"]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_replace() {
|
||||
let mut registry =
|
||||
StepRegistry::new(vec![DummyStep::boxed("auth"), DummyStep::boxed("provider")]);
|
||||
registry.register(
|
||||
DummyStep::boxed("custom_auth"),
|
||||
PipelinePhase::Replace("auth".to_string()),
|
||||
0,
|
||||
);
|
||||
assert_eq!(step_names(®istry), vec!["custom_auth", "provider"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_unregister() {
|
||||
let mut registry = StepRegistry::new(vec![DummyStep::boxed("core")]);
|
||||
registry.register(DummyStep::boxed("extra"), PipelinePhase::Last, 0);
|
||||
assert_eq!(step_names(®istry), vec!["core", "extra"]);
|
||||
|
||||
assert!(registry.unregister("extra"));
|
||||
assert_eq!(step_names(®istry), vec!["core"]);
|
||||
|
||||
// 移除不存在的步骤返回 false
|
||||
assert!(!registry.unregister("nonexistent"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_priority_ordering() {
|
||||
let mut registry = StepRegistry::new(vec![DummyStep::boxed("core")]);
|
||||
registry.register(DummyStep::boxed("low"), PipelinePhase::First, 10);
|
||||
registry.register(DummyStep::boxed("high"), PipelinePhase::First, 1);
|
||||
// priority 小的排前面
|
||||
assert_eq!(step_names(®istry), vec!["high", "low", "core"]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_complex_pipeline() {
|
||||
let mut registry = StepRegistry::new(vec![
|
||||
DummyStep::boxed("auth"),
|
||||
DummyStep::boxed("routing"),
|
||||
DummyStep::boxed("provider"),
|
||||
]);
|
||||
registry.register(DummyStep::boxed("init"), PipelinePhase::First, 0);
|
||||
registry.register(
|
||||
DummyStep::boxed("rate_limit"),
|
||||
PipelinePhase::Before("auth".to_string()),
|
||||
0,
|
||||
);
|
||||
registry.register(
|
||||
DummyStep::boxed("log_auth"),
|
||||
PipelinePhase::After("auth".to_string()),
|
||||
0,
|
||||
);
|
||||
registry.register(
|
||||
DummyStep::boxed("custom_routing"),
|
||||
PipelinePhase::Replace("routing".to_string()),
|
||||
0,
|
||||
);
|
||||
registry.register(DummyStep::boxed("telemetry"), PipelinePhase::Last, 0);
|
||||
|
||||
assert_eq!(
|
||||
step_names(®istry),
|
||||
vec![
|
||||
"init",
|
||||
"rate_limit",
|
||||
"auth",
|
||||
"log_auth",
|
||||
"custom_routing",
|
||||
"provider",
|
||||
"telemetry"
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
§ion_lines,
|
||||
) {
|
||||
entries.push(entry);
|
||||
index += 1;
|
||||
}
|
||||
}
|
||||
|
||||
current_title = Some(title.trim().to_string());
|
||||
section_lines.clear();
|
||||
continue;
|
||||
}
|
||||
|
||||
if current_title.is_some() {
|
||||
section_lines.push(line.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(previous_title) = current_title {
|
||||
if let Some(entry) = self.build_memory_entry(
|
||||
session_id,
|
||||
file_type,
|
||||
index,
|
||||
&previous_title,
|
||||
§ion_lines,
|
||||
) {
|
||||
entries.push(entry);
|
||||
}
|
||||
}
|
||||
|
||||
entries
|
||||
}
|
||||
|
||||
fn build_memory_entry(
|
||||
&self,
|
||||
session_id: &str,
|
||||
file_type: MemoryFileType,
|
||||
index: usize,
|
||||
title: &str,
|
||||
lines: &[String],
|
||||
) -> Option<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("应被归档的条目"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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("不要自动修复"));
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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}"))?;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
¬e,
|
||||
topic.as_deref(),
|
||||
)
|
||||
}
|
||||
|
||||
fn resolve_memory_dir() -> PathBuf {
|
||||
dirs::home_dir()
|
||||
.map(|p| p.join(".proxycast").join("memory"))
|
||||
.unwrap_or_else(|| PathBuf::from(".proxycast/memory"))
|
||||
}
|
||||
|
||||
fn resolve_working_dir(working_dir: Option<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 {
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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();
|
||||
|
||||
// 用户消息 = 原始输入 + 前序步骤的累积上下文
|
||||
|
||||
@@ -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 任务
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(¤t);
|
||||
let home_dir = dirs::home_dir();
|
||||
let mut depth = 0usize;
|
||||
|
||||
loop {
|
||||
dirs.push(current.clone());
|
||||
if let Some(root) = project_root.as_ref() {
|
||||
if ¤t == root {
|
||||
break;
|
||||
}
|
||||
}
|
||||
if let Some(home) = home_dir.as_ref() {
|
||||
if ¤t == home {
|
||||
break;
|
||||
}
|
||||
}
|
||||
// 兜底保护,避免跨层级扫描过深导致来源列表爆炸
|
||||
if depth >= 12 {
|
||||
break;
|
||||
}
|
||||
if !current.pop() {
|
||||
break;
|
||||
}
|
||||
depth += 1;
|
||||
}
|
||||
|
||||
dirs
|
||||
}
|
||||
|
||||
fn expand_path(path: &str, working_dir: Option<&Path>) -> PathBuf {
|
||||
let trimmed = path.trim();
|
||||
if trimmed.starts_with("~/") {
|
||||
if let Some(home) = dirs::home_dir() {
|
||||
return home.join(trimmed.trim_start_matches("~/"));
|
||||
}
|
||||
}
|
||||
|
||||
let p = PathBuf::from(trimmed);
|
||||
if p.is_absolute() {
|
||||
return p;
|
||||
}
|
||||
|
||||
if let Some(base) = working_dir {
|
||||
return base.join(p);
|
||||
}
|
||||
p
|
||||
}
|
||||
|
||||
fn default_user_memory_path() -> PathBuf {
|
||||
dirs::home_dir()
|
||||
.unwrap_or_else(|| PathBuf::from("."))
|
||||
.join(".proxycast")
|
||||
.join("AGENTS.md")
|
||||
}
|
||||
|
||||
fn default_managed_policy_path() -> PathBuf {
|
||||
#[cfg(target_os = "macos")]
|
||||
{
|
||||
return PathBuf::from("/Library/Application Support/ProxyCast/AGENTS.md");
|
||||
}
|
||||
#[cfg(target_os = "linux")]
|
||||
{
|
||||
return PathBuf::from("/etc/proxycast/AGENTS.md");
|
||||
}
|
||||
#[cfg(target_os = "windows")]
|
||||
{
|
||||
return PathBuf::from("C:/Program Files/ProxyCast/AGENTS.md");
|
||||
}
|
||||
#[allow(unreachable_code)]
|
||||
PathBuf::from("/etc/proxycast/AGENTS.md")
|
||||
}
|
||||
|
||||
fn normalize_path(path: &Path) -> PathBuf {
|
||||
path.canonicalize().unwrap_or_else(|_| path.to_path_buf())
|
||||
}
|
||||
|
||||
fn find_git_root(start: &Path) -> Option<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);
|
||||
}
|
||||
}
|
||||
@@ -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,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) =>
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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" />
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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,
|
||||
};
|
||||
}
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user