修复: gen_stream判重+重试分类+状态机最短路径+既有测试修复
This commit is contained in:
@@ -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", "同三元组应返回最新落定记录");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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));
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user