修复: 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

@@ -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 {