重构: crud.rs按表拆分(SMELL-P1-9) + F-09批1 PerConvState数据结构

- SMELL-P1-9: crud.rs 2212行→crud/6文件(mod/settings/project_repo/task_repo/conversation_repo/idea_repo)
  re-export pub use *_repo::* 零调用方改动,宏 pub(crate) use + 子模块 use super::impl_repo
  基线测试锁12表白名单+13Repo构造
- F-09 批1: PerConvState struct(9字段对齐AiSession::new)+AiSession.per_conv HashMap+conv()/conv_read()访问器+3单测
  纯新增无行为变化,b-1主代自主裁决采纳,批2迁移承接
主代统一兜底: cargo check --workspace 0 + df-storage 35+11 + devflow 96 passed
This commit is contained in:
2026-06-19 01:34:59 +08:00
parent d817386418
commit 2c8764abee
11 changed files with 2592 additions and 2214 deletions

View File

@@ -0,0 +1,871 @@
//! 想法/知识域 Repo:IdeaRepo / KnowledgeRepo / KnowledgeEventsRepo + 向量工具(embedding BLOB 序列化 + 余弦相似度)
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::{IdeaRecord, KnowledgeEventRecord, KnowledgeRecord};
use super::impl_repo;
use super::{now_millis_str, storage_err, validate_column_name};
// ============================================================
// 知识库 SELECT 列清单(防 COLS 漂移)
// ============================================================
/// `knowledges` 表对应 `KnowledgeRecord` 14 个字段的列名(顺序与结构体一致)。
///
/// 多处 `search`/`search_vector` 内联 COLS 串的 DRY 收口(CR-260615-03):集中一处定义,
/// 配合下方 `KNOWLEDGE_COL_COUNT` 断言,任一处加列漏改会被测试 `test_knowledge_cols_matches_record`
/// 立即捕获(`knowledge_from_row` 按 name 取列,SELECT 漏列会运行时 rusqlite 报错,故提前断言)。
const KNOWLEDGE_COLS: &str = "id,kind,title,content,tags,status,confidence,reuse_count,verified,source_project,source_ref,reasoning,created_at,updated_at";
/// `KnowledgeRecord` 字段数(与上面列清单的逗号分隔项数一致,被测试断言)。
/// 仅测试期消费(列漂移断言);保留为非 `cfg(test)` 以便测试外的阅读者一眼看到字段数。
#[cfg_attr(not(test), allow(dead_code))]
const KNOWLEDGE_COL_COUNT: usize = 14;
/// `search_vector` 用的列清单:KNOWLEDGE_COLS + embedding(余弦计算用,不入 KnowledgeRecord)。
const KNOWLEDGE_COLS_WITH_EMBEDDING: &str = concat!(
"id,kind,title,content,tags,status,confidence,reuse_count,verified,",
"source_project,source_ref,reasoning,created_at,updated_at,embedding"
);
// ============================================================
// 向量工具 — embedding BLOB 序列化 + 余弦相似度
// ============================================================
/// Vec<f32> → 小端字节 BLOB
fn f32s_to_blob(v: &[f32]) -> Vec<u8> {
v.iter().flat_map(|f| f.to_le_bytes()).collect()
}
/// BLOB → Vec<f32>(长度非 4 倍数时截断尾部残字节)
fn blob_to_f32s(blob: &[u8]) -> Vec<f32> {
blob.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
/// 余弦相似度(确定性数学,与任何实现结果一致)
fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
dot / (norm_a * norm_b + 1e-8)
}
// ============================================================
// from_row 辅助函数
// ============================================================
fn idea_from_row(row: &Row<'_>) -> std::result::Result<IdeaRecord, rusqlite::Error> {
Ok(IdeaRecord {
id: row.get("id")?,
title: row.get("title")?,
description: row.get("description")?,
status: row.get("status")?,
priority: row.get("priority")?,
score: row.get("score")?,
tags: row.get("tags")?,
source: row.get("source")?,
promoted_to: row.get("promoted_to")?,
ai_analysis: row.get("ai_analysis")?,
scores: row.get("scores")?,
created_at: row.get("created_at")?,
updated_at: row.get("updated_at")?,
})
}
fn knowledge_from_row(row: &Row<'_>) -> std::result::Result<KnowledgeRecord, rusqlite::Error> {
Ok(KnowledgeRecord {
id: row.get("id")?,
kind: row.get("kind")?,
title: row.get("title")?,
content: row.get("content")?,
tags: row.get("tags")?,
status: row.get("status")?,
confidence: row.get("confidence")?,
reuse_count: row.get("reuse_count")?,
verified: row.get::<_, i32>("verified")? != 0,
source_project: row.get("source_project")?,
source_ref: row.get("source_ref")?,
reasoning: row.get("reasoning")?,
created_at: row.get("created_at")?,
updated_at: row.get("updated_at")?,
})
}
fn knowledge_event_from_row(row: &Row<'_>) -> std::result::Result<KnowledgeEventRecord, rusqlite::Error> {
Ok(KnowledgeEventRecord {
id: row.get("id")?,
knowledge_id: row.get("knowledge_id")?,
event_type: row.get("event_type")?,
source_ref: row.get("source_ref")?,
context_json: row.get("context_json")?,
timestamp: row.get("timestamp")?,
})
}
// ============================================================
// Repo 实现
// ============================================================
impl_repo!(
/// 想法表 CRUD
IdeaRepo,
IdeaRecord,
"ideas",
from_row => |row| idea_from_row(row),
insert => |conn, rec| {
conn.execute(
"INSERT INTO ideas (id, title, description, status, priority, score, tags, source, promoted_to, ai_analysis, scores, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13)",
params![
rec.id, rec.title, rec.description, rec.status, rec.priority,
rec.score, rec.tags, rec.source, rec.promoted_to, rec.ai_analysis,
rec.scores, rec.created_at, rec.updated_at
],
)
},
update => |conn, rec| {
conn.execute(
"UPDATE ideas SET title = ?1, description = ?2, status = ?3, priority = ?4, score = ?5, tags = ?6, source = ?7, promoted_to = ?8, ai_analysis = ?9, scores = ?10, updated_at = ?11 WHERE id = ?12",
params![
rec.title, rec.description, rec.status, rec.priority,
rec.score, rec.tags, rec.source, rec.promoted_to, rec.ai_analysis,
rec.scores, rec.updated_at, rec.id
],
)
}
);
impl_repo!(
/// 知识库表 CRUD
KnowledgeRepo,
KnowledgeRecord,
"knowledges",
from_row => |row| knowledge_from_row(row),
insert => |conn, rec| {
let verified = if rec.verified { 1i32 } else { 0i32 };
conn.execute(
"INSERT INTO knowledges (id, kind, title, content, tags, status, confidence, reuse_count, verified, source_project, source_ref, reasoning, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)",
params![
rec.id, rec.kind, rec.title, rec.content, rec.tags, rec.status, rec.confidence,
rec.reuse_count, verified, rec.source_project, rec.source_ref, rec.reasoning,
rec.created_at, rec.updated_at
],
)
},
update => |conn, rec| {
let verified = if rec.verified { 1i32 } else { 0i32 };
conn.execute(
"UPDATE knowledges SET kind = ?1, title = ?2, content = ?3, tags = ?4, status = ?5, confidence = ?6, reuse_count = ?7, verified = ?8, source_project = ?9, source_ref = ?10, reasoning = ?11, updated_at = ?12 WHERE id = ?13",
params![
rec.kind, rec.title, rec.content, rec.tags, rec.status, rec.confidence,
rec.reuse_count, verified, rec.source_project, rec.source_ref, rec.reasoning,
rec.updated_at, rec.id
],
)
}
);
// KnowledgeRepo 的整体更新已由 impl_repo! 宏统一生成的 update_full 提供。
impl KnowledgeRepo {
/// 检索知识: title/content LIKE 匹配,可选 kind 过滤,按 reuse_count 降序,top-N
///
/// 克制检索: 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 pattern = format!("%{}%", query);
let kind = kind.map(|s| s.to_owned());
let limit_i = limit as i64;
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
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 OR content LIKE ?2) 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))
.map_err(storage_err)?;
for r in rows {
results.push(r.map_err(storage_err)?);
}
} else {
let mut stmt = guard
.prepare(&format!("SELECT {KNOWLEDGE_COLS} FROM knowledges WHERE status = 'published' AND (title LIKE ?1 OR content LIKE ?2) 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))
.map_err(storage_err)?;
for r in rows {
results.push(r.map_err(storage_err)?);
}
}
Ok(results)
})
.await
.map_err(storage_err)?
}
/// 按状态列出(审核收件箱用): 按 confidence 语义排序(high>medium>low),次按 created_at
///
/// 用 CASE WHEN 替代纯 TEXT 排序(字典序 high>low>medium 非预期语义)。
pub async fn list_by_status(&self, status: &str) -> Result<Vec<KnowledgeRecord>> {
let conn = self.conn.clone();
let status = status.to_owned();
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
let mut stmt = guard
.prepare(
"SELECT id,kind,title,content,tags,status,confidence,reuse_count,verified,\
source_project,source_ref,reasoning,created_at,updated_at \
FROM knowledges WHERE status = ?1
ORDER BY CASE confidence
WHEN 'high' THEN 3
WHEN 'medium' THEN 2
WHEN 'low' THEN 1
ELSE 0
END DESC, created_at DESC",
)
.map_err(storage_err)?;
let rows = stmt
.query_map(params![status], |row| knowledge_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)?
}
/// 复用计数 +1(SQL 行级原子操作,并发安全)
pub async fn increment_reuse_count(&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 knowledges SET reuse_count = reuse_count + 1, updated_at = ?1 WHERE id = ?2",
params![now, id],
)
.map_err(storage_err)?;
Ok(affected > 0)
})
.await
.map_err(storage_err)?
}
/// 写入向量嵌入(BLOB = Vec<f32> 小端字节序列化)
///
/// embedding 列不进 KnowledgeRecord(IPC 不需要传向量给前端),专用方法读写。
pub async fn set_embedding(&self, id: &str, embedding: &[f32]) -> Result<bool> {
let conn = self.conn.clone();
let id = id.to_owned();
let blob = f32s_to_blob(embedding);
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
let affected = guard
.execute(
"UPDATE knowledges SET embedding = ?1 WHERE id = ?2",
params![blob, id],
)
.map_err(storage_err)?;
Ok(affected > 0)
})
.await
.map_err(storage_err)?
}
/// 向量检索: 加载全部 published 且有 embedding 的记录,纯 Rust 余弦相似度取 top-N
///
/// 返回 (记录, 相似度分数)。数据量 <50k 时暴力遍历 <50ms,够用;
/// 更大规模再升 sqlite-vec HNSW(结果不变,只提速)。
///
/// 列限定: 显式列出所需列(与 search/list_by_status/top_used 一致),避免 SELECT *
/// 拉到未知新增列;embedding 单独取(不入 KnowledgeRecord)。
///
/// TODO(性能,低优先): 调用方(hybrid_search→merge_hybrid_results→build_knowledge_context)
/// 实际只消费 id/kind/title/content/reuse_count;reasoning(AI 生成大文本)、tags、
/// source_project/source_ref 等元字段未被使用却仍随每行读出。真正省 IO 需返回精简结构
/// (如 KnowledgeVectorHit { id, kind, title, content, reuse_count })替换返回类型,
/// 但这会改变 search_vector 签名与 merge_hybrid_results 调用契约——当前保守不动,
/// 待向量检索量级或 reasoning 文本体积成为瓶颈再单独立项。SELECT 列化本身不省字段,
/// 仅消除 SELECT * 的隐式依赖与未知列风险。
pub async fn search_vector(&self, query_vec: &[f32], limit: usize) -> Result<Vec<(KnowledgeRecord, f32)>> {
let conn = self.conn.clone();
let query_vec = query_vec.to_vec();
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
// 显式列: 14 个 KnowledgeRecord 字段 + embedding(余弦计算用,不入 KnowledgeRecord)
let mut stmt = guard
.prepare(&format!(
"SELECT {KNOWLEDGE_COLS_WITH_EMBEDDING} FROM knowledges WHERE status = 'published' AND embedding IS NOT NULL"
))
.map_err(storage_err)?;
let rows = stmt
.query_map([], |row| {
let rec = knowledge_from_row(row)?;
let blob: Vec<u8> = row.get("embedding")?;
Ok((rec, blob))
})
.map_err(storage_err)?;
let mut scored: Vec<(KnowledgeRecord, f32)> = Vec::new();
for r in rows {
let (rec, blob) = r.map_err(storage_err)?;
let emb = blob_to_f32s(&blob);
// 维度不匹配(换过 embedding 模型的旧向量)跳过
if emb.len() != query_vec.len() {
continue;
}
let score = cosine_similarity(&query_vec, &emb);
scored.push((rec, score));
}
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(limit);
Ok(scored)
})
.await
.map_err(storage_err)?
}
/// 列出非归档知识(全部 status != 'archived'),按 confidence 语义排序
pub async fn list_non_archived(&self) -> Result<Vec<KnowledgeRecord>> {
let conn = self.conn.clone();
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
let mut stmt = guard
.prepare(
"SELECT id,kind,title,content,tags,status,confidence,reuse_count,verified,\
source_project,source_ref,reasoning,created_at,updated_at \
FROM knowledges WHERE status != 'archived'
ORDER BY CASE confidence
WHEN 'high' THEN 3
WHEN 'medium' THEN 2
WHEN 'low' THEN 1
ELSE 0
END DESC, created_at DESC",
)
.map_err(storage_err)?;
let rows = stmt
.query_map([], |row| knowledge_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)?
}
/// 热门知识(已发布,按复用次数降序)
pub async fn top_used(&self, limit: usize) -> Result<Vec<KnowledgeRecord>> {
let conn = self.conn.clone();
let limit_i = limit as i64;
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
let mut stmt = guard
.prepare("SELECT id,kind,title,content,tags,status,confidence,reuse_count,verified,source_project,source_ref,reasoning,created_at,updated_at FROM knowledges WHERE status = 'published' ORDER BY reuse_count DESC LIMIT ?1")
.map_err(storage_err)?;
let rows = stmt
.query_map(params![limit_i], |row| knowledge_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)?
}
}
impl_repo!(
/// 知识生命线事件表 CRUD(追加型审计表)
KnowledgeEventsRepo,
KnowledgeEventRecord,
"knowledge_events",
from_row => |row| knowledge_event_from_row(row),
insert => |conn, rec| {
conn.execute(
"INSERT INTO knowledge_events (id, knowledge_id, event_type, source_ref, context_json, timestamp)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
params![rec.id, rec.knowledge_id, rec.event_type, rec.source_ref, rec.context_json, rec.timestamp],
)
},
update => |conn, rec| {
conn.execute(
"UPDATE knowledge_events SET knowledge_id = ?1, event_type = ?2, source_ref = ?3, context_json = ?4, timestamp = ?5 WHERE id = ?6",
params![rec.knowledge_id, rec.event_type, rec.source_ref, rec.context_json, rec.timestamp, rec.id],
)
}
);
impl KnowledgeEventsRepo {
/// 按知识 ID 查询全部事件(时间正序,构建生命线视图用)
pub async fn list_by_knowledge(&self, knowledge_id: &str) -> Result<Vec<KnowledgeEventRecord>> {
let conn = self.conn.clone();
let knowledge_id = knowledge_id.to_owned();
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
let mut stmt = guard
.prepare("SELECT id,knowledge_id,event_type,source_ref,context_json,timestamp FROM knowledge_events WHERE knowledge_id = ?1 ORDER BY timestamp ASC")
.map_err(storage_err)?;
let rows = stmt
.query_map(params![knowledge_id], |row| knowledge_event_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)?
}
/// 按知识 ID + event_type 查询最近 N 条(如引用记录翻页: type=referenced)
pub async fn list_by_knowledge_type(
&self,
knowledge_id: &str,
event_type: &str,
limit: usize,
) -> Result<Vec<KnowledgeEventRecord>> {
let conn = self.conn.clone();
let knowledge_id = knowledge_id.to_owned();
let event_type = event_type.to_owned();
let limit_i = limit as i64;
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
let mut stmt = guard
.prepare("SELECT id,knowledge_id,event_type,source_ref,context_json,timestamp FROM knowledge_events WHERE knowledge_id = ?1 AND event_type = ?2 ORDER BY timestamp DESC LIMIT ?3")
.map_err(storage_err)?;
let rows = stmt
.query_map(params![knowledge_id, event_type, limit_i], |row| knowledge_event_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)?
}
/// 跨知识列最近 N 条事件(全表 timestamp DESC,top-N)。
///
/// **专用兜底方法**:本表时间列名是 `timestamp` 而非 `created_at`,但
/// `impl_repo!` 宏生成的 `query()` / `list_all()` 硬编码 `ORDER BY created_at`
/// (见宏内 `ORDER BY created_at DESC` 字面量),误调 `state.knowledge_events.query(...)`
/// 会触发 SQLite "no such column: created_at"。本方法走专用 SELECT 绕过宏硬编码,
/// 供需要跨知识按时间倒序浏览事件的调用方使用(对标 AiToolExecutionRepo::list_recent
/// 对 ai_tool_executions 表的同款兜底处理——那张表同样无 created_at 列)。
/// limit 上限钳制 200,防前端恶意/失误传超大值。
pub async fn list_recent(&self, limit: u32) -> Result<Vec<KnowledgeEventRecord>> {
let conn = self.conn.clone();
// 钳制 limit 防滥用(最大 200)
let safe_limit = limit.min(200) as i64;
tokio::task::spawn_blocking(move || {
let guard = conn.blocking_lock();
let mut stmt = guard
.prepare(
"SELECT id,knowledge_id,event_type,source_ref,context_json,timestamp \
FROM knowledge_events ORDER BY timestamp DESC LIMIT ?1",
)
.map_err(storage_err)?;
let rows = stmt
.query_map(params![safe_limit], |row| knowledge_event_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)?
}
}
// ============================================================
// 单元测试 — 向量纯函数 + KnowledgeRepo 内存 DB + COLS 漂移防护
// ============================================================
#[cfg(test)]
mod tests {
use super::*;
use crate::db::Database;
// ---------- COLS 漂移防护(CR-260615-03) ----------
/// KNOWLEDGE_COLS 列数须等于 KNOWLEDGE_COL_COUNT(任一处漂移:加列漏改 / 串错位 → 立即失败)。
/// `knowledge_from_row` 按 name 取列,SELECT 漏列会在运行时被 rusqlite 报错;此断言提前到测试期捕获。
#[test]
fn test_knowledge_cols_matches_record() {
let count = KNOWLEDGE_COLS.split(',').count();
assert_eq!(
count, KNOWLEDGE_COL_COUNT,
"KNOWLEDGE_COLS({count}列) ≠ KNOWLEDGE_COL_COUNT({KNOWLEDGE_COL_COUNT}); \
修改一处须同步另一处"
);
// search_vector 多一列 embedding
let count_with_emb = KNOWLEDGE_COLS_WITH_EMBEDDING.split(',').count();
assert_eq!(
count_with_emb,
KNOWLEDGE_COL_COUNT + 1,
"KNOWLEDGE_COLS_WITH_EMBEDDING({count_with_emb}列) ≠ KNOWLEDGE_COL_COUNT+1({}); \
embedding 列应单独追加",
KNOWLEDGE_COL_COUNT + 1
);
// 每个列名须能被 split 出来(防末尾多逗号 / 空段)
for col in KNOWLEDGE_COLS.split(',') {
assert!(!col.is_empty(), "KNOWLEDGE_COLS 含空列名段");
}
}
// ---------- 向量纯函数 ----------
#[test]
fn f32s_blob_roundtrip_nonempty() {
let v = vec![0.0, 1.5, -2.25, 3.14159, -0.0001];
let blob = f32s_to_blob(&v);
assert_eq!(blob.len(), v.len() * 4);
let back = blob_to_f32s(&blob);
assert_eq!(back, v);
}
#[test]
fn f32s_blob_roundtrip_empty() {
let v: Vec<f32> = vec![];
let blob = f32s_to_blob(&v);
assert!(blob.is_empty());
assert!(blob_to_f32s(&blob).is_empty());
}
#[test]
fn blob_to_f32s_drops_trailing_partial_bytes() {
// 1 个完整 f32 + 3 残字节 → chunks_exact 截断尾部
let blob = f32s_to_blob(&[42.0]);
let with_garbage: Vec<u8> = blob.into_iter().chain([0xff, 0xff, 0xff]).collect();
assert_eq!(blob_to_f32s(&with_garbage), vec![42.0]);
}
#[test]
fn cosine_identical_vectors_near_one() {
let a = vec![1.0, 2.0, 3.0];
let sim = cosine_similarity(&a, &a);
assert!((sim - 1.0).abs() < 1e-5, "identical ≈ 1.0, got {}", sim);
}
#[test]
fn cosine_orthogonal_vectors_near_zero() {
let a = vec![1.0, 0.0];
let b = vec![0.0, 1.0];
let sim = cosine_similarity(&a, &b);
assert!(sim.abs() < 1e-5, "orthogonal ≈ 0.0, got {}", sim);
}
#[test]
fn cosine_opposite_vectors_near_neg_one() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![-1.0, -2.0, -3.0];
let sim = cosine_similarity(&a, &b);
assert!((sim + 1.0).abs() < 1e-5, "opposite ≈ -1.0, got {}", sim);
}
#[test]
fn cosine_dimension_mismatch_zips_to_shortest() {
// 实现用 zip(a,b),维度不匹配按较短的那个对齐,不 panic
let a = vec![1.0, 0.0, 0.0]; // 3 维
let b = vec![1.0, 0.0]; // 2 维 → zip 取前 2 个
let sim = cosine_similarity(&a, &b);
assert!((sim - 1.0).abs() < 1e-5, "zip 截断后应 ≈ 1.0, got {}", sim);
}
#[test]
fn cosine_zero_vector_does_not_panic() {
// 分母有 +1e-8 保护,零向量返回有限值而非 NaN/inf
let z = vec![0.0, 0.0, 0.0];
let sim = cosine_similarity(&z, &z);
assert!(sim.is_finite(), "零向量相似度应有限, got {}", sim);
}
// ---------- KnowledgeRepo 内存 DB ----------
/// 构造一条 KnowledgeRecord fixture
fn krec(
id: &str,
title: &str,
content: &str,
kind: &str,
status: &str,
confidence: Option<&str>,
reuse: i32,
) -> KnowledgeRecord {
KnowledgeRecord {
id: id.to_string(),
kind: kind.to_string(),
title: title.to_string(),
content: content.to_string(),
tags: Some("[]".to_string()),
status: status.to_string(),
confidence: confidence.map(|s| s.to_string()),
reuse_count: reuse,
verified: false,
source_project: Some("proj-1".to_string()),
source_ref: Some("conv:c1".to_string()),
reasoning: None,
created_at: "1700000000000".to_string(),
updated_at: "1700000000000".to_string(),
}
}
async fn setup_repo() -> KnowledgeRepo {
let db = Database::open_in_memory().await.expect("open_in_memory");
KnowledgeRepo::new(&db)
}
#[tokio::test]
async fn search_matches_title_and_filters_non_published() {
let repo = setup_repo().await;
// 一条命中(title 含关键词)、一条 published 不命中、一条 status 非 published 但命中
repo.insert(krec("k1", "Rust 异步并发模型", "tokio 运行时", "lesson", "published", Some("high"), 5))
.await
.unwrap();
repo.insert(krec("k2", "无关标题", "无关内容", "lesson", "published", None, 0))
.await
.unwrap();
repo.insert(krec("k3", "Rust 异步进阶", "...", "lesson", "candidate", Some("low"), 9))
.await
.unwrap();
let hits = repo.search("异步", None, 10).await.unwrap();
// 只有 k1 是 published 且命中;k3 命中但非 published 被过滤
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].id, "k1");
}
#[tokio::test]
async fn search_matches_content_branch() {
let repo = setup_repo().await;
// title 不含关键词,content 含 → 命中 content LIKE 分支
repo.insert(krec("k1", "标题", "深入理解 Rust 所有权与借用", "lesson", "published", None, 1))
.await
.unwrap();
let hits = repo.search("所有权", None, 10).await.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].id, "k1");
}
#[tokio::test]
async fn search_no_keyword_match_returns_empty() {
let repo = setup_repo().await;
repo.insert(krec("k1", "Rust", "tokio", "lesson", "published", None, 1))
.await
.unwrap();
let hits = repo.search("不存在的关键词xyz", None, 10).await.unwrap();
assert!(hits.is_empty());
}
#[tokio::test]
async fn search_kind_filter_narrows_results() {
let repo = setup_repo().await;
repo.insert(krec("k1", "Rust 规范", "...", "lesson", "published", None, 1))
.await
.unwrap();
repo.insert(krec("k2", "Rust 决策", "...", "decision", "published", None, 1))
.await
.unwrap();
// 不过滤 kind → 两条都命中
let all = repo.search("Rust", None, 10).await.unwrap();
assert_eq!(all.len(), 2);
// 只取 lesson → 仅 k1
let only_lesson = repo.search("Rust", Some("lesson"), 10).await.unwrap();
assert_eq!(only_lesson.len(), 1);
assert_eq!(only_lesson[0].id, "k1");
}
#[tokio::test]
async fn search_orders_by_reuse_count_desc() {
let repo = setup_repo().await;
repo.insert(krec("low", "Rust A", "...", "lesson", "published", None, 1))
.await
.unwrap();
repo.insert(krec("high", "Rust B", "...", "lesson", "published", None, 50))
.await
.unwrap();
repo.insert(krec("mid", "Rust C", "...", "lesson", "published", None, 10))
.await
.unwrap();
let hits = repo.search("Rust", None, 10).await.unwrap();
let ids: Vec<_> = hits.iter().map(|h| h.id.as_str()).collect();
assert_eq!(ids, vec!["high", "mid", "low"]);
}
#[tokio::test]
async fn search_truncates_to_limit() {
let repo = setup_repo().await;
for i in 0..5 {
repo.insert(krec(&format!("k{i}"), &format!("Rust-{i}"), "...", "lesson", "published", None, i))
.await
.unwrap();
}
let hits = repo.search("Rust", None, 2).await.unwrap();
assert_eq!(hits.len(), 2);
// reuse_count 最高的两条(k4=4, k3=3)
assert_eq!(hits[0].id, "k4");
assert_eq!(hits[1].id, "k3");
}
#[tokio::test]
async fn list_by_status_orders_by_confidence_semantics() {
// confidence 字典序 high>low>medium 非预期,实现用 CASE 强制 high>medium>low
let repo = setup_repo().await;
repo.insert(krec("low", "t", "c", "lesson", "pending_review", Some("low"), 0))
.await
.unwrap();
repo.insert(krec("high", "t", "c", "lesson", "pending_review", Some("high"), 0))
.await
.unwrap();
repo.insert(krec("medium", "t", "c", "lesson", "pending_review", Some("medium"), 0))
.await
.unwrap();
// created_at 相同,纯靠 confidence 排序
let list = repo.list_by_status("pending_review").await.unwrap();
let confidences: Vec<_> = list.iter().map(|r| r.confidence.as_deref().unwrap_or("")).collect();
assert_eq!(confidences, vec!["high", "medium", "low"]);
}
#[tokio::test]
async fn list_by_status_filters_other_status() {
let repo = setup_repo().await;
repo.insert(krec("a", "t", "c", "lesson", "pending_review", Some("high"), 0))
.await
.unwrap();
repo.insert(krec("b", "t", "c", "lesson", "published", Some("high"), 0))
.await
.unwrap();
let list = repo.list_by_status("pending_review").await.unwrap();
assert_eq!(list.len(), 1);
assert_eq!(list[0].id, "a");
}
#[tokio::test]
async fn increment_reuse_count_is_atomic_plus_one() {
let repo = setup_repo().await;
repo.insert(krec("k1", "Rust", "...", "lesson", "published", None, 0))
.await
.unwrap();
// 连续 +1 两次
assert!(repo.increment_reuse_count("k1").await.unwrap());
assert!(repo.increment_reuse_count("k1").await.unwrap());
let rec = repo.get_by_id("k1").await.unwrap().expect("记录存在");
assert_eq!(rec.reuse_count, 2);
// 不存在的 id → false
assert!(!repo.increment_reuse_count("nope").await.unwrap());
}
#[tokio::test]
async fn search_vector_orders_by_cosine_and_truncates() {
let repo = setup_repo().await;
// 三条 published 记录,embedding 维度均为 2
repo.insert(krec("exact", "e", "c", "lesson", "published", None, 0))
.await
.unwrap();
repo.insert(krec("orth", "e", "c", "lesson", "published", None, 0))
.await
.unwrap();
repo.insert(krec("neg", "e", "c", "lesson", "published", None, 0))
.await
.unwrap();
// 一条非 published(应被过滤)
repo.insert(krec("hidden", "e", "c", "lesson", "candidate", None, 0))
.await
.unwrap();
repo.set_embedding("exact", &[1.0, 0.0]).await.unwrap(); // 与 query 完全同向
repo.set_embedding("orth", &[0.0, 1.0]).await.unwrap(); // 正交 ≈ 0
repo.set_embedding("neg", &[-1.0, 0.0]).await.unwrap(); // 反向 ≈ -1
repo.set_embedding("hidden", &[1.0, 0.0]).await.unwrap(); // 非 published 过滤掉
let results = repo.search_vector(&[1.0, 0.0], 10).await.unwrap();
// hidden 被 status 过滤,剩 3 条
assert_eq!(results.len(), 3);
// 按相似度降序: exact(≈1) > orth(≈0) > neg(≈-1)
assert_eq!(results[0].0.id, "exact");
assert_eq!(results[1].0.id, "orth");
assert_eq!(results[2].0.id, "neg");
assert!((results[0].1 - 1.0).abs() < 1e-5);
assert!(results[1].1.abs() < 1e-5);
assert!((results[2].1 + 1.0).abs() < 1e-5);
}
#[tokio::test]
async fn search_vector_skips_dimension_mismatch() {
let repo = setup_repo().await;
repo.insert(krec("dim2", "e", "c", "lesson", "published", None, 0))
.await
.unwrap();
repo.insert(krec("dim3", "e", "c", "lesson", "published", None, 0))
.await
.unwrap();
repo.set_embedding("dim2", &[1.0, 0.0]).await.unwrap();
repo.set_embedding("dim3", &[1.0, 0.0, 0.0]).await.unwrap(); // 维度不匹配
// query 是 2 维,dim3 被跳过
let results = repo.search_vector(&[1.0, 0.0], 10).await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0.id, "dim2");
}
#[tokio::test]
async fn search_vector_truncates_to_limit() {
let repo = setup_repo().await;
for i in 0..4 {
repo.insert(krec(&format!("k{i}"), "e", "c", "lesson", "published", None, 0))
.await
.unwrap();
// 全部与 query 同向,相似度相同 → 仅验证 truncate 生效
repo.set_embedding(&format!("k{i}"), &[1.0, 0.0]).await.unwrap();
}
let results = repo.search_vector(&[1.0, 0.0], 2).await.unwrap();
assert_eq!(results.len(), 2);
}
#[tokio::test]
async fn search_vector_ignores_records_without_embedding() {
let repo = setup_repo().await;
// published 但未写 embedding
repo.insert(krec("noemb", "e", "c", "lesson", "published", None, 0))
.await
.unwrap();
let results = repo.search_vector(&[1.0, 0.0], 10).await.unwrap();
assert!(results.is_empty());
}
}