修复: AI工具OOM/路径遍历+续跑锁收敛+UX交互批

This commit is contained in:
2026-06-17 14:07:16 +08:00
parent 2065335c8c
commit 6b67214395
15 changed files with 390 additions and 98 deletions

View File

@@ -941,21 +941,32 @@ pub(crate) async fn run_agentic_loop(
/// 分支:审批已被 remove,改为以剩余 pending_approvals 任一 conversation_id 做一致性校验
/// (此处空,校验通过即沿用全局值,该期 generating=true 且 switch 为 readonly 不并发)。
pub(crate) async fn try_continue_agent_loop(app: &AppHandle, state: &AppState, start_iteration: usize) {
let (is_generating, has_pending, pending_conv_id) = {
// BUG-260617-05: 原 5 次独立 lock().await 造成 TOCTOU 竞态——should_continue=true 判出后、
// spawn 前用户点 stop(ai_chat_stop 复位 generating=false),续跑仍按过时快照继续 spawn。
// 修复:单次 lock 取结构化快照(所有续跑判定所需字段),无锁态判定;spawn 前单次 lock 原子
// 重检 generating 仍为 true 才续跑,stop 后中途插入的直接收敛退出。
let snap = {
let session = state.ai_session.lock().await;
// pending_approvals 中任一审批的 conversation_id:审批等待态(has_pending)下作为 conv_id 来源,
// 取第一个非空值(同一对话的审批 conversation_id 一致,见 process_tool_calls 写入路径)。
let pending_conv_id = session.pending_approvals.values()
.find_map(|a| a.conversation_id.clone());
(session.generating, !session.pending_approvals.is_empty(), pending_conv_id)
ContinueSnapshot {
is_generating: session.generating,
has_pending: !session.pending_approvals.is_empty(),
pending_conv_id,
active_conversation_id: session.active_conversation_id.clone(),
agent_language: session.agent_language.clone(),
model_override: session.model_override.clone(),
}
};
let should_continue = is_generating && !has_pending;
let should_continue = snap.is_generating && !snap.has_pending;
if !should_continue {
// generating=false(被 stop)或仍有审批(pending_approvals 非空):
// 统一 emit AiCompleted 标当前轮收敛,清前端 streaming。
// 轮 token 已在前序 AiCompleted/AiApprovalResult 流程落库,此处零 token 上报仅作收敛信号。
if is_generating {
if snap.is_generating {
// pending_approvals 非空但 generating 仍 true:转审批态,前端审批态 watchdog 已 clear,不卡
tracing::info!("[ai] try_continue 跳过:仍有待审批,转审批等待态");
} else {
@@ -963,12 +974,9 @@ pub(crate) async fn try_continue_agent_loop(app: &AppHandle, state: &AppState, s
tracing::info!("[ai] try_continue 跳过:generating 已复位(被 stop/已结束),补发 AiCompleted 清前端 streaming");
// R-PD-6: 优先用审批所属 conversation_id(审批等待态被 stop 触发,审批仍在 pending_approvals),
// 仅当无任何审批(has_pending=false 且 generating=false)时回退 active_conversation_id。
let conv_id = match pending_conv_id {
let conv_id = match snap.pending_conv_id.clone() {
Some(cid) => cid,
None => {
let session = state.ai_session.lock().await;
session.active_conversation_id.clone().unwrap_or_default()
}
None => snap.active_conversation_id.clone().unwrap_or_default(),
};
let _ = app.emit("ai-chat-event", AiChatEvent::AiCompleted {
total_tokens: 0,
@@ -1009,13 +1017,9 @@ pub(crate) async fn try_continue_agent_loop(app: &AppHandle, state: &AppState, s
};
// R-PD-6: 续生成路径 conv_id 解耦——has_pending=false 时审批已 remove,无审批 conversation_id 可取;
// 此期 generating=true 且 switchConversation 为 readonly 不并发改 active_conversation_id,
// 故读全局值安全(非竞态期);若 has_pending=true 已在上面 return,不会到此。
let (lang, conv_id) = {
let session = state.ai_session.lock().await;
let lang = session.agent_language.clone().unwrap_or_else(|| "zh-CN".to_string());
let conv_id = session.active_conversation_id.clone().unwrap_or_default();
(lang, conv_id)
};
// 故读快照值安全(非竞态期);若 has_pending=true 已在上面 return,不会到此。
let lang = snap.agent_language.clone().unwrap_or_else(|| "zh-CN".to_string());
let conv_id = snap.active_conversation_id.clone().unwrap_or_default();
let system_prompt = build_system_prompt(state, &lang).await;
let session_arc = state.ai_session.clone();
@@ -1029,10 +1033,22 @@ pub(crate) async fn try_continue_agent_loop(app: &AppHandle, state: &AppState, s
// F-260616-07: 流式失败重试次数快照
let max_retries = state.agent_max_retries.load(Ordering::SeqCst);
// F-01 阶段6: 续跑沿用同一主对话的 model_override(审批续跑/达 max 续跑保持一致)。
let model_override = {
let session = state.ai_session.lock().await;
session.model_override.clone()
};
let model_override = snap.model_override.clone();
// BUG-260617-05 续: provider 解析/build_system_prompt 期间用户可能点 stop。
// spawn 前单次 lock 原子重检 generating——若已被 stop 复位,收敛退出而非覆盖用户的 stop。
// (run_agentic_loop 入口 GeneratingGuard 会再次置 generating=true,若不重检会抹掉 stop。)
if !state.ai_session.lock().await.generating {
tracing::info!("[ai] try_continue 终止:spawn 前重检 generating 已被 stop 复位,补发 AiCompleted");
let _ = app.emit("ai-chat-event", AiChatEvent::AiCompleted {
total_tokens: 0,
prompt_tokens: 0,
completion_tokens: 0,
incomplete: None,
conversation_id: Some(conv_id.clone()),
});
return;
}
// 恢复循环前通知前端新建 assistant 消息:审批(通过/拒绝)后新一轮文本
// 不应追加到发起工具调用的旧消息,用 AiAgentRound 隔开
@@ -1045,3 +1061,15 @@ pub(crate) async fn try_continue_agent_loop(app: &AppHandle, state: &AppState, s
run_agentic_loop(session_arc, tools_arc, db, app_handle, provider_config, system_prompt, conv_id, knowledge_config, llm_concurrency, max_iterations, max_retries, start_iteration, model_override).await;
});
}
/// BUG-260617-05: try_continue_agent_loop 续跑判定所需 session 字段的一次性快照。
/// 单次 lock 取出后无锁态判定,消除多 lock 间其他 IPC(ai_chat_stop/clear/switch)改写 session
/// 致续跑判断基于过时快照的 TOCTOU 竞态。
struct ContinueSnapshot {
is_generating: bool,
has_pending: bool,
pending_conv_id: Option<String>,
active_conversation_id: Option<String>,
agent_language: Option<String>,
model_override: Option<String>,
}

View File

@@ -46,8 +46,10 @@ impl TokenAccumulator {
/// 把单轮增量叠加到 DB 的 Option<i64> 字段(读旧值+增量,跨 loop 实例防覆盖)
///
/// 纯函数:抽自 save_conversation 的 token 累加逻辑,None 起始当作 0。
/// saturating_add:长期对话累积接近 i64::MAX 时不再翻负,封顶在 i64::MAX(统计语义安全,
/// 溢出回绕成负值会污染前端用量展示与计费/限额判定)。
pub(crate) fn accumulate_tokens(old: Option<i64>, add: u32) -> Option<i64> {
Some(old.unwrap_or(0) + add as i64)
Some(old.unwrap_or(0).saturating_add(add as i64))
}
/// 持久化截断阈值:超过此长度的消息 content 落库前截断头尾各保 HEAD/TAIL 字符。

View File

@@ -29,12 +29,36 @@ const DEFAULT_RUN_COMMAND_TIMEOUT_SECS: u64 = 60;
///
/// AE-2025-03路径 Baudit.rs 挂起审批前预读旧文件复用此函数生成 diff
/// 供前端审批卡即时预览write_file 审批不再只看裸 content。故 pub(crate)。
///
/// BUG-260617-07: LCS DP 表 O(n*m) 内存,两文件均>1000 行时 DP 表可达数百 MB
/// (5000×5000×8B≈200MB)。虽然 1MB 字节上限能挡住多数情况,但短行高密度的源码
/// (大量空行/单字符行)仍可能突破。改:两文件均>1000 行时跳过 LCS,退化为
/// 「删旧全量 + 增新全量」朴素行对比 + 截断提示。朴素对比 O(n+m) 内存,安全。
/// (patch_file 调用方传入的 old/new 是局部替换,通常远小于全文,正常路径仍走 LCS)
pub(crate) fn generate_diff(old: &str, new: &str) -> String {
let a: Vec<&str> = old.lines().collect();
let b: Vec<&str> = new.lines().collect();
let (n, m) = (a.len(), b.len());
// LCS 动态规划表usize 即可;大文件已被 1MB 限制挡住,行数有限)
// 超长输入降级:两文件均>1000 行跳过 LCS DP(防 200MB+ 内存峰值),
// 退化为朴素全删全增 diff(配 300 行截断),足够审批卡看出"大范围改动"语义。
const LCS_MAX_LINES: usize = 1000;
if n > LCS_MAX_LINES && m > LCS_MAX_LINES {
let mut out = String::new();
let mut changes = 0usize;
for line in &a { out.push_str("-"); out.push_str(line); out.push('\n'); changes += 1; }
for line in &b { out.push_str("+"); out.push_str(line); out.push('\n'); changes += 1; }
if changes > 300 {
let kept: String = out.lines().take(300).collect::<Vec<_>>().join("\n");
return format!(
"{}\n... (输入过长({}/{} 行)已跳过 LCS 退化对比,diff 已截断,共 {} 处变更行)",
kept, n, m, changes
);
}
return out.trim_end_matches('\n').to_string() + "\n";
}
// LCS 动态规划表usize 即可;已由 LCS_MAX_LINES 挡住超大输入,行数有限)
let mut dp = vec![vec![0usize; m + 1]; n + 1];
for i in (0..n).rev() {
for j in (0..m).rev() {
@@ -74,14 +98,35 @@ pub(crate) fn generate_diff(old: &str, new: &str) -> String {
out.trim_end_matches('\n').to_string() + "\n"
}
/// 验证文件路径:禁止访问系统敏感目录
/// 验证文件路径:禁止访问系统敏感目录 + 防路径遍历绕过
///
/// 安全要点(BUG-260617-03):
/// 1. **先 URL 解码**再检查——防 `%2e%2e` / `%2f` 等 URL 编码绕过 `..` 检测。
/// LLM / 外部输入可能传编码串,若直接做字符串 `..` 检查会被 `%2e%2e` 骗过,
/// 解码后再走分段归一化校验。
/// 2. **按路径分隔符分段**做 `..` 归一化检测——纯子串 `contains("..")` 会误伤
/// 合法目录名如 `my..file`,改用分段(逐段判断是否有 `..` 段)更精准,且能识别
/// 解码后 `..%2f` / `..\` 等变体。
/// 3. 敏感系统目录黑名单仍保留(.ssh/.aws/.gnupg/AppData/ProgramData/Windows/System32)。
fn validate_path(path: &str) -> anyhow::Result<()> {
// 规范化为反斜杠LLM 可能传正斜杠绕过黑名单Windows tokio::fs 两种分隔符都吃)
let normalized = path.replace('/', "\\");
// 1) URL 解码:防 %2e%2e / %2f / %5c 等 URL 编码绕过 (LLM/外部输入可能传编码串)
// percent_decode_str 对非法 %XX 容错(保留原字节),decode_utf8_lossy 容错非 UTF-8。
use percent_encoding::percent_decode_str;
let decoded = percent_decode_str(path).decode_utf8_lossy().into_owned();
// 2) 规范化为反斜杠:LLM 可能传正斜杠绕过黑名单(Windows tokio::fs 两种分隔符都吃)
let normalized = decoded.replace('/', "\\");
let lower = normalized.to_lowercase();
if lower.contains("..") {
// 3) 分段归一化检查:按分隔符切分,任何一段 == ".." 视为路径遍历。
// 比纯 contains("..") 更精准(不误伤 my..file 这类合法名),且能识别解码后的 `..`。
let has_traversal = lower
.split(|c| c == '\\' || c == '/')
.any(|seg| seg == "..");
if has_traversal {
anyhow::bail!("禁止路径遍历 (..)");
}
if lower.contains("\\.ssh")
|| lower.contains("\\.aws")
|| lower.contains("\\.gnupg")
@@ -791,16 +836,30 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
}));
}
// 默认分页模式: limit 硬上限 2000 行(防 LLM 传超大 limit 读全文件,1MB 限下仍可能数万行)
let result = if let Some(offset) = args["offset"].as_u64() {
// BUG-260617-11: 无 offset 时旧实现 content.clone() 全量返回大文件,
// 虽 1MB 字节上限挡住极端情况,但万行级源码全量进 LLM context 仍易撑爆。
// 改:无 offset 默认返回前 500 行 + has_more 提示翻页(对齐 read 工具常规用法)。
let line_count = content.lines().count();
let (result, offset_used, has_more) = if let Some(offset) = args["offset"].as_u64() {
let lines: Vec<&str> = content.lines().collect();
let skip = offset as usize;
let limit = args["limit"].as_u64().unwrap_or(200).min(2000) as usize;
lines.into_iter().skip(skip).take(limit).collect::<Vec<&str>>().join("\n")
let page: Vec<&str> = lines.into_iter().skip(skip).take(limit).collect();
let more = (skip + page.len()) < line_count;
(page.join("\n"), Some(skip), more)
} else {
content.clone()
// 无 offset 默认前 500 行(大文件翻页友好,避免一次灌入全量)
const DEFAULT_PREVIEW_LINES: usize = 500;
let page: Vec<&str> = content.lines().take(DEFAULT_PREVIEW_LINES).collect();
let more = line_count > page.len();
(page.join("\n"), None, more)
};
let line_count = content.lines().count();
Ok(serde_json::json!({ "path": path, "content": result, "size": metadata.len(), "lines": line_count }))
Ok(serde_json::json!({
"path": path, "content": result, "size": metadata.len(), "lines": line_count,
"offset": offset_used,
"returned_lines": result.lines().count(),
"has_more": has_more,
}))
})),
);
registry.register(
@@ -814,7 +873,9 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
let path = resolved.to_str().ok_or_else(|| anyhow::anyhow!("路径含非法字符"))?;
let recursive = args["recursive"].as_bool().unwrap_or(false);
let skip_noise = args["skip_noise_dirs"].as_bool().unwrap_or(true);
let max_depth = args["max_depth"].as_u64().unwrap_or(3) as usize;
// BUG-260617-09: max_depth 由 LLM 参数控制,无上限时虽 entries 上限(1000)
// 隐式约束,但深递归仍可能大量 fs IO / 撑爆上下文。clamp 到合理范围 1-10。
let max_depth = args["max_depth"].as_u64().unwrap_or(3).clamp(1, 10) as usize;
let mut entries = Vec::new();
let truncated = list_dir_recursive(path, recursive, 0, max_depth, 1000, skip_noise, &mut entries).await?;
Ok(serde_json::json!({ "path": path, "entries": entries, "truncated": truncated }))
@@ -1120,10 +1181,25 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
let modified = metadata.modified()
.ok().and_then(|t| t.duration_since(std::time::UNIX_EPOCH).ok())
.map(|d| d.as_millis() as i64);
// is_binary: 读前 8KB 检测 \x00
// is_binary: 流式读前 8KB 检测 \x00 (BUG-260617-02)
// 旧实现 tokio::fs::read(path) 把整个文件读进内存再切片前 8192,
// >2MB 文件触发 OOM(注释">2MB 跳过"只作用于 lines, is_binary 无防护)。
// 改:File::open + BufReader + .take(8192) 只读前 N 字节做二进制检测,
// size/modified 等元信息另从 metadata 取(上面已取),不依赖全量读。
let is_binary = if !is_dir && size > 0 {
let sample = tokio::fs::read(path).await.unwrap_or_default();
sample[..sample.len().min(8192)].contains(&0x00)
use tokio::io::AsyncReadExt;
let file = match tokio::fs::File::open(path).await {
Ok(f) => f,
Err(_) => return Ok(serde_json::json!({
"path": path, "exists": true, "size": size, "lines": serde_json::Value::Null,
"modified": modified, "is_binary": false, "is_dir": is_dir,
"error": "读取文件头失败"
})),
};
let mut reader = tokio::io::BufReader::new(file);
let mut sample = vec![0u8; 8192];
let n = reader.read(&mut sample).await.unwrap_or(0);
sample[..n].contains(&0x00)
} else { false };
// lines: 文本文件 \n 计数(>2MB 跳过避免全量读)
let lines = if !is_dir && !is_binary && size <= 2_097_152 {