846 lines
37 KiB
Rust
846 lines
37 KiB
Rust
//! AI 对话域 Repo:AiProviderRepo / AiConversationRepo / AiToolExecutionRepo
|
|
|
|
use std::sync::Arc;
|
|
|
|
use rusqlite::{params, Connection, OptionalExtension, Row};
|
|
use tokio::sync::Mutex;
|
|
|
|
use df_types::error::Result;
|
|
|
|
use crate::db::Database;
|
|
use crate::models::{AiConversationRecord, AiProviderRecord, AiToolExecutionRecord};
|
|
|
|
use super::impl_repo;
|
|
use super::{now_millis_str, storage_err, validate_column_name};
|
|
|
|
// ============================================================
|
|
// from_row 辅助函数
|
|
// ============================================================
|
|
|
|
fn ai_provider_from_row(row: &Row<'_>) -> std::result::Result<AiProviderRecord, rusqlite::Error> {
|
|
// model_configs:DB TEXT 列存 JSON 字符串。读 Option<String> 兼容老库 NULL,
|
|
// 再经 deserialize_model_configs 解析(老字符串数组/新对象数组/空 → Vec<ModelConfig>)。
|
|
// 解析失败不致命:降级为空 Vec(防单行坏数据拖垮 list_all)。
|
|
let model_configs: Vec<df_ai_core::model::ModelConfig> = {
|
|
let raw: Option<String> = row.get("model_configs").ok();
|
|
match raw {
|
|
None => Vec::new(),
|
|
Some(s) => serde_json::from_str::<ModelConfigsWrap>(&format!(
|
|
r#"{{"v":{}}}"#,
|
|
if s.trim().is_empty() {
|
|
"null".to_string()
|
|
} else if s.trim_start().starts_with('[') || s.trim_start().starts_with('{') {
|
|
s
|
|
} else {
|
|
// 非 JSON 字面文本(理论不会出现)→ 包装为 JSON 字符串让 deserialize 兜底
|
|
serde_json::to_string(&s).unwrap_or_else(|_| "null".into())
|
|
}
|
|
))
|
|
.map(|w| w.v)
|
|
.unwrap_or_default(),
|
|
}
|
|
};
|
|
Ok(AiProviderRecord {
|
|
id: row.get("id")?,
|
|
name: row.get("name")?,
|
|
provider_type: row.get("provider_type")?,
|
|
api_key: row.get("api_key")?,
|
|
base_url: row.get("base_url")?,
|
|
default_model: row.get("default_model")?,
|
|
models: row.get("models")?,
|
|
model_configs,
|
|
is_default: row.get::<_, i32>("is_default")? != 0,
|
|
config: row.get("config")?,
|
|
created_at: row.get("created_at")?,
|
|
updated_at: row.get("updated_at")?,
|
|
// enabled/weight 列老库经 v19 迁移补建,DEFAULT 1 / DEFAULT 50。
|
|
// from_row 按 i32 取列值兼容(SQLite 无真 BOOLEAN),0→false/非0→true。
|
|
enabled: row.get::<_, i32>("enabled").unwrap_or(1) != 0,
|
|
// weight 读侧 clamp [0,100]:与 insert/update_full 落库的 `.min(100)` 对齐,
|
|
// 防老库(clamp 落地前写入的)或外部直改 DB 产生的越界值污染路由权重语义。
|
|
weight: row.get::<_, i32>("weight")
|
|
.unwrap_or(50)
|
|
.clamp(0, 100) as u32,
|
|
})
|
|
}
|
|
|
|
/// from_row 内部辅助:复用 deserialize_model_configs 解析 DB TEXT 列 JSON。
|
|
/// 包一层 { "v": <原始值> } 把任意 JSON 值送进 deserialize_model_configs。
|
|
#[derive(serde::Deserialize)]
|
|
struct ModelConfigsWrap {
|
|
#[serde(default, deserialize_with = "df_ai_core::model::deserialize_model_configs")]
|
|
v: Vec<df_ai_core::model::ModelConfig>,
|
|
}
|
|
|
|
fn ai_conversation_from_row(row: &Row<'_>) -> std::result::Result<AiConversationRecord, rusqlite::Error> {
|
|
Ok(AiConversationRecord {
|
|
id: row.get("id")?,
|
|
title: row.get("title")?,
|
|
messages: row.get("messages")?,
|
|
provider_id: row.get("provider_id")?,
|
|
model: row.get("model")?,
|
|
models: row.get("models")?,
|
|
archived: row.get::<_, i32>("archived")? != 0,
|
|
pinned: row.get::<_, i32>("pinned")? != 0,
|
|
prompt_tokens: row.get("prompt_tokens")?,
|
|
completion_tokens: row.get("completion_tokens")?,
|
|
pinned_goals: row.get("pinned_goals")?,
|
|
pending_approvals: row.get("pending_approvals")?,
|
|
created_at: row.get("created_at")?,
|
|
updated_at: row.get("updated_at")?,
|
|
})
|
|
}
|
|
|
|
fn ai_tool_execution_from_row(row: &Row<'_>) -> std::result::Result<AiToolExecutionRecord, rusqlite::Error> {
|
|
Ok(AiToolExecutionRecord {
|
|
id: row.get("id")?,
|
|
conversation_id: row.get("conversation_id")?,
|
|
// message_id 列老库经 v21 迁移补建。unwrap_or(None) 兜底:
|
|
// 新库空表直接有列;老库行 ALTER 后 NULL;极端情况(迁移未跑/手工删列)防御。
|
|
message_id: row.get("message_id").unwrap_or(None),
|
|
tool_call_id: row.get("tool_call_id")?,
|
|
tool_name: row.get("tool_name")?,
|
|
arguments: row.get("arguments")?,
|
|
result: row.get("result")?,
|
|
status: row.get("status")?,
|
|
risk_level: row.get("risk_level")?,
|
|
requested_at: row.get("requested_at")?,
|
|
executed_at: row.get("executed_at")?,
|
|
decided_by: row.get("decided_by")?,
|
|
})
|
|
}
|
|
|
|
// ============================================================
|
|
// Repo 实现
|
|
// ============================================================
|
|
|
|
impl_repo!(
|
|
/// AI 提供商配置表 CRUD
|
|
AiProviderRepo,
|
|
AiProviderRecord,
|
|
"ai_providers",
|
|
from_row => |row| ai_provider_from_row(row),
|
|
insert => |conn, rec| {
|
|
let is_default = if rec.is_default { 1i32 } else { 0i32 };
|
|
// model_configs:Vec<ModelConfig> → JSON 字符串落 TEXT 列
|
|
let model_configs_json = serde_json::to_string(&rec.model_configs).unwrap_or_else(|_| "[]".into());
|
|
// enabled/weight 落库(SQLite 无 BOOLEAN,i32 承载)。
|
|
let enabled_i = if rec.enabled { 1i32 } else { 0i32 };
|
|
let weight_i = rec.weight.min(100) as i32;
|
|
conn.execute(
|
|
"INSERT OR REPLACE INTO ai_providers (id, name, provider_type, api_key, base_url, default_model, models, model_configs, is_default, config, created_at, updated_at, enabled, weight)
|
|
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)",
|
|
params![
|
|
rec.id, rec.name, rec.provider_type, rec.api_key, rec.base_url,
|
|
rec.default_model, rec.models, model_configs_json, is_default, rec.config, rec.created_at, rec.updated_at,
|
|
enabled_i, weight_i
|
|
],
|
|
)
|
|
},
|
|
update => |conn, rec| {
|
|
let is_default = if rec.is_default { 1i32 } else { 0i32 };
|
|
let model_configs_json = serde_json::to_string(&rec.model_configs).unwrap_or_else(|_| "[]".into());
|
|
let enabled_i = if rec.enabled { 1i32 } else { 0i32 };
|
|
let weight_i = rec.weight.min(100) as i32;
|
|
conn.execute(
|
|
"UPDATE ai_providers SET name = ?1, provider_type = ?2, api_key = ?3, base_url = ?4, default_model = ?5, models = ?6, model_configs = ?7, is_default = ?8, config = ?9, updated_at = ?10, enabled = ?11, weight = ?12 WHERE id = ?13",
|
|
params![
|
|
rec.name, rec.provider_type, rec.api_key, rec.base_url,
|
|
rec.default_model, rec.models, model_configs_json, is_default, rec.config, rec.updated_at,
|
|
enabled_i, weight_i, rec.id
|
|
],
|
|
)
|
|
}
|
|
);
|
|
|
|
impl_repo!(
|
|
/// AI 对话历史表 CRUD
|
|
AiConversationRepo,
|
|
AiConversationRecord,
|
|
"ai_conversations",
|
|
from_row => |row| ai_conversation_from_row(row),
|
|
insert => |conn, rec| {
|
|
conn.execute(
|
|
"INSERT INTO ai_conversations (id, title, messages, provider_id, model, models, archived, pinned, prompt_tokens, completion_tokens, pinned_goals, pending_approvals, created_at, updated_at)
|
|
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)",
|
|
params![
|
|
rec.id, rec.title, rec.messages, rec.provider_id, rec.model, rec.models, rec.archived,
|
|
if rec.pinned { 1i32 } else { 0i32 },
|
|
rec.prompt_tokens, rec.completion_tokens,
|
|
rec.pinned_goals, rec.pending_approvals, rec.created_at, rec.updated_at
|
|
],
|
|
)
|
|
},
|
|
update => |conn, rec| {
|
|
conn.execute(
|
|
"UPDATE ai_conversations SET title = ?1, messages = ?2, provider_id = ?3, model = ?4, models = ?5, archived = ?6, pinned = ?7, prompt_tokens = ?8, completion_tokens = ?9, pinned_goals = ?10, pending_approvals = ?11, updated_at = ?12 WHERE id = ?13",
|
|
params![
|
|
rec.title, rec.messages, rec.provider_id, rec.model, rec.models, rec.archived,
|
|
if rec.pinned { 1i32 } else { 0i32 },
|
|
rec.prompt_tokens, rec.completion_tokens,
|
|
rec.pinned_goals, rec.pending_approvals, rec.updated_at, rec.id
|
|
],
|
|
)
|
|
}
|
|
);
|
|
|
|
// ============================================================
|
|
// AuditQuery — 审批历史多条件查询入参(status / risk / 工具名关键词)
|
|
// ============================================================
|
|
|
|
/// 审批历史多条件查询入参(对标 [`IdeaQuery`] 的可选字段 struct 设计)。
|
|
///
|
|
/// 所有字段可选;全 None → 等价 `list_recent`(向后兼容)。设计对齐 `查询能力补全方案`:
|
|
/// 可选字段 struct 而非逐个加 IPC 参数,复用 [`IdeaRepo::list_by_query`] 的动态 WHERE 拼接
|
|
/// 模式(if-let 分支拼 SQL + 分支化参数绑定)。
|
|
///
|
|
/// - `status`:状态精确匹配(pending/approved/rejected/executing/completed/failed/interrupted)
|
|
/// - `risk_level`:风险等级精确匹配(low/medium/high)
|
|
/// - `tool_keyword`:`tool_name LIKE %kw%`(对齐 idea_repo 关键词 LIKE 检索,不上 FTS5)
|
|
/// - `limit`/`offset`:钳制上限 200(对齐 [`AiToolExecutionRepo::list_recent`])
|
|
///
|
|
/// `Deserialize`:Tauri IPC 从前端 JSON 反序列化为命令参数。
|
|
/// `Default`:命令层兼容旧全量调用(`AuditQuery::default()` 等价无条件)。
|
|
#[derive(Debug, Clone, Default, serde::Deserialize)]
|
|
pub struct AuditQuery {
|
|
pub status: Option<String>,
|
|
pub risk_level: Option<String>,
|
|
pub tool_keyword: Option<String>,
|
|
pub limit: Option<u32>,
|
|
pub offset: Option<u32>,
|
|
}
|
|
|
|
impl_repo!(
|
|
/// AI 工具执行审计表 CRUD
|
|
AiToolExecutionRepo,
|
|
AiToolExecutionRecord,
|
|
"ai_tool_executions",
|
|
from_row => |row| ai_tool_execution_from_row(row),
|
|
insert => |conn, rec| {
|
|
conn.execute(
|
|
"INSERT INTO ai_tool_executions (id, conversation_id, message_id, tool_call_id, tool_name, arguments, result, status, risk_level, requested_at, executed_at, decided_by)
|
|
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12)",
|
|
params![
|
|
rec.id, rec.conversation_id, rec.message_id, rec.tool_call_id, rec.tool_name,
|
|
rec.arguments, rec.result, rec.status, rec.risk_level,
|
|
rec.requested_at, rec.executed_at, rec.decided_by
|
|
],
|
|
)
|
|
},
|
|
update => |conn, rec| {
|
|
conn.execute(
|
|
"UPDATE ai_tool_executions SET conversation_id = ?1, message_id = ?2, tool_call_id = ?3, tool_name = ?4, arguments = ?5, result = ?6, status = ?7, risk_level = ?8, requested_at = ?9, executed_at = ?10, decided_by = ?11 WHERE id = ?12",
|
|
params![
|
|
rec.conversation_id, rec.message_id, rec.tool_call_id, rec.tool_name,
|
|
rec.arguments, rec.result, rec.status, rec.risk_level,
|
|
rec.requested_at, rec.executed_at, rec.decided_by, rec.id
|
|
],
|
|
)
|
|
}
|
|
);
|
|
|
|
// ai_tool_executions 无 created_at 列(用 requested_at/executed_at 计时),
|
|
// 通用 query 宏硬编码 ORDER BY created_at 会触发 "no such column" → 调用方 unwrap_or_default 吞错。
|
|
// 故为此表提供专用查询,绕过通用 query。详见 ai.rs audit_finalize。
|
|
impl AiToolExecutionRepo {
|
|
/// 批量插入审计记录(单事务多行 INSERT,砍 N 次串行 INSERT 尾巴)。
|
|
///
|
|
/// 对比 [`insert`](`impl_repo!` 生成,每次 spawn_blocking + 单行 execute):
|
|
/// 本方法单次 `spawn_blocking` + 单事务,`prepare` 一次 INSERT stmt 循环 bind N 行,
|
|
/// 一次 `COMMIT`(原子性:全插或全不插,审计留痕可追溯)。空 `records` 直接返回(无操作)。
|
|
///
|
|
/// **用途**:audit/mod.rs `process_tool_calls` 低风险工具 join_all 并行执行后的回填循环
|
|
/// (每工具一条审计),把 N 次串行 INSERT 合并为一次事务批量(治 aichat 效率 AC-EFF-T1-1)。
|
|
///
|
|
/// 安全:全部值走参数绑定(同 `insert` 宏体),无 SQL 拼接注入面;单连接 Mutex 持锁整段,
|
|
/// 与单行 insert 的锁粒度相同(一次持锁换 N 次持锁)。
|
|
pub async fn insert_batch(&self, records: Vec<AiToolExecutionRecord>) -> Result<()> {
|
|
if records.is_empty() {
|
|
return Ok(());
|
|
}
|
|
let conn = self.conn.clone();
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
guard.execute_batch("BEGIN").map_err(storage_err)?;
|
|
let result = (|| -> std::result::Result<(), rusqlite::Error> {
|
|
let mut stmt = guard.prepare(
|
|
"INSERT INTO ai_tool_executions (id, conversation_id, message_id, tool_call_id, tool_name, arguments, result, status, risk_level, requested_at, executed_at, decided_by)
|
|
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12)",
|
|
)?;
|
|
for rec in &records {
|
|
stmt.execute(params![
|
|
rec.id, rec.conversation_id, rec.message_id, rec.tool_call_id, rec.tool_name,
|
|
rec.arguments, rec.result, rec.status, rec.risk_level,
|
|
rec.requested_at, rec.executed_at, rec.decided_by
|
|
])?;
|
|
}
|
|
Ok(())
|
|
})();
|
|
match result {
|
|
Ok(()) => {
|
|
guard.execute_batch("COMMIT").map_err(storage_err)?;
|
|
}
|
|
Err(e) => {
|
|
// 回滚失败静默(尽量保一致性;rollback 失败通常是连接已坏,交给上层)
|
|
let _ = guard.execute_batch("ROLLBACK");
|
|
return Err(storage_err(e));
|
|
}
|
|
}
|
|
Ok(())
|
|
})
|
|
.await
|
|
.map_err(storage_err)??;
|
|
Ok(())
|
|
}
|
|
|
|
/// 按 tool_call_id 查最新一条审计记录(审批回填定位用)。
|
|
pub async fn find_by_tool_call_id(
|
|
&self,
|
|
tool_call_id: &str,
|
|
) -> Result<Option<AiToolExecutionRecord>> {
|
|
let conn = self.conn.clone();
|
|
let tid = tool_call_id.to_owned();
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
let mut stmt = guard
|
|
.prepare(
|
|
"SELECT * FROM ai_tool_executions WHERE tool_call_id = ?1 ORDER BY requested_at DESC LIMIT 1",
|
|
)
|
|
.map_err(storage_err)?;
|
|
let row = stmt
|
|
.query_row(params![tid], |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,而本表无该列)。
|
|
pub async fn list_pending(&self) -> Result<Vec<AiToolExecutionRecord>> {
|
|
let conn = self.conn.clone();
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
let mut stmt = guard
|
|
.prepare("SELECT * FROM ai_tool_executions WHERE status = 'pending' ORDER BY requested_at ASC")
|
|
.map_err(storage_err)?;
|
|
let rows = stmt
|
|
.query_map([], |row| ai_tool_execution_from_row(row))
|
|
.map_err(storage_err)?;
|
|
let mut results = Vec::new();
|
|
for r in rows {
|
|
results.push(r.map_err(storage_err)?);
|
|
}
|
|
Ok(results)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// 清理超期的残留 pending 工具调用(旧会话遗留)。
|
|
///
|
|
/// `max_age_secs`: 超过此秒数的 pending 记录被标记为 interrupted(不硬删,保留审计痕迹)。
|
|
pub async fn cleanup_stale_pending(&self, max_age_secs: u64) -> Result<u64> {
|
|
let conn = self.conn.clone();
|
|
let cutoff_ms = (df_types::now_millis() as i64 - (max_age_secs as i64 * 1000)).to_string();
|
|
let affected = tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
guard.execute(
|
|
"UPDATE ai_tool_executions SET status = 'interrupted' \
|
|
WHERE status = 'pending' AND CAST(requested_at AS INTEGER) < ?1",
|
|
params![cutoff_ms],
|
|
).map_err(storage_err)
|
|
})
|
|
.await
|
|
.map_err(storage_err)??;
|
|
Ok(affected as u64)
|
|
}
|
|
|
|
/// 审批历史面板分页查询:按 requested_at 倒序(最新在前),limit 默认 50。
|
|
///
|
|
/// 与 list_pending 同理走专用 SELECT,绕过通用 query 宏(后者硬编码
|
|
/// ORDER BY created_at,本表无该列)。limit/offset 上限钳制(limit ≤ 200),
|
|
/// 防前端恶意/失误传超大值。
|
|
pub async fn list_recent(
|
|
&self,
|
|
limit: u32,
|
|
offset: u32,
|
|
) -> Result<Vec<AiToolExecutionRecord>> {
|
|
let conn = self.conn.clone();
|
|
// 钳制 limit 防滥用(默认 50,最大 200)
|
|
let safe_limit = limit.min(200) as i64;
|
|
let safe_offset = offset as i64;
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
let mut stmt = guard
|
|
.prepare(
|
|
"SELECT * FROM ai_tool_executions ORDER BY requested_at DESC LIMIT ?1 OFFSET ?2",
|
|
)
|
|
.map_err(storage_err)?;
|
|
let rows = stmt
|
|
.query_map(params![safe_limit, safe_offset], |row| {
|
|
ai_tool_execution_from_row(row)
|
|
})
|
|
.map_err(storage_err)?;
|
|
let mut results = Vec::new();
|
|
for r in rows {
|
|
results.push(r.map_err(storage_err)?);
|
|
}
|
|
Ok(results)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// 多条件查询:动态 WHERE 拼接(status / risk_level / 工具名关键词) + 分页。
|
|
///
|
|
/// 复用 [`IdeaRepo::list_by_query`] 的动态 WHERE 模式:if-let 分支按可选条件拼 SQL 片段,
|
|
/// 各分支化参数绑定到 `?N` 占位符。limit 钳制上限 200(对齐 [`Self::list_recent`])。
|
|
///
|
|
/// **向后兼容**:空 query(全 None)→ 无 WHERE 子句,等价 `list_recent`。
|
|
/// 与 list_pending/list_recent 同理走专用 SELECT,绕过通用 query 宏(后者硬编码
|
|
/// ORDER BY created_at,本表无该列)。
|
|
pub async fn list_by_query(&self, q: &AuditQuery) -> Result<Vec<AiToolExecutionRecord>> {
|
|
let conn = self.conn.clone();
|
|
let status = q.status.clone();
|
|
let risk = q.risk_level.clone();
|
|
let kw = q.tool_keyword.clone();
|
|
let limit_i: i64 = q.limit.unwrap_or(50).min(200) as i64;
|
|
let offset_i: i64 = q.offset.unwrap_or(0) as i64;
|
|
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
|
|
let mut where_clauses: Vec<String> = Vec::new();
|
|
let mut params_vec: Vec<Box<dyn rusqlite::ToSql>> = Vec::new();
|
|
|
|
if let Some(s) = &status {
|
|
where_clauses.push(format!("status = ?{}", params_vec.len() + 1));
|
|
params_vec.push(Box::new(s.clone()));
|
|
}
|
|
if let Some(r) = &risk {
|
|
where_clauses.push(format!("risk_level = ?{}", params_vec.len() + 1));
|
|
params_vec.push(Box::new(r.clone()));
|
|
}
|
|
if let Some(k) = &kw {
|
|
let escaped = k.replace('%', "\\%").replace('_', "\\_");
|
|
let pat = format!("%{escaped}%");
|
|
where_clauses.push(format!("tool_name LIKE ?{} ESCAPE '\\'", params_vec.len() + 1));
|
|
params_vec.push(Box::new(pat));
|
|
}
|
|
|
|
let where_sql = if where_clauses.is_empty() {
|
|
String::new()
|
|
} else {
|
|
format!(" WHERE {}", where_clauses.join(" AND "))
|
|
};
|
|
|
|
let where_param_count = params_vec.len();
|
|
let sql = format!(
|
|
"SELECT * FROM ai_tool_executions{where_sql} \
|
|
ORDER BY requested_at DESC LIMIT ?{lim} OFFSET ?{off}",
|
|
lim = where_param_count + 1,
|
|
off = where_param_count + 2,
|
|
);
|
|
|
|
let mut stmt = guard.prepare(&sql).map_err(storage_err)?;
|
|
params_vec.push(Box::new(limit_i));
|
|
params_vec.push(Box::new(offset_i));
|
|
let param_refs: Vec<&dyn rusqlite::ToSql> =
|
|
params_vec.iter().map(|p| p.as_ref()).collect();
|
|
let rows = stmt
|
|
.query_map(param_refs.as_slice(), |row| ai_tool_execution_from_row(row))
|
|
.map_err(storage_err)?;
|
|
let mut results = Vec::new();
|
|
for r in rows {
|
|
results.push(r.map_err(storage_err)?);
|
|
}
|
|
Ok(results)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// 按 [`AuditQuery`] 条件计数(不含 limit/offset,用于分页 total)。
|
|
///
|
|
/// 复用 [`Self::list_by_query`] 的 WHERE 构造逻辑(仅 WHERE,无 ORDER BY/LIMIT),
|
|
/// 返回满足条件的总行数(忽略分页裁剪)。对标 [`TaskRepo::count_by_query`]。
|
|
pub async fn count_by_query(&self, q: &AuditQuery) -> Result<i64> {
|
|
let conn = self.conn.clone();
|
|
let status = q.status.clone();
|
|
let risk = q.risk_level.clone();
|
|
let kw = q.tool_keyword.clone();
|
|
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
|
|
let mut where_clauses: Vec<String> = Vec::new();
|
|
let mut params_vec: Vec<Box<dyn rusqlite::ToSql>> = Vec::new();
|
|
|
|
if let Some(s) = &status {
|
|
where_clauses.push(format!("status = ?{}", params_vec.len() + 1));
|
|
params_vec.push(Box::new(s.clone()));
|
|
}
|
|
if let Some(r) = &risk {
|
|
where_clauses.push(format!("risk_level = ?{}", params_vec.len() + 1));
|
|
params_vec.push(Box::new(r.clone()));
|
|
}
|
|
if let Some(k) = &kw {
|
|
let escaped = k.replace('%', "\\%").replace('_', "\\_");
|
|
let pat = format!("%{escaped}%");
|
|
where_clauses.push(format!("tool_name LIKE ?{} ESCAPE '\\'", params_vec.len() + 1));
|
|
params_vec.push(Box::new(pat));
|
|
}
|
|
|
|
let sql = if where_clauses.is_empty() {
|
|
"SELECT COUNT(*) FROM ai_tool_executions".to_string()
|
|
} else {
|
|
format!(
|
|
"SELECT COUNT(*) FROM ai_tool_executions WHERE {}",
|
|
where_clauses.join(" AND ")
|
|
)
|
|
};
|
|
let param_refs: Vec<&dyn rusqlite::ToSql> =
|
|
params_vec.iter().map(|p| p.as_ref()).collect();
|
|
let count: i64 = guard
|
|
.query_row(&sql, param_refs.as_slice(), |row| row.get(0))
|
|
.map_err(storage_err)?;
|
|
Ok(count)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// 按工具聚合执行统计(status 分布计数,AC-5 运行时失败率画像数据源)。
|
|
///
|
|
/// 单条 GROUP BY 取 `(tool_name, status, count)` 三元组;内存聚合与失败率口径
|
|
/// (failed_rate = failed / (completed + failed))在命令层完成(record.rs
|
|
/// `tool_failure_stats`),本层只负责取数,不掺展示逻辑。
|
|
/// `from`:可选时间下限(millis,`requested_at >= from`),None = 全量。
|
|
///
|
|
/// 与本表其他查询同理走专用 SELECT(通用 query 宏硬编码 ORDER BY created_at,
|
|
/// 本表无该列)。`requested_at` 存毫秒数字符串,`CAST AS INTEGER` 数值比较
|
|
/// (对齐 `cleanup_stale_pending` 同口径)。参数化绑定防注入。
|
|
pub async fn stats_by_tool(&self, from: Option<i64>) -> Result<Vec<(String, String, i64)>> {
|
|
let conn = self.conn.clone();
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
// 条件分支仅差 WHERE + 参数,一条 GROUP BY 复用(match 分支互斥,stmt 借用安全)
|
|
let (sql, param): (&str, Option<i64>) = match from {
|
|
Some(f) => (
|
|
"SELECT tool_name, status, COUNT(*) FROM ai_tool_executions \
|
|
WHERE CAST(requested_at AS INTEGER) >= ?1 GROUP BY tool_name, status",
|
|
Some(f),
|
|
),
|
|
None => (
|
|
"SELECT tool_name, status, COUNT(*) FROM ai_tool_executions \
|
|
GROUP BY tool_name, status",
|
|
None,
|
|
),
|
|
};
|
|
let row_map = |row: &rusqlite::Row| {
|
|
Ok((
|
|
row.get::<_, String>(0)?,
|
|
row.get::<_, String>(1)?,
|
|
row.get::<_, i64>(2)?,
|
|
))
|
|
};
|
|
let mut stmt = guard.prepare(sql).map_err(storage_err)?;
|
|
let mut rows = match param {
|
|
Some(f) => stmt.query_map(params![f], row_map).map_err(storage_err)?,
|
|
None => stmt.query_map([], row_map).map_err(storage_err)?,
|
|
};
|
|
let mut out = Vec::new();
|
|
for r in &mut rows {
|
|
out.push(r.map_err(storage_err)?);
|
|
}
|
|
Ok(out)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
}
|
|
|
|
// AiConversationRepo 的整体更新已由 impl_repo! 宏统一生成的 update_full 提供。
|
|
|
|
impl AiConversationRepo {
|
|
/// 写入对话版本化快照 checkpoint(INSERT OR IGNORE,同 id 已存在则跳过)。
|
|
///
|
|
/// 参数化绑定(替代原调用方的 format! 拼 SQL + execute_batch),防 snapshot 含引号/
|
|
/// 特殊字符致注入或损坏。列对齐 conversation_checkpoints(id, conv_id, snapshot,
|
|
/// token_total, created_at)。
|
|
pub async fn insert_checkpoint(
|
|
&self,
|
|
id: &str,
|
|
conv_id: &str,
|
|
snapshot: &str,
|
|
token_total: i64,
|
|
created_at: &str,
|
|
) -> Result<()> {
|
|
let conn = self.conn.clone();
|
|
let id = id.to_string();
|
|
let conv_id = conv_id.to_string();
|
|
let snapshot = snapshot.to_string();
|
|
let created_at = created_at.to_string();
|
|
let _ = tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
guard
|
|
.execute(
|
|
"INSERT OR IGNORE INTO conversation_checkpoints \
|
|
(id, conv_id, snapshot, token_total, created_at) \
|
|
VALUES (?1, ?2, ?3, ?4, ?5)",
|
|
params![id, conv_id, snapshot, token_total, created_at],
|
|
)
|
|
.map_err(storage_err)
|
|
})
|
|
.await
|
|
.map_err(storage_err)??;
|
|
Ok(())
|
|
}
|
|
|
|
/// 清空对话消息内容(保留 conversation 记录本身,只清 messages JSON + 清零 token 计数)
|
|
///
|
|
/// "清空对话"语义:对话壳保留(侧栏仍可见,可继续在该对话内聊),仅清空历史消息。
|
|
/// messages 是 ai_conversations 表内的 JSON 列而非独立行,故"删 messages"= 置空该列。
|
|
pub async fn clear_messages(&self, id: &str) -> Result<bool> {
|
|
let conn = self.conn.clone();
|
|
let id = id.to_owned();
|
|
let now = now_millis_str();
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
let affected = guard
|
|
.execute(
|
|
"UPDATE ai_conversations SET messages = '[]', prompt_tokens = 0, completion_tokens = 0, updated_at = ?1 WHERE id = ?2",
|
|
params![now, id],
|
|
)
|
|
.map_err(storage_err)?;
|
|
Ok(affected > 0)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// 清空对话消息内容(单事务原子:ai_conversations.messages 置 '[]' + ai_messages 表全删)。
|
|
///
|
|
/// A2-B9(G3.2 clearChat 裁决):原 `clear_messages` + `delete_range` 两条独立 DB 写非原子,
|
|
/// DB 失败会致 messages JSON 列与 ai_messages 表不一致(如仅一条成功)。本方法一次 transaction
|
|
/// 覆盖两条写(① UPDATE ai_conversations 置空消息 + 清零 token;② DELETE ai_messages 该 conv
|
|
/// 全部行),成功全成功 / 失败回滚全失败。供 `ai_chat_clear` 先停 loop 再单事务清空。
|
|
///
|
|
/// 对话壳保留(侧栏仍可见,可继续在该对话内聊);返回 Ok(())——调用方只关心成功与否
|
|
/// (对齐 replace_conversation 语义,不返回受影响行数)。
|
|
pub async fn clear_conversation_atomic(&self, id: &str) -> Result<()> {
|
|
let conn = self.conn.clone();
|
|
let id = id.to_owned();
|
|
let now = now_millis_str();
|
|
tokio::task::spawn_blocking(move || -> Result<()> {
|
|
let mut guard = conn.blocking_lock();
|
|
let tx = guard.transaction().map_err(storage_err)?;
|
|
{
|
|
// ① ai_conversations.messages 置空 + token 清零(对话壳保留)
|
|
tx.execute(
|
|
"UPDATE ai_conversations SET messages = '[]', prompt_tokens = 0, completion_tokens = 0, updated_at = ?1 WHERE id = ?2",
|
|
params![now, id],
|
|
)
|
|
.map_err(storage_err)?;
|
|
// ② ai_messages 表全删(等价 delete_range min_seq=0 max=None:seq 恒 >= 0)
|
|
tx.execute(
|
|
"DELETE FROM ai_messages WHERE conversation_id = ?1",
|
|
params![id],
|
|
)
|
|
.map_err(storage_err)?;
|
|
}
|
|
tx.commit().map_err(storage_err)?;
|
|
Ok(())
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// 设置归档标记(仅改 archived,不动 updated_at)
|
|
///
|
|
/// 区别于 update_field(后者强制 SET updated_at=now,会把归档/取消归档误判为内容更新,
|
|
/// 导致侧栏相对时间跳变为"刚刚")。归档是纯元数据标记,应保持时间不变。
|
|
pub async fn set_archived(&self, id: &str, archived: bool) -> Result<bool> {
|
|
let conn = self.conn.clone();
|
|
let id = id.to_owned();
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
let affected = guard
|
|
.execute(
|
|
"UPDATE ai_conversations SET archived = ?1 WHERE id = ?2",
|
|
params![if archived { 1 } else { 0 }, id],
|
|
)
|
|
.map_err(storage_err)?;
|
|
Ok(affected > 0)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// 设置标题(仅改 title,不动 updated_at)
|
|
///
|
|
/// 区别于 update_field(强制 SET updated_at=now,会把标题生成误判为内容更新,
|
|
/// 导致侧栏时间分组/排序跳变)。标题生成是系统后台操作,应保持会话相对时间不变。
|
|
pub async fn set_title(&self, id: &str, title: &str) -> Result<bool> {
|
|
let conn = self.conn.clone();
|
|
let id = id.to_owned();
|
|
let title = title.to_owned();
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
let affected = guard
|
|
.execute(
|
|
"UPDATE ai_conversations SET title = ?1 WHERE id = ?2",
|
|
params![title, id],
|
|
)
|
|
.map_err(storage_err)?;
|
|
Ok(affected > 0)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// 设置置顶标记(仅改 pinned,不动 updated_at) — UX-17
|
|
///
|
|
/// 同 set_archived:置顶是纯元数据标记,不应改变相对时间。前端排序读 pinned DESC, updated_at DESC。
|
|
pub async fn set_pinned(&self, id: &str, pinned: bool) -> Result<bool> {
|
|
let conn = self.conn.clone();
|
|
let id = id.to_owned();
|
|
tokio::task::spawn_blocking(move || {
|
|
let guard = conn.blocking_lock();
|
|
let affected = guard
|
|
.execute(
|
|
"UPDATE ai_conversations SET pinned = ?1 WHERE id = ?2",
|
|
params![if pinned { 1 } else { 0 }, id],
|
|
)
|
|
.map_err(storage_err)?;
|
|
Ok(affected > 0)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
|
|
/// G1.3: 删除对话 + 其全部 ai_messages 子行(单事务原子)。
|
|
///
|
|
/// 背景:原宏生成 `delete` 只删 ai_conversations 主行,而 ai_messages 表无外键级联
|
|
/// (conversation_id 仅普通索引),子行孤儿累积。本方法在同一事务内**先删子行
|
|
/// (ai_messages)再删主行(ai_conversations)**,要么全删要么全不删。
|
|
///
|
|
/// 顺序注意:先删数据再摘内存(命令层 per_conv.remove 在其后),防后台在途
|
|
/// save_conversation 在删主行后把孤儿消息写回复活。与 save_conversation 共享同一
|
|
/// conn(Mutex),事务原子性保证删除期间无中间态(半删半留)。
|
|
pub async fn delete_with_messages(&self, id: &str) -> Result<bool> {
|
|
let conn = self.conn.clone();
|
|
let id = id.to_owned();
|
|
tokio::task::spawn_blocking(move || {
|
|
let mut guard = conn.blocking_lock();
|
|
let tx = guard.transaction().map_err(storage_err)?;
|
|
// 先删子行(ai_messages)再删主行(ai_conversations),单事务原子
|
|
tx.execute(
|
|
"DELETE FROM ai_messages WHERE conversation_id = ?1",
|
|
params![id],
|
|
)
|
|
.map_err(storage_err)?;
|
|
let conv_affected = tx
|
|
.execute("DELETE FROM ai_conversations WHERE id = ?1", params![id])
|
|
.map_err(storage_err)?;
|
|
tx.commit().map_err(storage_err)?;
|
|
Ok(conv_affected > 0)
|
|
})
|
|
.await
|
|
.map_err(storage_err)?
|
|
}
|
|
}
|
|
|
|
// ============================================================
|
|
// 单元测试 — AiProviderRepo model_configs DB roundtrip + 老库兼容
|
|
// ============================================================
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use crate::db::Database;
|
|
use crate::models::AiProviderRecord;
|
|
use df_ai_core::model::{Capability, IntelligenceTier, Modality, ModelConfig};
|
|
|
|
/// model_configs DB roundtrip + 老库空兼容
|
|
#[tokio::test]
|
|
async fn ai_provider_model_configs_roundtrip_and_old_db_compat() {
|
|
let db = Database::open_in_memory().await.expect("open_in_memory");
|
|
let repo = AiProviderRepo::new(&db);
|
|
|
|
// 新格式:带多模型 + 多维度配置
|
|
let configs = vec![
|
|
ModelConfig::with_defaults("glm-4-flash"),
|
|
ModelConfig {
|
|
model_id: "glm-4v".into(),
|
|
modalities: vec![Modality::Text, Modality::Vision],
|
|
capabilities: vec![Capability::ToolUse],
|
|
intelligence: IntelligenceTier::Plus,
|
|
weight: 70,
|
|
context_window: 128_000,
|
|
..ModelConfig::with_defaults("glm-4v")
|
|
},
|
|
];
|
|
let rec = AiProviderRecord {
|
|
id: "p1".into(),
|
|
name: "测试".into(),
|
|
provider_type: "openai_compat".into(),
|
|
api_key: String::new(),
|
|
base_url: "https://x".into(),
|
|
default_model: "glm-4-flash".into(),
|
|
models: None,
|
|
model_configs: configs.clone(),
|
|
is_default: false,
|
|
config: None,
|
|
created_at: "0".into(),
|
|
updated_at: "0".into(),
|
|
enabled: true,
|
|
weight: 50,
|
|
};
|
|
repo.insert(rec).await.expect("insert");
|
|
|
|
let got = repo.get_by_id("p1").await.expect("get").expect("row exists");
|
|
assert_eq!(got.model_configs.len(), 2);
|
|
assert_eq!(got.model_configs[0].model_id, "glm-4-flash");
|
|
assert_eq!(got.model_configs[1].model_id, "glm-4v");
|
|
assert_eq!(got.model_configs[1].intelligence, IntelligenceTier::Plus);
|
|
assert_eq!(got.model_configs[1].context_window, 128_000);
|
|
|
|
// 老库空兼容:直接写 model_configs=NULL 的行(模拟 V18 之前的老库行)
|
|
// 然后 from_row 应得空 Vec
|
|
{
|
|
let conn = db.conn();
|
|
let g = conn.lock().await;
|
|
g.execute(
|
|
"INSERT OR REPLACE INTO ai_providers \
|
|
(id,name,provider_type,api_key,base_url,default_model,models,model_configs,is_default,config,created_at,updated_at) \
|
|
VALUES ('old','','openai_compat','','','','{}',NULL,0,NULL,'0','0')",
|
|
[],
|
|
)
|
|
.expect("raw insert old row");
|
|
}
|
|
let old = repo.get_by_id("old").await.expect("get").expect("old row");
|
|
assert!(old.model_configs.is_empty(), "NULL 列应得空 Vec");
|
|
|
|
// 老格式字符串数组 JSON(向后兼容 deserialize_model_configs)
|
|
{
|
|
let conn = db.conn();
|
|
let g = conn.lock().await;
|
|
g.execute(
|
|
"INSERT OR REPLACE INTO ai_providers \
|
|
(id,name,provider_type,api_key,base_url,default_model,models,model_configs,is_default,config,created_at,updated_at) \
|
|
VALUES ('legacy','','openai_compat','','','','{}','[\"glm-4-flash\",\"glm-4v\"]',0,NULL,'0','0')",
|
|
[],
|
|
)
|
|
.expect("raw insert legacy row");
|
|
}
|
|
let legacy = repo.get_by_id("legacy").await.expect("get").expect("legacy row");
|
|
assert_eq!(legacy.model_configs.len(), 2, "老字符串数组应转 2 个默认 ModelConfig");
|
|
assert_eq!(legacy.model_configs[0].model_id, "glm-4-flash");
|
|
}
|
|
}
|