修复: 全库走查 P0+P1(df-mcp 数据丢失/evaluate 拆 + df-execute probe_pwsh 死缓存/超时/单源 + ai_self_review 注入隔离)

- df-mcp P0: update_idea/project/task 缺省回退 existing(治部分更新清空 title 丢数据)+ P1: evaluate_idea 拆只读(Low)+score_idea(Medium 写),read-only 不再改库
- df-execute P0: probe_pwsh 死缓存(两 OnceLock 合并 PWSH_CACHE 单源,治 Windows 永走 PS5)+ P1: probe_pwsh 3s 超时防挂起 + detect_shell 复用 shell.rs 单源(探测与执行对齐)
- df-nodes P1: ai_self_review prompt 注入隔离(truncate + XML 标签 <task_output> 数据/指令隔离 + system 声明)
- generate_image: SSRF 集成测试 + b64 OOM 估算纯函数测试(走查测试增强)
This commit is contained in:
lxy
2026-08-02 03:11:44 +08:00
parent dffc4e4851
commit fc249adf17
7 changed files with 1000 additions and 82 deletions
+188 -4
View File
@@ -62,6 +62,22 @@ const MAX_IMAGE_BYTES: u64 = 50 * 1024 * 1024;
/// 生成图片 POST 请求超时(秒)。图像生成模型耗时较高(高分辨率 10-30s),给 120s 余量。
const GENERATE_TIMEOUT_SECS: u64 = 120;
/// 按 base64 长度估算解码后字节数(base64 每 4 字符编码 3 字节)。
///
/// 用于 b64_json 解码前的 OOM 防护:恶意 provider 返超长 b64_json(如 200MB base64 → ~150MB
/// 解码字节),若先 `STANDARD.decode` 全载入内存再检查上限,瞬时 OOM。故解码前先用本函数
/// 估算,超 MAX_IMAGE_BYTES 直接 bail 不解码。
///
/// 公式 `len * 3 / 4`:忽略 padding(`=`)和 whitespace 误差,估值**略小于实际** b64 长度对应
/// 的理论解码字节(实际含 padding 会更少)。安全方向:估算偏保守(偏小),配合解码后兜底校验
/// (`decoded.len() > MAX_IMAGE_BYTES`)双保险。`saturating_mul` 防 usize 溢出(64 位平台
/// b64.len() 不可能逼近 usize::MAX,但防御性编程)。
///
/// 抽为独立纯函数供单测覆盖边界(2026-08-02 走查修复,对齐 test-strategy 防回归)。
fn estimate_decoded_bytes(b64_len: usize) -> u64 {
(b64_len as u64).saturating_mul(3) / 4
}
/// generate_image 工具 handler 入口(供 tools/generate_image.rs register 调用)。
///
/// 参数:
@@ -229,9 +245,8 @@ pub(crate) async fn execute_generate_image(
let b64 = b64_opt.as_ref().unwrap();
// OOM 防护(2026-08-02 走查修复):解码前先按 base64 长度估算解码后字节数,超 MAX_IMAGE_BYTES
// 直接 bail 不解码。恶意 provider 返超长 b64_json(如 200MB base64 → ~150MB 解码字节),
// 若先 STANDARD.decode 全载入内存再检查,瞬时 OOM。估算公式 len * 3 / 4(base64 每 4 字符
// 编码 3 字节),忽略 padding 误差(估值略大于实际,安全方向偏向拒)。
let estimated_decoded = (b64.len() as u64).saturating_mul(3) / 4;
// 若先 STANDARD.decode 全载入内存再检查,瞬时 OOM。详见 `estimate_decoded_bytes` 文档。
let estimated_decoded = estimate_decoded_bytes(b64.len());
if estimated_decoded > MAX_IMAGE_BYTES {
anyhow::bail!(
"b64_json 估算解码后约 {} 字节超过 {} 上限(原始 base64 长度 {})",
@@ -451,7 +466,9 @@ fn body_snippet(v: &Value) -> String {
// 覆盖:
// ① build_images_url:端点拼接规则(纯函数,零依赖)
// ② select_image_provider:provider 筛选逻辑(纯函数,内存构造 record)
// ③ handler 参数边界(prompt 缺失/空、provider_id 不存在)
// ③ estimate_decoded_bytes:b64 解码字节数估算(纯函数,OOM 防护核心)
// ④ handler 参数边界(prompt 缺失/空、provider_id 不存在)
// ⑤ handler SSRF 防护(endpoint 走 validate_url 前置拒,不发真实请求,CI 可跑)
// 真实 API 调用走 #[ignore](CI 无凭证/无网时跳过,本地手跑)。
// ============================================================
@@ -581,6 +598,80 @@ mod tests {
assert_eq!(got.id, "first");
}
// ── estimate_decoded_bytes:b64 解码字节数估算(OOM 防护核心纯函数) ──
#[test]
fn estimate_decoded_bytes_normal_small() {
// 1MB base64 → 估算 ~0.75MB(base64 每 4 字符编码 3 字节)
let one_mb = 1024 * 1024;
assert_eq!(estimate_decoded_bytes(one_mb), (one_mb as u64) * 3 / 4);
assert_eq!(estimate_decoded_bytes(one_mb), 786_432); // 1MB * 3 / 4
}
#[test]
fn estimate_decoded_bytes_over_limit() {
// 200MB 解码后场景:b64 长约 267MB → 估算 ~200MB > 50MB 上限,OOM 防护应挡
let two_hundred_mb_b64 = (200 * 1024 * 1024) as usize; // 假设这是 b64 长度
let estimated = estimate_decoded_bytes(two_hundred_mb_b64);
assert!(
estimated > MAX_IMAGE_BYTES,
"200MB b64 估算 {} 应超 {} 上限,OOM 防护靠此拒",
estimated,
MAX_IMAGE_BYTES
);
}
#[test]
fn estimate_decoded_bytes_boundary_at_max() {
// 边界:估算正好等于 MAX_IMAGE_BYTES → 比较用 `>`(超过才拒),正好等于算合法通过
// 反推 b64 长度:要使 len*3/4 == MAX_IMAGE_BYTES,len = MAX_IMAGE_BYTES*4/3
let b64_len = (MAX_IMAGE_BYTES as u128 * 4 / 3) as usize;
let estimated = estimate_decoded_bytes(b64_len);
// 估算值应 ≤ MAX_IMAGE_BYTES(因整数除法向下取整,可能略小于上限)
assert!(
estimated <= MAX_IMAGE_BYTES,
"估算 {} 应 ≤ 上限 {}(正好或略小于,`>` 比较下不拒)",
estimated,
MAX_IMAGE_BYTES
);
}
#[test]
fn estimate_decoded_bytes_just_over_max() {
// 边界:估算值 = MAX_IMAGE_BYTES + 1 → 应被拒
// 反推 b64 长度:要使 len*3/4 == MAX+1,len = (MAX+1)*4/3
let target = MAX_IMAGE_BYTES + 1;
let b64_len = (target as u128 * 4 / 3) as usize;
let estimated = estimate_decoded_bytes(b64_len);
// 因整数除法,估算可能略小于 target,但应保证 > MAX_IMAGE_BYTES
// 用更精确的 b64_len:target*4/3 + 4(向上补一个 base64 块),保证估算严格 > MAX
let b64_len_safe = ((target as u128 * 4 + 2) / 3) as usize;
let estimated_safe = estimate_decoded_bytes(b64_len_safe);
assert!(
estimated_safe > MAX_IMAGE_BYTES,
"b64_len={} 估算 {} 应严格 > 上限 {}(超限应拒)",
b64_len_safe,
estimated_safe,
MAX_IMAGE_BYTES
);
}
#[test]
fn estimate_decoded_bytes_saturating_no_overflow() {
// 极大值(u64::MAX 的 usize 等价)→ saturating_mul 防 panic,返回 u64 内合理值
// 注:usize 在 64 位平台 = u64,saturating_mul(u64::MAX, 3) 会 saturate 到 u64::MAX
let huge = usize::MAX;
let estimated = estimate_decoded_bytes(huge);
// saturating_mul(MAX, 3) = MAX,再 / 4 = MAX/4,不 panic 即通过
assert_eq!(estimated, u64::MAX / 4);
}
#[test]
fn estimate_decoded_bytes_zero() {
// 空输入 → 0(边界,虽 b64_opt 已 filter 空串不进此路径,纯函数仍应优雅)
assert_eq!(estimate_decoded_bytes(0), 0);
}
// ── handler 参数边界(走 db 需 in-memory db) ──
#[tokio::test]
@@ -611,4 +702,97 @@ mod tests {
let err = execute_generate_image(args, &db, &allowed).await.unwrap_err();
assert!(format!("{}", err).contains("无可用") || format!("{}", err).contains("provider"));
}
// ── handler SSRF 防护(endpoint 走 validate_url 前置拒,不发真实请求) ──
//
// base_url 来自 DB 用户配置(设置页可填任意 URL),provider 攻陷 / 配置错误 / 恶意 base_url
// 即可打内网。endpoint POST 须走 SSRF 防护(validate_url + resolve_and_check_host),与图片
// URL 下载同源(2026-08-02 走查修复)。本组测试构造 enabled + openai_compat + api_key 非空
// 的 provider 写入 in-memory db,handler 走到 validate_url 阶段拒,不发真实请求(CI 可跑)。
// 风格对齐 http.rs 的 test_handler_rejects_localhost_url / test_handler_rejects_private_ip_url。
/// 构造带 api_key 的 openai_compat provider(复用 mk_provider 但填 api_key,跳过 keyring)。
fn mk_provider_with_key(id: &str, base_url: &str) -> AiProviderRecord {
AiProviderRecord {
api_key: "sk-test-dummy-not-real".into(),
..mk_provider(id, true, "openai_compat", base_url)
}
}
#[tokio::test]
async fn handler_rejects_localhost_base_url() {
// base_url = http://localhost:8080 → validate_url L117-119 拒 localhost(SSRF 防护)
// 在 build_client 前拒,不发真实请求
let db = Arc::new(Database::open_in_memory().await.unwrap());
AiProviderRepo::new(&db)
.insert(mk_provider_with_key("p1", "http://localhost:8080"))
.await
.unwrap();
let allowed = Arc::new(RwLock::new(AllowedDirs::default()));
let args = json!({ "prompt": "一只猫", "provider_id": "p1" });
let err = execute_generate_image(args, &db, &allowed).await.unwrap_err();
let msg = format!("{}", err);
assert!(
msg.contains("localhost") || msg.contains("SSRF"),
"期望 SSRF 拒绝 localhost,实际: {}",
msg
);
}
#[tokio::test]
async fn handler_rejects_private_ip_base_url() {
// base_url 含 169.254.169.254(云元数据服务,SSRF 头号目标)→ validate_url 字面量 IP 私网拒
let db = Arc::new(Database::open_in_memory().await.unwrap());
AiProviderRepo::new(&db)
.insert(mk_provider_with_key("p1", "http://169.254.169.254/latest"))
.await
.unwrap();
let allowed = Arc::new(RwLock::new(AllowedDirs::default()));
let args = json!({ "prompt": "一只猫", "provider_id": "p1" });
let err = execute_generate_image(args, &db, &allowed).await.unwrap_err();
let msg = format!("{}", err);
assert!(
msg.contains("私网") || msg.contains("SSRF") || msg.contains("169.254"),
"期望 SSRF 拒绝元数据 IP,实际: {}",
msg
);
}
#[tokio::test]
async fn handler_rejects_loopback_base_url() {
// base_url = http://127.0.0.1 → validate_url 字面量 IP 环回拒
let db = Arc::new(Database::open_in_memory().await.unwrap());
AiProviderRepo::new(&db)
.insert(mk_provider_with_key("p1", "http://127.0.0.1:9000"))
.await
.unwrap();
let allowed = Arc::new(RwLock::new(AllowedDirs::default()));
let args = json!({ "prompt": "一只猫", "provider_id": "p1" });
let err = execute_generate_image(args, &db, &allowed).await.unwrap_err();
let msg = format!("{}", err);
assert!(
msg.contains("私网") || msg.contains("SSRF") || msg.contains("127.0.0.1"),
"期望 SSRF 拒绝环回 IP,实际: {}",
msg
);
}
#[tokio::test]
async fn handler_rejects_non_http_scheme_base_url() {
// base_url = file:///etc/passwd → validate_url L104-109 拒非 http 协议
let db = Arc::new(Database::open_in_memory().await.unwrap());
AiProviderRepo::new(&db)
.insert(mk_provider_with_key("p1", "file:///etc/passwd"))
.await
.unwrap();
let allowed = Arc::new(RwLock::new(AllowedDirs::default()));
let args = json!({ "prompt": "一只猫", "provider_id": "p1" });
let err = execute_generate_image(args, &db, &allowed).await.unwrap_err();
let msg = format!("{}", err);
assert!(
msg.contains("协议"),
"期望拒绝非 http 协议,实际: {}",
msg
);
}
}