feat: enhance memory management tools; improve lesson tracking and update descriptions for clarity
This commit is contained in:
+60
-12
@@ -2,13 +2,9 @@
|
||||
//! before it executes.
|
||||
|
||||
/// Outcome of gating a tool call: whether it's allowed to run.
|
||||
///
|
||||
/// Why: `Block` carries a reason string for surfacing to the user/log, even
|
||||
/// though nothing currently produces `Block` (classify() always allows).
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum Verdict {
|
||||
Allow,
|
||||
#[allow(dead_code)]
|
||||
Block(String),
|
||||
}
|
||||
|
||||
@@ -19,23 +15,20 @@ impl Harness {
|
||||
/// Decide whether a tool call is allowed to execute.
|
||||
///
|
||||
/// Flow: if the tool isn't flagged risky, allow immediately → basic
|
||||
/// content checks (path traversal) → defer to `classify`.
|
||||
///
|
||||
/// Why: `classify` is currently a stub that always allows; the basic
|
||||
/// checks here serve as defense-in-depth alongside the shell filters
|
||||
/// and `resolve_path` in the tool modules.
|
||||
/// content checks (path traversal in paths AND command args) →
|
||||
/// workspace-root validation for output paths → classify.
|
||||
///
|
||||
/// Return: `Verdict::Allow` or `Verdict::Block(reason)`.
|
||||
pub fn gate_tool_call(
|
||||
tool_name: &str,
|
||||
args: &serde_json::Value,
|
||||
_workspace_roots: &[&std::path::Path],
|
||||
workspace_roots: &[&std::path::Path],
|
||||
) -> Verdict {
|
||||
|
||||
if !crate::tool::tool_is_risky(tool_name) {
|
||||
return Verdict::Allow;
|
||||
}
|
||||
// Basic path traversal check for file-mutating tools.
|
||||
// Path traversal check for file-mutating tools.
|
||||
if matches!(tool_name, "write" | "edit" | "delete") {
|
||||
if let Some(path) = args.get("path").and_then(|v| v.as_str()) {
|
||||
if path.contains("..") {
|
||||
@@ -43,11 +36,66 @@ impl Harness {
|
||||
}
|
||||
}
|
||||
}
|
||||
// Path traversal and dangerous content check for bash commands.
|
||||
if tool_name == "bash" {
|
||||
let cmd = args.get("command").and_then(|v| v.as_str()).unwrap_or("");
|
||||
if cmd.contains("..") {
|
||||
return Verdict::Block("path traversal detected in bash command".to_string());
|
||||
}
|
||||
let dangerous_patterns = [
|
||||
"rm -rf /", "rm -rf --no-preserve-root",
|
||||
"rm -rf ~", "rm -fr /", "mkfs.", "dd if=",
|
||||
":(){", "> /dev/sda", "chmod -R 000 /",
|
||||
];
|
||||
for pat in &dangerous_patterns {
|
||||
if cmd.contains(pat) {
|
||||
return Verdict::Block(format!("destructive command pattern blocked: {}", pat));
|
||||
}
|
||||
}
|
||||
}
|
||||
if let Some(out_path) = Self::find_output_path(tool_name, args) {
|
||||
if !workspace_roots.is_empty()
|
||||
&& !out_path.starts_with("/tmp")
|
||||
&& !out_path.is_absolute()
|
||||
{
|
||||
let allowed = workspace_roots.iter().any(|r| out_path.starts_with(r));
|
||||
if !allowed {
|
||||
return Verdict::Block(format!(
|
||||
"output path '{:?}' is outside all workspace roots", out_path
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
Self::classify(tool_name)
|
||||
}
|
||||
|
||||
/// Extract a candidate output path from a tool call, if one exists.
|
||||
///
|
||||
/// Used to verify that writes and file mutations stay inside workspace roots.
|
||||
fn find_output_path(tool_name: &str, args: &serde_json::Value) -> Option<std::path::PathBuf> {
|
||||
match tool_name {
|
||||
"write" | "edit" | "delete" | "read" => {
|
||||
args.get("path").and_then(|v| v.as_str()).map(std::path::PathBuf::from)
|
||||
}
|
||||
"bash" => {
|
||||
let cmd = args.get("command").and_then(|v| v.as_str())?;
|
||||
let lower = cmd.to_lowercase();
|
||||
for prefix in &["cp ", "mv ", "install ", "ln -s ", "cat >", "cat >>"] {
|
||||
if let Some(rest) = lower.strip_prefix(prefix) {
|
||||
if let Some(target) = rest.split_whitespace().last() {
|
||||
if !target.starts_with('-') {
|
||||
return Some(std::path::PathBuf::from(target));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn classify(cmd: &str) -> Verdict {
|
||||
// Classify known-dangerous patterns beyond path traversal.
|
||||
match cmd {
|
||||
"bash" | "write" | "edit" | "delete" | "git_operator" => Verdict::Allow,
|
||||
_ => Verdict::Allow,
|
||||
|
||||
@@ -130,6 +130,9 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) {
|
||||
state.push_transcript(ChatMessageDisplay::new(Role::User, text.clone()));
|
||||
if let Some(ref mut rt) = state.session_runtime {
|
||||
rt.push_message(ChatMessage::user(text));
|
||||
refresh_lesson_counters(&state.memory_dir, rt);
|
||||
} else {
|
||||
let _ = std::fs::create_dir_all(&state.memory_dir);
|
||||
}
|
||||
state.misc.thinking = true;
|
||||
spawn_turn(state);
|
||||
@@ -372,19 +375,15 @@ pub fn apply_action(state: &mut AppStateRest, action: Action) {
|
||||
let _ = trigger_review(state);
|
||||
}
|
||||
} else if kind == "review" {
|
||||
let lessons_found = if message.contains("lesson") || message.contains("Lesson") {
|
||||
message.rsplit(' ').next().and_then(|w| {
|
||||
w.trim_end_matches(')').trim_end_matches('s')
|
||||
.split('(').next_back()
|
||||
.and_then(|n| n.parse::<u32>().ok())
|
||||
}).unwrap_or(0)
|
||||
let counted = if let Some(ref mut rt) = state.session_runtime {
|
||||
refresh_lesson_counters(&state.memory_dir, rt);
|
||||
true
|
||||
} else {
|
||||
0
|
||||
false
|
||||
};
|
||||
if let Some(ref mut rt) = state.session_runtime {
|
||||
if lessons_found > 0 {
|
||||
if counted {
|
||||
rt.consecutive_empty_reviews = 0;
|
||||
rt.lesson_count += lessons_found;
|
||||
} else {
|
||||
rt.consecutive_empty_reviews += 1;
|
||||
}
|
||||
@@ -634,6 +633,85 @@ fn generate_workspace_tree(roots: &[std::path::PathBuf]) -> String {
|
||||
out
|
||||
}
|
||||
|
||||
/// Load all memory entries from `memory_dir` and format them as a compact
|
||||
/// section appended to the system prompt, so the AI is always aware of
|
||||
/// stored lessons and project knowledge.
|
||||
///
|
||||
/// Flow: list memory slugs → for each, read + parse the file → collect
|
||||
/// entries whose lifecycle is not "stale" → cap total output at 3000 chars
|
||||
/// to avoid dominating the prompt budget.
|
||||
///
|
||||
/// Why: previously, lessons existed on disk but the AI never saw them
|
||||
/// unless it explicitly called `recall()`. This makes the memory system
|
||||
/// actually useful by surfacing relevant knowledge automatically.
|
||||
///
|
||||
/// Return: a formatted string (may be empty if no memory entries exist).
|
||||
fn build_memory_section(memory_dir: &std::path::Path) -> String {
|
||||
let names = crate::model::memory::Memory::list(memory_dir);
|
||||
if names.is_empty() {
|
||||
return String::new();
|
||||
}
|
||||
|
||||
let mut section = String::from("\n\n--- Persistent Memory ---\n");
|
||||
section.push_str(&format!("Total entries: {}\n\n", names.len()));
|
||||
|
||||
for name in &names {
|
||||
if section.len() > 3000 {
|
||||
section.push_str("... (more entries omitted, use recall() to see all)\n");
|
||||
break;
|
||||
}
|
||||
if let Ok(mem) = crate::model::memory::Memory::read(memory_dir, name) {
|
||||
if mem.lifecycle == "stale" {
|
||||
continue;
|
||||
}
|
||||
section.push_str(&format!("## [{}] {}\n{}\n\n", mem.kind, mem.name, mem.content));
|
||||
}
|
||||
}
|
||||
section.push_str("---");
|
||||
section
|
||||
}
|
||||
|
||||
/// Scan `memory_dir` and update every lesson counter in `SessionRuntime`
|
||||
/// from real on-disk data.
|
||||
///
|
||||
/// Flow: list all memory slugs → read+parse each → increment the matching
|
||||
/// kind counter (user/feedback/project/reference), lifecycle counter
|
||||
/// (active/stale/contradicted), and the total. If a memory cannot be read
|
||||
/// (e.g. a race with deletion) it is silently skipped.
|
||||
///
|
||||
/// Why: previously the UI showed all zeros because nothing ever set the
|
||||
/// breakdown counters. This runs on every user submit so the dashboard
|
||||
/// reflects actual memory state.
|
||||
fn refresh_lesson_counters(memory_dir: &std::path::Path, rt: &mut crate::app::state::runtime::SessionRuntime) {
|
||||
let names = crate::model::memory::Memory::list(memory_dir);
|
||||
rt.lesson_count = 0;
|
||||
rt.lessons_user = 0;
|
||||
rt.lessons_feedback = 0;
|
||||
rt.lessons_project = 0;
|
||||
rt.lessons_reference = 0;
|
||||
rt.lessons_active = 0;
|
||||
rt.lessons_stale = 0;
|
||||
rt.lessons_contradicted = 0;
|
||||
for name in &names {
|
||||
if let Ok(mem) = crate::model::memory::Memory::read(memory_dir, name) {
|
||||
rt.lesson_count += 1;
|
||||
match mem.kind.as_str() {
|
||||
"user" => rt.lessons_user += 1,
|
||||
"feedback" => rt.lessons_feedback += 1,
|
||||
"project" => rt.lessons_project += 1,
|
||||
"reference" => rt.lessons_reference += 1,
|
||||
_ => {}
|
||||
}
|
||||
match mem.lifecycle.as_str() {
|
||||
"active" => rt.lessons_active += 1,
|
||||
"stale" => rt.lessons_stale += 1,
|
||||
"contradicted" => rt.lessons_contradicted += 1,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Persist a `ChatMessage` to the SQLite message log, if a database
|
||||
/// connection is available.
|
||||
///
|
||||
@@ -679,11 +757,13 @@ fn run_agent_turn(
|
||||
let mut prev_shaped = false;
|
||||
|
||||
let tree_info = generate_workspace_tree(&tc.workspace_roots);
|
||||
let memory_section = build_memory_section(&tc.ctx.memory_dir);
|
||||
let system_text = format!(
|
||||
"{}\n\n{}\n\n{}",
|
||||
"{}\n\n{}\n\n{}{}",
|
||||
crate::resources::SYSTEM_PROMPT,
|
||||
crate::resources::SYSTEM_TOOLS,
|
||||
tree_info
|
||||
tree_info,
|
||||
memory_section,
|
||||
);
|
||||
if !msgs.iter().any(|m| matches!(m.role, crate::dto::chat::message::Role::System)) {
|
||||
let sys = ChatMessage::system(system_text);
|
||||
|
||||
Reference in New Issue
Block a user