修复: gen_stream判重+重试分类+状态机最短路径+既有测试修复

This commit is contained in:
lxy
2026-08-12 00:29:06 +08:00
parent 9407f821b9
commit d31a27512c
36 changed files with 2614 additions and 264 deletions
+103 -4
View File
@@ -319,6 +319,38 @@ impl AiToolExecutionRepo {
.map_err(storage_err)?
}
/// 按 (conversation_id, tool_name, arguments) 查最新一条审计记录(重试防护用)。
///
/// 语义:同一会话内同一工具同一参数的重试判重。区别于 `find_by_tool_call_id`
/// (裸 id 匹配,弱模型 provider 的 tool_call_id 每轮重排,裸 id 判重会误杀跨轮合法调用,
/// 实证 2026-08-11)。
pub async fn find_by_conv_tool_args(
&self,
conversation_id: &str,
tool_name: &str,
arguments: &str,
) -> Result<Option<AiToolExecutionRecord>> {
let conn = self.conn.clone();
let cid = conversation_id.to_owned();
let tname = tool_name.to_owned();
let args = arguments.to_owned();
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
let mut stmt = guard
.prepare(
"SELECT * FROM ai_tool_executions WHERE conversation_id = ?1 AND tool_name = ?2 AND arguments = ?3 ORDER BY requested_at DESC LIMIT 1",
)
.map_err(storage_err)?;
let row = stmt
.query_row(params![cid, tname, args], |row| ai_tool_execution_from_row(row))
.optional()
.map_err(storage_err)?;
Ok(row)
})
.await
.map_err(storage_err)?
}
/// 列出所有 status=pending 的审计行(启动重建 pending_approvals 用)
///
/// 专用 SELECT(非 query 宏——后者硬编码 ORDER BY created_at,而本表无该列)。
@@ -428,9 +460,9 @@ impl AiToolExecutionRepo {
params_vec.push(Box::new(r.clone()));
}
if let Some(k) = &kw {
let escaped = k.replace('%', "\\%").replace('_', "\\_");
let escaped = k.replace('|', "||").replace('%', "|%").replace('_', "|_");
let pat = format!("%{escaped}%");
where_clauses.push(format!("tool_name LIKE ?{} ESCAPE '\\'", params_vec.len() + 1));
where_clauses.push(format!("tool_name LIKE ?{} ESCAPE '|'", params_vec.len() + 1));
params_vec.push(Box::new(pat));
}
@@ -491,9 +523,9 @@ impl AiToolExecutionRepo {
params_vec.push(Box::new(r.clone()));
}
if let Some(k) = &kw {
let escaped = k.replace('%', "\\%").replace('_', "\\_");
let escaped = k.replace('|', "||").replace('%', "|%").replace('_', "|_");
let pat = format!("%{escaped}%");
where_clauses.push(format!("tool_name LIKE ?{} ESCAPE '\\'", params_vec.len() + 1));
where_clauses.push(format!("tool_name LIKE ?{} ESCAPE '|'", params_vec.len() + 1));
params_vec.push(Box::new(pat));
}
@@ -844,4 +876,71 @@ mod tests {
assert_eq!(legacy.model_configs.len(), 2, "老字符串数组应转 2 个默认 ModelConfig");
assert_eq!(legacy.model_configs[0].model_id, "glm-4-flash");
}
// ---------- find_by_conv_tool_args(重试判重,SC-260811-P0-1) ----------
/// 构造一条审计记录(参数自洽,满足 NOT NULL 约束)。
fn audit_rec(
id: &str, conv: &str, tc_id: &str, tool: &str, args: &str, status: &str,
) -> AiToolExecutionRecord {
AiToolExecutionRecord {
id: id.into(),
conversation_id: Some(conv.into()),
message_id: None,
tool_call_id: tc_id.into(),
tool_name: tool.into(),
arguments: args.into(),
result: None,
status: status.into(),
risk_level: "low".into(),
requested_at: "0".into(),
executed_at: None,
decided_by: None,
}
}
/// P0-1 核心语义:弱模型跨轮复用同 tool_call_id(gen_stream_0),但工具/参数不同 →
/// find_by_conv_tool_args 不得误判为重试(旧 find_by_tool_call_id 会误命中)。
#[tokio::test]
async fn find_by_conv_tool_args_ignores_same_id_different_tool() {
let db = Database::open_in_memory().await.expect("open_in_memory");
let repo = AiToolExecutionRepo::new(&db);
// 会话开头已执行过 list_tasks(gen_stream_0),之后同会话又出现 gen_stream_0 的 advance_task
repo.insert(audit_rec("a1", "conv1", "gen_stream_0", "list_tasks", "{\"limit\":50}", "completed")).await.expect("insert");
repo.insert(audit_rec("a2", "conv1", "gen_stream_0", "advance_task", "{\"id\":\"t1\",\"target_status\":\"blocked\"}", "completed")).await.expect("insert");
// 查 advance_task 同三元组 → 命中(a2)
let got = repo.find_by_conv_tool_args("conv1", "advance_task", "{\"id\":\"t1\",\"target_status\":\"blocked\"}").await.expect("query");
assert!(got.is_some(), "同 (conv,tool,args) 应命中");
assert_eq!(got.unwrap().id, "a2");
// 查 list_tasks 同三元组 → 命中(a1),互不串扰
let got2 = repo.find_by_conv_tool_args("conv1", "list_tasks", "{\"limit\":50}").await.expect("query");
assert!(got2.is_some());
assert_eq!(got2.unwrap().id, "a1");
// 跨会话同 tool_call_id 不命中(不同 conv)
let got3 = repo.find_by_conv_tool_args("conv2", "advance_task", "{\"id\":\"t1\",\"target_status\":\"blocked\"}").await.expect("query");
assert!(got3.is_none(), "不同会话同参数不应命中");
}
/// 同 (conv,tool,args) 且已落定(completed) → 判重返回记录(调用方 retry_count≥1);
/// pending 视为"尚未落定"(首次挂起审批)不判重。
#[tokio::test]
async fn find_by_conv_tool_args_pending_not_retry() {
let db = Database::open_in_memory().await.expect("open_in_memory");
let repo = AiToolExecutionRepo::new(&db);
repo.insert(audit_rec("a1", "conv1", "call_x", "write_file", "{\"path\":\"a.go\"}", "pending")).await.expect("insert");
// pending 记录不视为重试
let got = repo.find_by_conv_tool_args("conv1", "write_file", "{\"path\":\"a.go\"}").await.expect("query");
assert!(got.is_some(), "pending 记录也应能查到(状态判定在调用方 detect_retry_count 做,此处只负责按三元组定位)");
assert_eq!(got.unwrap().status, "pending");
// 同三元组 latest(有 completed 应返回最新一条 requested_at DESC)
let mut rec2 = audit_rec("a2", "conv1", "call_x", "write_file", "{\"path\":\"a.go\"}", "completed");
rec2.requested_at = "1".into(); // 晚于 a1 的 "0",保证 DESC 序可辨
repo.insert(rec2).await.expect("insert");
let got2 = repo.find_by_conv_tool_args("conv1", "write_file", "{\"path\":\"a.go\"}").await.expect("query");
assert_eq!(got2.unwrap().id, "a2", "同三元组应返回最新落定记录");
}
}
+6 -5
View File
@@ -301,11 +301,12 @@ impl IdeaRepo {
params_vec.push(Box::new(s.clone()));
}
if let Some(kw) = &keyword {
let escaped = kw.replace('%', "\\%").replace('_', "\\_");
let escaped = kw.replace('|', "||").replace('%', "|%").replace('_', "|_");
let pat = format!("%{escaped}%");
let p1 = params_vec.len() + 1;
let p2 = p1 + 1;
where_clauses.push(format!("(title LIKE ?{p1} OR description LIKE ?{p2}) ESCAPE '\\'"));
// ESCAPE 跟单个 LIKE(不能跟括号分组,否则 near "ESCAPE" syntax error)
where_clauses.push(format!("(title LIKE ?{p1} ESCAPE '|' OR description LIKE ?{p2} ESCAPE '|')"));
params_vec.push(Box::new(pat.clone()));
params_vec.push(Box::new(pat));
}
@@ -638,7 +639,7 @@ impl KnowledgeRepo {
/// 克制检索: top-N≤3(由调用方 limit 控制),精确匹配优先(语义模糊后做)。
pub async fn search(&self, query: &str, kind: Option<&str>, limit: usize) -> Result<Vec<KnowledgeRecord>> {
let conn = self.conn.clone();
let escaped = query.replace('%', "\\%").replace('_', "\\_");
let escaped = query.replace('|', "||").replace('%', "|%").replace('_', "|_");
let pattern = format!("%{escaped}%");
let kind = kind.map(|s| s.to_owned());
let limit_i = limit as i64;
@@ -647,7 +648,7 @@ impl KnowledgeRepo {
let mut results = Vec::new();
if let Some(k) = &kind {
let mut stmt = guard
.prepare(&format!("SELECT {KNOWLEDGE_COLS} FROM knowledges WHERE status = 'published' AND (title LIKE ?1 ESCAPE '\\' OR content LIKE ?2 ESCAPE '\\') AND kind = ?3 ORDER BY reuse_count DESC LIMIT ?4"))
.prepare(&format!("SELECT {KNOWLEDGE_COLS} FROM knowledges WHERE status = 'published' AND (title LIKE ?1 ESCAPE '|' OR content LIKE ?2 ESCAPE '|') AND kind = ?3 ORDER BY reuse_count DESC LIMIT ?4"))
.map_err(storage_err)?;
let rows = stmt
.query_map(params![pattern, pattern, k, limit_i], |row| knowledge_from_row(row))
@@ -657,7 +658,7 @@ impl KnowledgeRepo {
}
} else {
let mut stmt = guard
.prepare(&format!("SELECT {KNOWLEDGE_COLS} FROM knowledges WHERE status = 'published' AND (title LIKE ?1 ESCAPE '\\' OR content LIKE ?2 ESCAPE '\\') ORDER BY reuse_count DESC LIMIT ?3"))
.prepare(&format!("SELECT {KNOWLEDGE_COLS} FROM knowledges WHERE status = 'published' AND (title LIKE ?1 ESCAPE '|' OR content LIKE ?2 ESCAPE '|') ORDER BY reuse_count DESC LIMIT ?3"))
.map_err(storage_err)?;
let rows = stmt
.query_map(params![pattern, pattern, limit_i], |row| knowledge_from_row(row))
@@ -215,19 +215,30 @@ use rusqlite::OptionalExtension;
#[cfg(test)]
mod tests {
use super::*;
use crate::crud::project_repo::ProjectRepo;
use crate::db::Database;
use crate::models::{ProjectRecord, ProjectModuleRecord};
use df_types::types::ProjectStatus;
/// 建库 + 建占位 project 满足 FK 约束 + 返回 repo(对标 project_service_repo::setup)。
async fn setup() -> (Database, ProjectModuleRepo, String) {
let db = Database::open_in_memory().await.expect("open_in_memory");
let project_id = "proj-test".to_string();
db.conn()
.blocking_lock()
.execute(
"INSERT INTO projects (id, name, status, path, stack, created_at, updated_at) \
VALUES (?1, ?2, 'active', '/tmp', 'rust', '0', '0')",
params![project_id, "Test Project"],
)
// 用异步 repo API 建占位 project(勿在 async 上下文直接 blocking_lock,会触发
// "Cannot block the current thread from within a runtime" panic,实证 2026-08-11)。
ProjectRepo::new(&db)
.insert(ProjectRecord {
id: project_id.clone(),
name: "Test Project".to_string(),
description: String::new(),
status: ProjectStatus::InProgress,
idea_id: None,
path: Some("/tmp".into()),
stack: Some("rust".into()),
created_at: "0".to_string(),
updated_at: "0".to_string(),
})
.await
.expect("insert placeholder project");
let repo = ProjectModuleRepo::new(&db);
(db, repo, project_id)
+3 -2
View File
@@ -315,8 +315,9 @@ impl ProjectRepo {
if let Some(kw) = &q.keyword {
let trimmed = kw.trim();
if !trimmed.is_empty() {
let escaped = trimmed.replace('%', "\\%").replace('_', "\\_");
sql.push_str(" AND (name LIKE ? OR description LIKE ?) ESCAPE '\\'");
let escaped = trimmed.replace('|', "||").replace('%', "|%").replace('_', "|_");
// ESCAPE 跟单个 LIKE(不能跟括号分组,否则 near "ESCAPE" syntax error)
sql.push_str(" AND (name LIKE ? ESCAPE '|' OR description LIKE ? ESCAPE '|')");
let pattern = format!("%{escaped}%");
params_vec.push(Box::new(pattern.clone()));
params_vec.push(Box::new(pattern));
+50 -4
View File
@@ -391,12 +391,17 @@ impl TaskRepo {
params_vec.push(Box::new(mid.clone()));
}
// keyword: title/description LIKE %kw%(P2,对齐知识库 search 的 LIKE 模式)
// ESCAPE 字符用 |(管道符,任务标题/描述几乎不含),不用反斜杠——反斜杠在
// Rust format! → rusqlite 绑定 → SQLite 多层转义里极易出错(SQLite 报
// "ESCAPE expression must be a single character"),改 | 一劳永逸。
// 注意:ESCAPE 只能跟单个 LIKE,不能跟括号分组(实测 `(a OR b) ESCAPE 'x'`
// 报 near "ESCAPE" syntax error),故每个 LIKE 各自 ESCAPE。
if let Some(kw) = &keyword {
let escaped = kw.replace('%', "\\%").replace('_', "\\_");
let escaped = kw.replace('|', "||").replace('%', "|%").replace('_', "|_");
let pat = format!("%{escaped}%");
let p1 = params_vec.len() + 1;
let p2 = p1 + 1;
where_clauses.push(format!("(title LIKE ?{p1} OR description LIKE ?{p2}) ESCAPE '\\'"));
where_clauses.push(format!("(title LIKE ?{p1} ESCAPE '|' OR description LIKE ?{p2} ESCAPE '|')"));
params_vec.push(Box::new(pat.clone()));
params_vec.push(Box::new(pat));
}
@@ -494,11 +499,11 @@ impl TaskRepo {
params_vec.push(Box::new(a.clone()));
}
if let Some(ref kw) = keyword {
let escaped = kw.replace('%', "\\%").replace('_', "\\_");
let escaped = kw.replace('|', "||").replace('%', "|%").replace('_', "|_");
let pat = format!("%{escaped}%");
let p1 = params_vec.len() + 1;
let p2 = p1 + 1;
where_clauses.push(format!("(title LIKE ?{p1} OR description LIKE ?{p2}) ESCAPE '\\'"));
where_clauses.push(format!("(title LIKE ?{p1} ESCAPE '|' OR description LIKE ?{p2} ESCAPE '|')"));
params_vec.push(Box::new(pat.clone()));
params_vec.push(Box::new(pat));
}
@@ -1358,4 +1363,45 @@ mod tests {
let after = repo.get_by_id("t1").await.unwrap().unwrap();
assert_eq!(after.title, "新标题", "软删后字段不应被改动");
}
// ============================================================
// keyword LIKE 查询(ESCAPE 转义)—— 防 2026-08-11 语法回归
// ============================================================
// 背景:旧写法 `(a LIKE ?1 OR b LIKE ?2) ESCAPE '|'`(ESCAPE 跟括号分组)在 SQLite
// 报 near "ESCAPE" syntax error,keyword 查询全挂。正确写法:ESCAPE 跟每个 LIKE。
// 本测试锁两种语义:普通子串匹配 + 含 %/_ 通配符字面匹配(转义生效)。
#[tokio::test]
async fn list_by_query_keyword_matches_substring() {
let repo = setup().await;
repo.insert(trec("t1", "todo", None)).await.unwrap();
repo.insert(trec("t2", "todo", None)).await.unwrap();
// 定制 title:t1 含「支付」,t2 不含
repo.update_field_active("t1", "title", "海外支付集成").await.unwrap();
let rows = repo
.list_by_query(&TaskQuery { keyword: Some("支付".into()), ..Default::default() })
.await
.unwrap();
assert_eq!(rows.len(), 1, "keyword 子串应只命中 t1");
assert_eq!(rows[0].id, "t1");
}
#[tokio::test]
async fn list_by_query_keyword_escapes_wildcards() {
let repo = setup().await;
repo.insert(trec("t1", "todo", None)).await.unwrap();
repo.insert(trec("t2", "todo", None)).await.unwrap();
// t1 标题含字面 % 与 _(通配符需转义,按字面匹配)
repo.update_field_active("t1", "title", "比率 100%_cache").await.unwrap();
repo.update_field_active("t2", "title", "比率 100x_cache").await.unwrap();
// 查询字面 "%_"(含两个通配符,转义后应按字面匹配 t1;t2 的 x 不匹配 %)
let rows = repo
.list_by_query(&TaskQuery { keyword: Some("100%_".into()), ..Default::default() })
.await
.unwrap();
assert_eq!(rows.len(), 1, "% 与 _ 应被转义为字面,只命中 t1");
assert_eq!(rows[0].id, "t1");
}
}