mod.rs 加 pub fn err_str<E: ToString>,commands/ 10 文件 87 处 .map_err(|e| e.to_string()) → .map_err(err_str),残留 0(10 处复杂表达式 e 用于 format!/anyhow 保留)。行为零变化,统一入口为未来加日志/分类留点。批5,cargo workspace 0
564 lines
21 KiB
Rust
564 lines
21 KiB
Rust
//! 知识库集成 — 注入 + 提炼(嵌入生成 / 混合检索 / 上下文构建 / 对话提炼)
|
||
|
||
use std::sync::Arc;
|
||
|
||
use tokio::sync::Mutex;
|
||
|
||
use df_ai::provider::{ChatMessage, CompletionRequest, LlmProvider, MessageRole};
|
||
use df_storage::crud::{AiConversationRepo, KnowledgeRepo};
|
||
use df_storage::db::Database;
|
||
use df_storage::models::{AiProviderRecord, KnowledgeRecord};
|
||
|
||
use df_core::types::new_id;
|
||
|
||
use crate::state::{AppState, ExtractTrigger, LlmConcurrency};
|
||
|
||
use crate::commands::err_str;
|
||
|
||
use super::{AiSession};
|
||
|
||
/// 按配置构建 embedding provider + model。None = 配置缺失/provider 不存在。
|
||
async fn resolve_embed_provider(
|
||
state: &AppState,
|
||
config: &crate::state::KnowledgeConfig,
|
||
) -> Option<(Box<dyn LlmProvider>, String)> {
|
||
let id = config.embedding_provider_id.as_ref()?;
|
||
let model = config.embedding_model.clone().unwrap_or_else(|| "embedding-3".to_string());
|
||
let rec = match state.ai_providers.get_by_id(id).await {
|
||
Ok(Some(r)) => r,
|
||
_ => {
|
||
tracing::warn!("embedding provider 不存在: {}", id);
|
||
return None;
|
||
}
|
||
};
|
||
// build_provider_for 含空 key 早失败:Err → 返回 None 触发 LIKE 降级(与原 embed 失败降级行为一致)
|
||
let provider = match super::secret::build_provider_for(&rec) {
|
||
Ok(p) => p,
|
||
Err(e) => {
|
||
tracing::warn!("embedding provider 密钥不可用(降级 LIKE): {}", e);
|
||
return None;
|
||
}
|
||
};
|
||
Some((provider, model))
|
||
}
|
||
|
||
/// 生成文本嵌入(向量检索用)
|
||
///
|
||
/// 用配置指定的 embedding provider(必须 openai_compat 类型),失败返回 None(降级 LIKE)。
|
||
async fn generate_embedding(
|
||
state: &AppState,
|
||
text: &str,
|
||
config: &crate::state::KnowledgeConfig,
|
||
) -> Option<Vec<f32>> {
|
||
let (provider, model) = resolve_embed_provider(state, config).await?;
|
||
// 截断防超 token 上限(8192 token ≈ 8000 中文字)
|
||
let input: String = text.chars().take(8000).collect();
|
||
match provider.embed(&model, vec![input]).await {
|
||
Ok(mut vecs) if !vecs.is_empty() => Some(vecs.remove(0)),
|
||
Ok(_) => None,
|
||
Err(e) => {
|
||
tracing::warn!("embedding 生成失败(降级 LIKE): {}", e);
|
||
None
|
||
}
|
||
}
|
||
}
|
||
|
||
/// 知识条目发布时后台生成嵌入(fire-and-forget,失败仅 log)
|
||
///
|
||
/// 由 knowledge_update_status(发布路径)调用。vector_enabled 关闭时直接跳过。
|
||
pub async fn spawn_embedding_for_knowledge(
|
||
state: &AppState,
|
||
record: &df_storage::models::KnowledgeRecord,
|
||
) {
|
||
let config = state.knowledge_config.lock().await.clone();
|
||
if !config.vector_enabled {
|
||
return;
|
||
}
|
||
let Some((provider, model)) = resolve_embed_provider(state, &config).await else { return };
|
||
let text = format!("{} {}", record.title, record.content);
|
||
let id = record.id.clone();
|
||
let db = state.db.clone();
|
||
tauri::async_runtime::spawn(async move {
|
||
let input: String = text.chars().take(8000).collect();
|
||
match provider.embed(&model, vec![input]).await {
|
||
Ok(vecs) if !vecs.is_empty() => {
|
||
let repo = KnowledgeRepo::new(&db);
|
||
if let Err(e) = repo.set_embedding(&id, &vecs[0]).await {
|
||
tracing::warn!("嵌入写入失败(非阻断): {}", e);
|
||
} else {
|
||
tracing::info!("知识嵌入完成: {}", id);
|
||
}
|
||
}
|
||
Ok(_) => {}
|
||
Err(e) => tracing::warn!("知识嵌入生成失败(非阻断,走 LIKE 降级): {}", e),
|
||
}
|
||
});
|
||
}
|
||
|
||
/// 混合检索: LIKE 关键词 + 向量语义,合并去重加权
|
||
///
|
||
/// 双信号(两路都命中)排最前,LIKE 单信号次之,向量单信号第三。
|
||
/// vector_enabled 关闭或 embed 失败时纯 LIKE(零外部依赖降级)。
|
||
async fn hybrid_search(
|
||
state: &AppState,
|
||
query: &str,
|
||
limit: usize,
|
||
config: &crate::state::KnowledgeConfig,
|
||
) -> Vec<df_storage::models::KnowledgeRecord> {
|
||
let keyword_results = state.knowledge.search(query, None, limit).await.unwrap_or_default();
|
||
|
||
if !config.vector_enabled {
|
||
return keyword_results;
|
||
}
|
||
let query_vec = match generate_embedding(state, query, config).await {
|
||
Some(v) => v,
|
||
None => return keyword_results, // embed 失败降级
|
||
};
|
||
let vector_results = state.knowledge.search_vector(&query_vec, limit).await.unwrap_or_default();
|
||
|
||
merge_hybrid_results(keyword_results, vector_results, limit)
|
||
}
|
||
|
||
/// 混合检索三层合并去重(纯函数,抽自 hybrid_search)
|
||
///
|
||
/// 排序:双信号(LIKE + 向量均命中)> 仅 LIKE 单信号 > 仅向量单信号(且 cos≥0.3)。
|
||
/// cos<0.3 的向量单信号结果丢弃防噪音;limit 截断;按 id 去重。
|
||
pub(crate) fn merge_hybrid_results(
|
||
keyword_results: Vec<KnowledgeRecord>,
|
||
vector_results: Vec<(KnowledgeRecord, f32)>,
|
||
limit: usize,
|
||
) -> Vec<KnowledgeRecord> {
|
||
// 合并去重: 双信号 > LIKE 单信号 > 向量单信号(相似度<0.3 的向量结果丢弃防噪音)
|
||
let keyword_ids: std::collections::HashSet<String> = keyword_results.iter().map(|r| r.id.clone()).collect();
|
||
let mut merged = Vec::new();
|
||
let mut seen = std::collections::HashSet::new();
|
||
// 1. 双信号
|
||
for (rec, score) in &vector_results {
|
||
if keyword_ids.contains(&rec.id) && seen.insert(rec.id.clone()) {
|
||
tracing::debug!("混合检索双信号: {} (cos={:.2})", rec.title, score);
|
||
merged.push(rec.clone());
|
||
}
|
||
}
|
||
// 2. LIKE 单信号
|
||
for rec in &keyword_results {
|
||
if seen.insert(rec.id.clone()) {
|
||
merged.push(rec.clone());
|
||
}
|
||
}
|
||
// 3. 向量单信号(过滤低相似度)
|
||
for (rec, score) in &vector_results {
|
||
if *score >= 0.3 && seen.insert(rec.id.clone()) {
|
||
merged.push(rec.clone());
|
||
}
|
||
}
|
||
merged.truncate(limit);
|
||
merged
|
||
}
|
||
|
||
/// 构建知识库上下文片段,拼入 Chat system prompt
|
||
///
|
||
/// 流程: 开关检查 → 混合检索 top-3(克制) → 命中条目 reuse_count +1 + 记录引用事件(fire-and-forget) → markdown 格式化
|
||
/// 关闭时返回空串(零开销);无结果返回空串。
|
||
pub(crate) async fn build_knowledge_context(
|
||
state: &AppState,
|
||
conv_id: &str,
|
||
query: &str,
|
||
config: &crate::state::KnowledgeConfig,
|
||
) -> String {
|
||
if !config.auto_inject {
|
||
return String::new();
|
||
}
|
||
let results = hybrid_search(state, query, 3, config).await;
|
||
if results.is_empty() {
|
||
return String::new();
|
||
}
|
||
// 命中条目:复用计数 +1 + 记录引用事件(fire-and-forget,单个 spawn 任务批量处理)
|
||
let db = state.db.clone();
|
||
let ids: Vec<String> = results.iter().map(|r| r.id.clone()).collect();
|
||
let conv_id = conv_id.to_string();
|
||
let query_clone = query.to_string();
|
||
tauri::async_runtime::spawn(async move {
|
||
let repo = KnowledgeRepo::new(&db);
|
||
let timeline = crate::commands::knowledge_timeline::KnowledgeTimeline::new(&db);
|
||
for id in &ids {
|
||
if let Err(e) = repo.increment_reuse_count(id).await {
|
||
tracing::warn!("reuse_count +1 失败(非阻断): {}", e);
|
||
}
|
||
timeline.record_referenced(id, &conv_id, &query_clone).await;
|
||
}
|
||
});
|
||
let mut out = String::from("## 相关知识库\n");
|
||
for r in &results {
|
||
let kind = r.kind.clone();
|
||
let title = r.title.clone();
|
||
// 截断 content 防膨胀(注入侧最多 500 字符)
|
||
let snippet: String = r.content.chars().take(500).collect();
|
||
out.push_str(&format!("- [{}] {}: {} (复用 {} 次)\n", kind, title, snippet, r.reuse_count));
|
||
}
|
||
out
|
||
}
|
||
|
||
/// 判断是否应触发提炼,满足则后台 spawn 提炼 task(非阻断)
|
||
///
|
||
/// 守卫: auto_extract 开 + trigger_mode == OnComplete + 消息数 ≥ min_messages
|
||
pub(crate) async fn maybe_spawn_extraction(
|
||
session_arc: &Arc<Mutex<AiSession>>,
|
||
db: &Arc<Database>,
|
||
conv_id: &str,
|
||
provider_config: &AiProviderRecord,
|
||
config: &crate::state::KnowledgeConfig,
|
||
llm_concurrency: LlmConcurrency,
|
||
) -> anyhow::Result<()> {
|
||
if !config.auto_extract {
|
||
return Ok(());
|
||
}
|
||
if config.trigger_mode != ExtractTrigger::OnComplete {
|
||
return Ok(());
|
||
}
|
||
// 消息数守卫(总消息数,含 system/assistant/tool)
|
||
let msg_count = session_arc.lock().await.messages.len();
|
||
if (msg_count as u32) < config.min_messages {
|
||
return Ok(());
|
||
}
|
||
|
||
let db = db.clone();
|
||
let conv_id = conv_id.to_string();
|
||
let provider_config = provider_config.clone();
|
||
let llm_concurrency = llm_concurrency.clone();
|
||
tauri::async_runtime::spawn(async move {
|
||
if let Err(e) = extract_knowledge_from_conversation(&db, &conv_id, &provider_config, &llm_concurrency).await {
|
||
tracing::warn!("知识提取失败(非阻断): {}", e);
|
||
}
|
||
});
|
||
Ok(())
|
||
}
|
||
|
||
/// 手动触发提炼(ManualOnly 模式 / 前端按钮调用)
|
||
///
|
||
/// fire-and-forget:立即返回,后台执行 LLM 提炼(避免 IPC 长时间阻塞)。
|
||
pub async fn trigger_extraction_now(state: &AppState) -> Result<bool, String> {
|
||
let conv_id = {
|
||
let session = state.ai_session.lock().await;
|
||
session.active_conversation_id.clone()
|
||
};
|
||
let conv_id = conv_id.ok_or_else(|| "当前无活跃对话".to_string())?;
|
||
let provider_config = super::prompt::get_active_provider(state).await.map_err(err_str)?;
|
||
let db = state.db.clone();
|
||
let llm_concurrency = state.llm_concurrency.clone();
|
||
tauri::async_runtime::spawn(async move {
|
||
if let Err(e) = extract_knowledge_from_conversation(&db, &conv_id, &provider_config, &llm_concurrency).await {
|
||
tracing::warn!("手动提炼失败(非阻断): {}", e);
|
||
}
|
||
});
|
||
Ok(true)
|
||
}
|
||
|
||
/// 知识提炼提示词 — 强制 JSON 输出,含矛盾知识约束
|
||
const EXTRACTION_SYSTEM_PROMPT: &str = "你是知识提炼引擎,从 AI 对话中识别可复用的经验。\
|
||
只提取真正通用、可被未来对话复用的知识,过滤一次性闲聊/项目特定的临时内容。\n\n\
|
||
输出严格的 JSON 数组(不要 markdown 代码块包裹),每个元素 schema:\n\
|
||
{\"kind\": \"review_rule|prompt_template|pitfall|architecture_pattern|diagnosis|deployment_note|workflow_optimization\", \
|
||
\"title\": \"简短标题\", \"content\": \"完整可复用内容\", \
|
||
\"tags\": [\"标签\"], \"confidence\": \"high|medium|low\", \"reasoning\": \"为何值得沉淀\"}\n\n\
|
||
规则:\n\
|
||
1. 如果适用范围有限制(如仅适用特定语言/框架/场景),必须在 content 或 tags 中明确标注\n\
|
||
2. confidence: high=对话中可直接观察的明确模式, medium=合理推断, low=推测性弱信号\n\
|
||
3. 无可提炼内容时返回空数组 []\n\
|
||
4. 输出纯 JSON,无任何额外文字";
|
||
|
||
/// 从对话中提炼知识,产出 candidate 写入知识库
|
||
///
|
||
/// 流程: 读对话消息 → 过滤 user/assistant 取最后 6 条 → LLM 提炼(强制 JSON) → 解析 → 批量插入 candidate
|
||
async fn extract_knowledge_from_conversation(
|
||
db: &Arc<Database>,
|
||
conv_id: &str,
|
||
provider_config: &AiProviderRecord,
|
||
llm_concurrency: &LlmConcurrency,
|
||
) -> anyhow::Result<()> {
|
||
let conv_repo = AiConversationRepo::new(db);
|
||
let conv = conv_repo
|
||
.get_by_id(conv_id)
|
||
.await?
|
||
.ok_or_else(|| anyhow::anyhow!("对话不存在: {}", conv_id))?;
|
||
// 对话标题(生命线溯源用,空标题降级为占位,避免字节切片风险)
|
||
let conv_title = conv
|
||
.title
|
||
.clone()
|
||
.filter(|t| !t.trim().is_empty())
|
||
.unwrap_or_else(|| "未命名对话".to_string());
|
||
|
||
let messages: Vec<ChatMessage> = serde_json::from_str(&conv.messages).unwrap_or_default();
|
||
// 过滤 user/assistant,取最后 6 条
|
||
let recent: Vec<&ChatMessage> = messages
|
||
.iter()
|
||
.filter(|m| matches!(m.role, MessageRole::User | MessageRole::Assistant))
|
||
.rev()
|
||
.take(6)
|
||
.collect();
|
||
if recent.len() < 4 {
|
||
return Ok(()); // 太短,不值得提炼
|
||
}
|
||
|
||
// 构造提炼消息: system 指令 + 对话内容(user 角色)
|
||
let mut conv_text = String::new();
|
||
for m in recent.iter().rev() {
|
||
let role = match m.role {
|
||
MessageRole::User => "用户",
|
||
MessageRole::Assistant => "助手",
|
||
_ => continue,
|
||
};
|
||
conv_text.push_str(&format!("[{}]: {}\n\n", role, m.content));
|
||
}
|
||
|
||
let extract_messages = vec![
|
||
ChatMessage::system(EXTRACTION_SYSTEM_PROMPT),
|
||
ChatMessage::user(&format!("请从以下对话中提炼可复用知识:\n\n{}", conv_text)),
|
||
];
|
||
let request = CompletionRequest {
|
||
model: provider_config.default_model.clone(),
|
||
messages: extract_messages,
|
||
temperature: Some(0.3),
|
||
max_tokens: Some(2048),
|
||
stream: false,
|
||
tools: None,
|
||
tool_choice: None,
|
||
};
|
||
|
||
let provider: Box<dyn LlmProvider> = match super::secret::build_provider_for(provider_config) {
|
||
Ok(p) => p,
|
||
// 空 key 早失败:上层(maybe_spawn_extraction/trigger_extraction_now)已 warn log 降级,语义一致
|
||
Err(e) => return Err(anyhow::anyhow!("provider 密钥不可用: {}", e)),
|
||
};
|
||
// LLM 并发限流(知识提炼属独立调用,纳入双层 Semaphore)
|
||
let _global_permit = llm_concurrency.acquire_global().await;
|
||
let _per_conv_permit = llm_concurrency.acquire_per_conv().await;
|
||
let resp = provider.complete(request).await?;
|
||
let raw = resp.text.trim();
|
||
|
||
// 容错:剥离可能的 ```json ... ``` 包裹
|
||
let json_str = strip_code_fence(raw);
|
||
let items: Vec<serde_json::Value> = match serde_json::from_str(json_str) {
|
||
Ok(v) => v,
|
||
Err(e) => {
|
||
tracing::warn!("知识提炼 JSON 解析失败,整批丢弃(非阻断): {} | 原始: {}", e, raw);
|
||
return Ok(());
|
||
}
|
||
};
|
||
|
||
let knowledge_repo = KnowledgeRepo::new(db);
|
||
let timeline = crate::commands::knowledge_timeline::KnowledgeTimeline::new(db);
|
||
let mut inserted = 0;
|
||
for item in &items {
|
||
let kind = match item.get("kind").and_then(|v| v.as_str()) {
|
||
Some(k) => k.to_string(),
|
||
None => continue,
|
||
};
|
||
let title = item.get("title").and_then(|v| v.as_str()).unwrap_or("").to_string();
|
||
let content = item.get("content").and_then(|v| v.as_str()).unwrap_or("").to_string();
|
||
if title.is_empty() || content.is_empty() {
|
||
continue;
|
||
}
|
||
let tags = item.get("tags").map(|v| serde_json::to_string(v).unwrap_or_else(|_| "[]".into()));
|
||
let confidence = item.get("confidence").and_then(|v| v.as_str()).map(|s| s.to_string());
|
||
// 回填 AI 判断依据(prompt 要求的 reasoning 字段,此前被丢弃)
|
||
let reasoning = item.get("reasoning").and_then(|v| v.as_str()).map(|s| s.to_string());
|
||
|
||
let now = crate::commands::now_millis();
|
||
let record = KnowledgeRecord {
|
||
id: new_id(),
|
||
kind,
|
||
title,
|
||
content,
|
||
tags,
|
||
status: "candidate".to_string(),
|
||
confidence,
|
||
reuse_count: 0,
|
||
verified: false,
|
||
source_project: None,
|
||
source_ref: Some(format!("conv:{}", conv_id)),
|
||
reasoning: reasoning.clone(),
|
||
created_at: now.clone(),
|
||
updated_at: now,
|
||
};
|
||
match knowledge_repo.insert(record.clone()).await {
|
||
Ok(_) => {
|
||
inserted += 1;
|
||
tracing::info!(
|
||
"AI 提炼知识候选: {} [confidence={}]",
|
||
record.title,
|
||
record.confidence.as_deref().unwrap_or("?")
|
||
);
|
||
// 生命线:AI 提炼产生(fire-and-forget)
|
||
timeline
|
||
.record_extracted(
|
||
&record.id,
|
||
conv_id,
|
||
&conv_title,
|
||
reasoning.as_deref().unwrap_or(""),
|
||
)
|
||
.await;
|
||
}
|
||
Err(e) => tracing::warn!("知识候选插入失败(非阻断): {}", e),
|
||
}
|
||
}
|
||
if inserted > 0 {
|
||
tracing::info!("知识提炼完成: 对话 {} 产出 {} 条 candidate", conv_id, inserted);
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
/// 剥离 LLM 输出可能的 ```json ... ``` 代码块包裹
|
||
fn strip_code_fence(s: &str) -> &str {
|
||
let s = s.trim();
|
||
if let Some(rest) = s.strip_prefix("```json") {
|
||
return rest.trim().trim_end_matches("```").trim();
|
||
}
|
||
if let Some(rest) = s.strip_prefix("```") {
|
||
return rest.trim().trim_end_matches("```").trim();
|
||
}
|
||
s
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
use df_storage::models::KnowledgeRecord;
|
||
|
||
// ---------- merge_hybrid_results ----------
|
||
|
||
fn kr(id: &str, title: &str) -> KnowledgeRecord {
|
||
KnowledgeRecord {
|
||
id: id.to_string(),
|
||
kind: "snippet".to_string(),
|
||
title: title.to_string(),
|
||
content: String::new(),
|
||
tags: None,
|
||
status: "published".to_string(),
|
||
confidence: None,
|
||
reuse_count: 0,
|
||
verified: false,
|
||
source_project: None,
|
||
source_ref: None,
|
||
reasoning: None,
|
||
created_at: "2026-01-01".to_string(),
|
||
updated_at: "2026-01-01".to_string(),
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn merge_empty_inputs() {
|
||
let out = merge_hybrid_results(vec![], vec![], 5);
|
||
assert!(out.is_empty());
|
||
}
|
||
|
||
#[test]
|
||
fn merge_dual_signal_ranks_first() {
|
||
// r1 同时命中双信号 → 应排在首位
|
||
let kw = vec![kr("r1", "kw1"), kr("r2", "kw2")];
|
||
let vec_results = vec![(kr("r1", "kw1-vec"), 0.8)];
|
||
let out = merge_hybrid_results(kw, vec_results, 5);
|
||
assert_eq!(out.len(), 2);
|
||
assert_eq!(out[0].id, "r1", "双信号 r1 必须排首");
|
||
assert_eq!(out[1].id, "r2");
|
||
}
|
||
|
||
#[test]
|
||
fn merge_keyword_only_after_dual() {
|
||
// r2 仅 LIKE,应在双信号之后
|
||
let kw = vec![kr("only-kw", "kw-only")];
|
||
let vec_results = vec![(kr("dual", "dual-vec"), 0.7)];
|
||
// dual 不在 kw 集合 → 非双信号,走向量单信号(0.7≥0.3)
|
||
let out = merge_hybrid_results(kw, vec_results, 5);
|
||
// 顺序:无双信号 → LIKE 单信号(only-kw)→ 向量单信号(dual)
|
||
assert_eq!(out.len(), 2);
|
||
assert_eq!(out[0].id, "only-kw");
|
||
assert_eq!(out[1].id, "dual");
|
||
}
|
||
|
||
#[test]
|
||
fn merge_vector_threshold_filters_below_03() {
|
||
// cos=0.29 < 0.3 → 向量单信号结果被滤掉
|
||
let kw = vec![];
|
||
let vec_results = vec![(kr("low", "low-vec"), 0.29)];
|
||
let out = merge_hybrid_results(kw, vec_results, 5);
|
||
assert!(out.is_empty(), "cos=0.29 应被过滤");
|
||
}
|
||
|
||
#[test]
|
||
fn merge_vector_threshold_keeps_at_031() {
|
||
// cos=0.31 ≥ 0.3 → 保留
|
||
let kw = vec![];
|
||
let vec_results = vec![(kr("ok", "ok-vec"), 0.31)];
|
||
let out = merge_hybrid_results(kw, vec_results, 5);
|
||
assert_eq!(out.len(), 1);
|
||
assert_eq!(out[0].id, "ok");
|
||
}
|
||
|
||
#[test]
|
||
fn merge_vector_threshold_boundary_exact_03() {
|
||
// 边界:cos 恰好 0.3 → 保留(>= 比较)
|
||
let kw = vec![];
|
||
let vec_results = vec![(kr("edge", "edge-vec"), 0.3)];
|
||
let out = merge_hybrid_results(kw, vec_results, 5);
|
||
assert_eq!(out.len(), 1, "cos=0.3 边界应保留(>= 比较)");
|
||
assert_eq!(out[0].id, "edge");
|
||
}
|
||
|
||
#[test]
|
||
fn merge_truncates_to_limit() {
|
||
// limit 截断
|
||
let kw: Vec<KnowledgeRecord> = (0..10).map(|i| kr(&format!("k{i}"), "t")).collect();
|
||
let out = merge_hybrid_results(kw, vec![], 3);
|
||
assert_eq!(out.len(), 3);
|
||
}
|
||
|
||
#[test]
|
||
fn merge_dedups_across_signals() {
|
||
// 同一 id 多路命中只出现一次(双信号路径优先)
|
||
let kw = vec![kr("dup", "dup-kw")];
|
||
let vec_results = vec![(kr("dup", "dup-vec"), 0.9), (kr("v2", "v2-vec"), 0.5)];
|
||
let out = merge_hybrid_results(kw, vec_results, 5);
|
||
assert_eq!(out.len(), 2, "dup 去重只出现一次");
|
||
assert_eq!(out[0].id, "dup", "dup 双信号排首");
|
||
assert_eq!(out[1].id, "v2");
|
||
}
|
||
|
||
#[test]
|
||
fn merge_dual_signal_not_duplicated_in_keyword_pass() {
|
||
// 双信号记录已被 seen 标记,LIKE 单信号遍历时不会重复入列
|
||
let kw = vec![kr("both", "both-kw"), kr("kwonly", "kwo")];
|
||
let vec_results = vec![(kr("both", "both-vec"), 0.6)];
|
||
let out = merge_hybrid_results(kw, vec_results, 5);
|
||
let both_count = out.iter().filter(|r| r.id == "both").count();
|
||
assert_eq!(both_count, 1);
|
||
assert_eq!(out.len(), 2);
|
||
}
|
||
|
||
#[test]
|
||
fn merge_all_three_signal_types_present() {
|
||
// 三类信号齐全:dual(双)+ kw-only(LIKE)+ vec-only(向量)
|
||
let kw = vec![kr("dual", "d-kw"), kr("kwonly", "k-kw")];
|
||
let vec_results = vec![
|
||
(kr("dual", "d-vec"), 0.85),
|
||
(kr("veconly", "v-vec"), 0.45),
|
||
];
|
||
let out = merge_hybrid_results(kw, vec_results, 10);
|
||
assert_eq!(out.len(), 3);
|
||
// 排序:dual → kwonly → veconly
|
||
assert_eq!(out[0].id, "dual");
|
||
assert_eq!(out[1].id, "kwonly");
|
||
assert_eq!(out[2].id, "veconly");
|
||
}
|
||
|
||
#[test]
|
||
fn merge_limit_truncates_after_sorting() {
|
||
// 截断发生在排序之后:limit=1 时即便有双信号也只留首条
|
||
let kw = vec![kr("kw1", "k1")];
|
||
let vec_results = vec![(kr("dual", "dv"), 0.9)];
|
||
// dual 不在 kw,故无双信号;顺序 kw1 → dual
|
||
let out = merge_hybrid_results(kw, vec_results, 1);
|
||
assert_eq!(out.len(), 1);
|
||
assert_eq!(out[0].id, "kw1");
|
||
}
|
||
}
|