Files
DevFlow/src-tauri/src/commands/ai/knowledge_inject.rs
绝尘 892a642bd4 优化: R-PD-10 commands err_str helper 统一错误格式化(87 处)
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
2026-06-15 05:45:04 +08:00

564 lines
21 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 知识库集成 — 注入 + 提炼(嵌入生成 / 混合检索 / 上下文构建 / 对话提炼)
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");
}
}