修复: 全库走查 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:
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user