新增: Phase2 阶段收尾(Sprint 1-20)
重构:删 5 零引用 crate(df-evolve/plugin/stages/task/traceability)+ 清死模块、ai.rs 拆 11 子 module、ai.ts 拆 6 composable、i18n 拆目录 功能:知识库全栈(df-project/scan + CRUD + 时间线 + 前端)、Settings 拆分、appSettings KV 迁移、模型池、LLM 并发 Semaphore 修复:审批持久化根治、ConditionEngine 默认拒绝、NodeRegistry unimplemented 清除、promote 补偿删除、工具结果截断 50KB、路径校验防 symlink 逃逸 文档:B-03 人工审批设计、决策记录三分档、规格契约自检、经验记录、todo 看板、PROGRESS 更新 详见 PROGRESS.md。src-tauri/儿童每日打卡应用/ 与本项目无关,已排除。
This commit is contained in:
271
src-tauri/src/commands/ai/agentic.rs
Normal file
271
src-tauri/src/commands/ai/agentic.rs
Normal file
@@ -0,0 +1,271 @@
|
||||
//! Agentic 循环 — 流式接收 → 工具执行 → 结果回传 LLM → 循环
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::Ordering;
|
||||
|
||||
use tauri::{AppHandle, Emitter};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use df_ai::ai_tools::AiToolRegistry;
|
||||
use df_ai::context::TokenEstimator;
|
||||
use df_ai::provider::{ChatMessage, CompletionRequest, LlmProvider};
|
||||
|
||||
use df_storage::db::Database;
|
||||
use df_storage::models::AiProviderRecord;
|
||||
|
||||
use crate::state::{AppState, LlmConcurrency};
|
||||
|
||||
use super::conversation::{save_conversation, TokenAccumulator};
|
||||
use super::knowledge_inject::maybe_spawn_extraction;
|
||||
use super::prompt::{build_system_prompt, get_active_provider};
|
||||
use super::stream_recv::stream_llm;
|
||||
use super::title::{ensure_conversation_title, spawn_ensure_title};
|
||||
use super::audit::process_tool_calls;
|
||||
|
||||
use super::{AiChatEvent, AiSession};
|
||||
|
||||
/// Agentic 循环最大迭代次数
|
||||
pub(crate) const MAX_AGENT_ITERATIONS: usize = 10;
|
||||
|
||||
/// Agentic 循环:流式接收 → 工具执行 → 结果回传 LLM → 循环
|
||||
///
|
||||
/// 退出条件:
|
||||
/// - LLM 只返回文本(无 tool_calls)→ 正常结束
|
||||
/// - 有工具需要审批 → 暂停循环(generating 保持 true),等 ai_approve 恢复
|
||||
/// - 达到最大迭代次数 → 正常结束
|
||||
pub(crate) async fn run_agentic_loop(
|
||||
session_arc: Arc<Mutex<AiSession>>,
|
||||
tools_arc: Arc<AiToolRegistry>,
|
||||
db: Arc<Database>,
|
||||
app_handle: AppHandle,
|
||||
provider_config: AiProviderRecord,
|
||||
system_prompt: String,
|
||||
conv_id: String,
|
||||
knowledge_config: crate::state::KnowledgeConfig,
|
||||
llm_concurrency: LlmConcurrency,
|
||||
) {
|
||||
let provider: Box<dyn LlmProvider> = df_ai::build_provider(
|
||||
&provider_config.provider_type,
|
||||
&provider_config.base_url,
|
||||
&provider_config.api_key,
|
||||
&provider_config.default_model,
|
||||
);
|
||||
let tool_defs = tools_arc.tool_definitions();
|
||||
// 停止信号副本:stream_llm 与每轮迭代共享读取,避免重复加锁
|
||||
let stop_flag = session_arc.lock().await.stop_flag.clone();
|
||||
|
||||
// token 累加器:loop 生命周期内各轮叠加,退出时传 save_conversation(累加模式落库)
|
||||
let mut tokens = TokenAccumulator::default();
|
||||
|
||||
for iteration in 0..MAX_AGENT_ITERATIONS {
|
||||
// 用户请求停止 → 收尾退出(已生成文本已在上一轮入库)
|
||||
if stop_flag.load(Ordering::SeqCst) {
|
||||
let usage = df_ai::provider::TokenUsage {
|
||||
prompt_tokens: tokens.prompt(),
|
||||
completion_tokens: tokens.completion(),
|
||||
total_tokens: tokens.total(),
|
||||
};
|
||||
// 入口 stop:本轮可能尚未 stream(首轮即停),不记 model——避免把未实际生成的 model 写入 models 数组
|
||||
save_conversation(&session_arc, &db, &conv_id, Some(&usage), None).await;
|
||||
// 标题生成后台化:不阻塞 Completed emit(失败有 extract_title 兜底)
|
||||
spawn_ensure_title(&provider_config, &db, &conv_id, &app_handle, &session_arc, &llm_concurrency);
|
||||
let mut session = session_arc.lock().await;
|
||||
session.generating = false;
|
||||
// generating 复位后再 emit Completed:保证前端收事件时后端已可接下一条(发送队列续发不被"正在生成中"拒绝)
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiCompleted { total_tokens: usage.total_tokens, prompt_tokens: tokens.prompt(), completion_tokens: tokens.completion(), conversation_id: Some(conv_id.clone()) });
|
||||
return;
|
||||
}
|
||||
|
||||
// 新一轮通知前端(第二轮起),前端需新建 assistant 消息
|
||||
if iteration > 0 {
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiAgentRound {
|
||||
round: (iteration + 1) as u32,
|
||||
conversation_id: Some(conv_id.clone()),
|
||||
});
|
||||
}
|
||||
|
||||
// 构建请求消息(超预算时自动裁剪旧消息,保护工具调用三元组 + 最近 6 条)
|
||||
let messages = {
|
||||
let session = session_arc.lock().await;
|
||||
let sys_tokens = TokenEstimator::default().estimate_text(&system_prompt);
|
||||
let (history_msgs, _trimmed) = session.messages.build_for_request(sys_tokens);
|
||||
let mut msgs = vec![ChatMessage::system(&system_prompt)];
|
||||
msgs.extend(history_msgs);
|
||||
msgs
|
||||
};
|
||||
|
||||
// 预估输入 token(兜底:部分 provider 如 GLM 流式 usage 不报 prompt_tokens,后段用它补)
|
||||
let estimated_prompt: u32 = {
|
||||
let est = TokenEstimator::default();
|
||||
messages.iter().map(|m| est.estimate_message(m)).sum()
|
||||
};
|
||||
|
||||
let request = CompletionRequest {
|
||||
model: provider_config.default_model.clone(),
|
||||
messages,
|
||||
temperature: Some(0.7),
|
||||
max_tokens: Some(8192),
|
||||
stream: true,
|
||||
tools: if tool_defs.is_empty() { None } else { Some(tool_defs.clone()) },
|
||||
tool_choice: None,
|
||||
};
|
||||
|
||||
// LLM 并发限流(全局 + 单对话双层),仅覆盖 stream_llm 调用本身;
|
||||
// 工具执行(process_tool_calls)是本地操作无 RPM 成本,permit 在 stream 后立即释放避免占槽
|
||||
let _global_permit = llm_concurrency.acquire_global().await;
|
||||
let _per_conv_permit = llm_concurrency.acquire_per_conv().await;
|
||||
// 流式接收(内部处理 idle timeout / 断连检测 / 停止信号)
|
||||
let (full_text, tool_calls_acc, round_usage) = match stream_llm(&*provider, request, &app_handle, &stop_flag, &conv_id).await {
|
||||
Some(result) => result,
|
||||
None => {
|
||||
// 错误已在 stream_llm 中 emit,直接结束
|
||||
let mut session = session_arc.lock().await;
|
||||
session.generating = false;
|
||||
return;
|
||||
}
|
||||
};
|
||||
// stream 结束立即释放 permit,后续工具执行不受限流(本地操作无 RPM 成本)
|
||||
drop(_global_permit);
|
||||
drop(_per_conv_permit);
|
||||
|
||||
// 累加本轮 token:provider 流式 usage 的 prompt_tokens 为 0 时(GLM 等),用预估输入兜底
|
||||
let round_prompt = if round_usage.prompt_tokens == 0 { estimated_prompt } else { round_usage.prompt_tokens };
|
||||
tokens.add(round_prompt, round_usage.completion_tokens);
|
||||
|
||||
// 追加 assistant 消息到历史
|
||||
let has_tool_calls = !tool_calls_acc.is_empty();
|
||||
{
|
||||
let mut session = session_arc.lock().await;
|
||||
if has_tool_calls {
|
||||
let mut order: Vec<u32> = tool_calls_acc.keys().copied().collect();
|
||||
order.sort_unstable();
|
||||
let ai_tool_calls: Vec<df_ai::provider::ToolCall> = order.iter()
|
||||
.map(|i| {
|
||||
let draft = &tool_calls_acc[i];
|
||||
df_ai::provider::ToolCall::new(&draft.id, &draft.name, &draft.args)
|
||||
})
|
||||
.collect();
|
||||
let mut msg = ChatMessage::assistant_with_tools(&full_text, ai_tool_calls);
|
||||
msg.model = Some(provider_config.default_model.clone());
|
||||
session.messages.push(msg);
|
||||
} else if !full_text.is_empty() {
|
||||
let mut msg = ChatMessage::assistant(&full_text);
|
||||
msg.model = Some(provider_config.default_model.clone());
|
||||
session.messages.push(msg);
|
||||
}
|
||||
}
|
||||
|
||||
// 停止信号:已生成文本入库后退出,不再执行后续工具调用
|
||||
if stop_flag.load(Ordering::SeqCst) {
|
||||
let usage = df_ai::provider::TokenUsage {
|
||||
prompt_tokens: tokens.prompt(),
|
||||
completion_tokens: tokens.completion(),
|
||||
total_tokens: tokens.total(),
|
||||
};
|
||||
save_conversation(&session_arc, &db, &conv_id, Some(&usage), Some(&provider_config.default_model)).await;
|
||||
// 标题生成后台化:不阻塞 Completed emit(失败有 extract_title 兜底)
|
||||
spawn_ensure_title(&provider_config, &db, &conv_id, &app_handle, &session_arc, &llm_concurrency);
|
||||
let mut session = session_arc.lock().await;
|
||||
session.generating = false;
|
||||
// generating 复位后再 emit Completed:保证前端收事件时后端已可接下一条(发送队列续发不被"正在生成中"拒绝)
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiCompleted { total_tokens: usage.total_tokens, prompt_tokens: tokens.prompt(), completion_tokens: tokens.completion(), conversation_id: Some(conv_id.clone()) });
|
||||
return;
|
||||
}
|
||||
|
||||
// 无工具调用 → 最终文本响应,循环结束
|
||||
if !has_tool_calls { break; }
|
||||
|
||||
// 处理工具调用(Low 自动执行 / Medium+High 待审批)
|
||||
let pending_count = {
|
||||
let mut session = session_arc.lock().await;
|
||||
process_tool_calls(&mut session, tool_calls_acc, &tools_arc, &db, &app_handle, &conv_id).await
|
||||
};
|
||||
|
||||
// 有待审批 → 暂停循环,等待用户审批后通过 ai_approve → try_continue_agent_loop 恢复
|
||||
if pending_count > 0 {
|
||||
let usage = df_ai::provider::TokenUsage {
|
||||
prompt_tokens: tokens.prompt(),
|
||||
completion_tokens: tokens.completion(),
|
||||
total_tokens: tokens.total(),
|
||||
};
|
||||
save_conversation(&session_arc, &db, &conv_id, Some(&usage), Some(&provider_config.default_model)).await;
|
||||
return; // generating 保持 true
|
||||
}
|
||||
|
||||
// 全部自动执行完成 → 继续下一轮
|
||||
}
|
||||
|
||||
// 正常完成
|
||||
let usage = df_ai::provider::TokenUsage {
|
||||
prompt_tokens: tokens.prompt(),
|
||||
completion_tokens: tokens.completion(),
|
||||
total_tokens: tokens.total(),
|
||||
};
|
||||
// 落库 + 标题 + 知识提炼打包后台化:不阻塞 generating 复位与 Completed 事件
|
||||
// save 先行(extract/title 都读已落库消息);extract 内部 fire-and-forget,与 title 可能并发
|
||||
// (均受 per_conv 信号量约束,读写不同字段互不干扰)
|
||||
// 并发取舍:与新对话新 loop 的 save 存在低概率并发 upsert,最多丢少量 token 累加(非功能错误,可接受)
|
||||
let usage_total = usage.total_tokens;
|
||||
{
|
||||
let session_arc = session_arc.clone();
|
||||
let db = db.clone();
|
||||
let conv_id = conv_id.clone();
|
||||
let provider_config = provider_config.clone();
|
||||
let knowledge_config = knowledge_config.clone();
|
||||
let app_handle = app_handle.clone();
|
||||
let llm_concurrency = llm_concurrency.clone();
|
||||
tauri::async_runtime::spawn(async move {
|
||||
save_conversation(&session_arc, &db, &conv_id, Some(&usage), Some(&provider_config.default_model)).await;
|
||||
// 知识提炼:需读已落库的对话消息,故在 save 之后
|
||||
if let Err(e) = maybe_spawn_extraction(&session_arc, &db, &conv_id, &provider_config, &knowledge_config, llm_concurrency.clone()).await {
|
||||
tracing::warn!("知识提炼触发失败(非阻断): {}", e);
|
||||
}
|
||||
ensure_conversation_title(&provider_config, &db, &conv_id, &app_handle, &session_arc, llm_concurrency).await;
|
||||
});
|
||||
}
|
||||
|
||||
let mut session = session_arc.lock().await;
|
||||
session.generating = false;
|
||||
// generating 复位后再 emit Completed:落库/标题/提炼已在后台,前端立即感知完成
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiCompleted { total_tokens: usage_total, prompt_tokens: tokens.prompt(), completion_tokens: tokens.completion(), conversation_id: Some(conv_id.clone()) });
|
||||
}
|
||||
|
||||
/// 检查是否所有待审批已处理,如果是则恢复 agentic 循环
|
||||
pub(crate) async fn try_continue_agent_loop(app: &AppHandle, state: &AppState) {
|
||||
let should_continue = {
|
||||
let session = state.ai_session.lock().await;
|
||||
session.generating && session.pending_approvals.is_empty()
|
||||
};
|
||||
|
||||
if !should_continue { return; }
|
||||
|
||||
let provider_config = match get_active_provider(state).await {
|
||||
Ok(p) => p,
|
||||
Err(_) => return,
|
||||
};
|
||||
let (lang, conv_id) = {
|
||||
let session = state.ai_session.lock().await;
|
||||
let lang = session.agent_language.clone().unwrap_or_else(|| "zh-CN".to_string());
|
||||
let conv_id = session.active_conversation_id.clone().unwrap_or_default();
|
||||
(lang, conv_id)
|
||||
};
|
||||
let system_prompt = build_system_prompt(state, &lang).await;
|
||||
|
||||
let session_arc = state.ai_session.clone();
|
||||
let tools_arc = state.ai_tools.clone();
|
||||
let db = state.db.clone();
|
||||
let app_handle = app.clone();
|
||||
let knowledge_config = state.knowledge_config.lock().await.clone();
|
||||
let llm_concurrency = state.llm_concurrency.clone();
|
||||
|
||||
// 恢复循环前通知前端新建 assistant 消息:审批(通过/拒绝)后新一轮文本
|
||||
// 不应追加到发起工具调用的旧消息,用 AiAgentRound 隔开
|
||||
let _ = app.emit("ai-chat-event", AiChatEvent::AiAgentRound {
|
||||
round: 0,
|
||||
conversation_id: Some(conv_id.clone()),
|
||||
});
|
||||
|
||||
tauri::async_runtime::spawn(async move {
|
||||
run_agentic_loop(session_arc, tools_arc, db, app_handle, provider_config, system_prompt, conv_id, knowledge_config, llm_concurrency).await;
|
||||
});
|
||||
}
|
||||
238
src-tauri/src/commands/ai/audit.rs
Normal file
238
src-tauri/src/commands/ai/audit.rs
Normal file
@@ -0,0 +1,238 @@
|
||||
//! 工具调用审计 + pending 审批恢复 + 工具调用处理
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use tauri::{AppHandle, Emitter};
|
||||
|
||||
use df_ai::ai_tools::{AiToolRegistry, RiskLevel};
|
||||
use df_ai::provider::ChatMessage;
|
||||
use df_storage::crud::AiToolExecutionRepo;
|
||||
use df_storage::db::Database;
|
||||
use df_storage::models::AiToolExecutionRecord;
|
||||
|
||||
use df_core::types::new_id;
|
||||
|
||||
use crate::state::AppState;
|
||||
|
||||
use crate::commands::now_millis;
|
||||
|
||||
use super::{AiChatEvent, AiSession, PendingApproval, ToolCallDraft};
|
||||
|
||||
/// RiskLevel → 审计记录字符串(low/medium/high)
|
||||
pub(crate) fn risk_str(r: RiskLevel) -> &'static str {
|
||||
match r {
|
||||
RiskLevel::Low => "low",
|
||||
RiskLevel::Medium => "medium",
|
||||
RiskLevel::High => "high",
|
||||
}
|
||||
}
|
||||
|
||||
/// 审计记录字符串 → RiskLevel(启动重建 pending_approvals 用,未知串返回 None 跳过)
|
||||
pub(crate) fn risk_from_str(s: &str) -> Option<RiskLevel> {
|
||||
match s {
|
||||
"low" => Some(RiskLevel::Low),
|
||||
"medium" => Some(RiskLevel::Medium),
|
||||
"high" => Some(RiskLevel::High),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 启动恢复:从审计表重建 pending_approvals(重启前未审批的工具调用,内存态已丢)
|
||||
///
|
||||
/// pending_approvals 是 AiSession 内存 HashMap,重启必丢。ai_tool_executions 表已存
|
||||
/// status=pending 的行(持久化真相源),此处读回重建内存态,使重启后待审批不丢。
|
||||
/// 前端经 ai_pending_tool_calls 查询 + switchConversation 恢复 toolCard 的 pending_approval 态。
|
||||
pub async fn restore_pending_approvals(state: &AppState) {
|
||||
let pending = match state.ai_tool_executions.list_pending().await {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
tracing::warn!("启动恢复 pending 审批失败(非阻断): {}", e);
|
||||
return;
|
||||
}
|
||||
};
|
||||
if pending.is_empty() {
|
||||
return;
|
||||
}
|
||||
let mut session = state.ai_session.lock().await;
|
||||
for rec in pending {
|
||||
let args: serde_json::Value = serde_json::from_str(&rec.arguments).unwrap_or_default();
|
||||
let Some(risk) = risk_from_str(&rec.risk_level) else { continue };
|
||||
session.pending_approvals.insert(
|
||||
rec.tool_call_id.clone(),
|
||||
PendingApproval {
|
||||
tool_call_id: rec.tool_call_id,
|
||||
tool_name: rec.tool_name,
|
||||
arguments: args,
|
||||
risk_level: risk,
|
||||
conversation_id: rec.conversation_id,
|
||||
recovered: true,
|
||||
},
|
||||
);
|
||||
}
|
||||
tracing::info!("启动恢复: {} 条 pending 工具审批重建到内存", session.pending_approvals.len());
|
||||
}
|
||||
|
||||
/// 写一条工具执行审计记录(insert 失败不阻断主流程,故 `let _ =`)
|
||||
///
|
||||
/// `decided_by` 有值(auto/human)= 已决策执行 → 记 executed_at;
|
||||
/// `None`(pending 待审批)→ executed_at 留空,待 audit_finalize 回填。
|
||||
pub(crate) async fn audit_tool_call(
|
||||
repo: &AiToolExecutionRepo,
|
||||
conv_id: &str,
|
||||
tool_call_id: &str,
|
||||
tool_name: &str,
|
||||
arguments: &str,
|
||||
status: &str,
|
||||
risk_level: RiskLevel,
|
||||
result: Option<String>,
|
||||
decided_by: Option<&str>,
|
||||
) {
|
||||
let executed_at = if decided_by.is_some() { Some(now_millis()) } else { None };
|
||||
let _ = repo
|
||||
.insert(AiToolExecutionRecord {
|
||||
id: new_id(),
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
tool_call_id: tool_call_id.to_string(),
|
||||
tool_name: tool_name.to_string(),
|
||||
arguments: arguments.to_string(),
|
||||
result,
|
||||
status: status.to_string(),
|
||||
risk_level: risk_str(risk_level).to_string(),
|
||||
requested_at: now_millis(),
|
||||
executed_at,
|
||||
decided_by: decided_by.map(|s| s.to_string()),
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
/// 审批后更新审计记录状态(按 tool_call_id 定位 pending 记录,回填 status/decided_by=human/executed_at/result)
|
||||
///
|
||||
/// 走专用 find_by_tool_call_id —— 通用 query 宏硬编码 ORDER BY created_at DESC,
|
||||
/// 而 ai_tool_executions 无该列,调用会报 "no such column: created_at" 被 unwrap_or_default 吞掉,
|
||||
/// 导致审批后审计记录永久卡 pending。
|
||||
pub(crate) async fn audit_finalize(state: &AppState, tool_call_id: &str, status: &str, result: Option<String>) {
|
||||
let Some(mut rec) = state
|
||||
.ai_tool_executions
|
||||
.find_by_tool_call_id(tool_call_id)
|
||||
.await
|
||||
.unwrap_or_default()
|
||||
else {
|
||||
tracing::warn!("audit_finalize: 未找到 tool_call_id={} 的审计记录", tool_call_id);
|
||||
return;
|
||||
};
|
||||
rec.status = status.to_string();
|
||||
rec.decided_by = Some("human".to_string());
|
||||
rec.executed_at = Some(now_millis());
|
||||
if let Some(r) = result {
|
||||
rec.result = Some(r);
|
||||
}
|
||||
let _ = state.ai_tool_executions.update_full(&rec).await;
|
||||
}
|
||||
|
||||
/// 处理流式接收的工具调用:Low 风险并行执行(join_all),Med/High 进审批门控
|
||||
/// 返回待审批的工具数量(0 = 全部自动执行完成)
|
||||
pub(crate) async fn process_tool_calls(
|
||||
session: &mut AiSession,
|
||||
tool_calls_acc: HashMap<u32, ToolCallDraft>,
|
||||
tools_arc: &Arc<AiToolRegistry>,
|
||||
db: &Arc<Database>,
|
||||
app_handle: &AppHandle,
|
||||
conv_id: &str,
|
||||
) -> usize {
|
||||
let mut tc_list: Vec<_> = tool_calls_acc.into_iter().collect();
|
||||
tc_list.sort_unstable_by_key(|(i, _)| *i);
|
||||
let mut pending_count = 0usize;
|
||||
let audit_repo = AiToolExecutionRepo::new(db);
|
||||
|
||||
// 解析 args + 批量发 Started(前端骨架按原始 index 顺序展示)
|
||||
let drafts: Vec<(u32, ToolCallDraft, serde_json::Value)> = tc_list.into_iter()
|
||||
.map(|(idx, draft)| {
|
||||
let args = serde_json::from_str(&draft.args).unwrap_or(serde_json::Value::Object(Default::default()));
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiToolCallStarted {
|
||||
id: draft.id.clone(),
|
||||
name: draft.name.clone(),
|
||||
args: args.clone(),
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
});
|
||||
(idx, draft, args)
|
||||
})
|
||||
.collect();
|
||||
|
||||
// 分类:Low 收集并行执行,Med/High 立即进审批门控(push 占位 tool_result)
|
||||
let mut low_risk: Vec<(ToolCallDraft, serde_json::Value)> = Vec::new();
|
||||
for (_, draft, args) in drafts {
|
||||
let risk_level = tools_arc.get(&draft.name).map(|t| t.risk_level).unwrap_or(RiskLevel::High);
|
||||
match risk_level {
|
||||
RiskLevel::Low => low_risk.push((draft, args)),
|
||||
RiskLevel::Medium | RiskLevel::High => {
|
||||
pending_count += 1;
|
||||
session.pending_approvals.insert(draft.id.clone(), PendingApproval {
|
||||
tool_call_id: draft.id.clone(),
|
||||
tool_name: draft.name.clone(),
|
||||
arguments: args.clone(),
|
||||
risk_level,
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
recovered: false,
|
||||
});
|
||||
session.messages.push(ChatMessage::tool_result(&draft.id, "需要用户审批,等待确认"));
|
||||
let reason = match risk_level {
|
||||
RiskLevel::High => "高风险操作,必须人工批准".to_string(),
|
||||
_ => "创建操作,请确认是否执行".to_string(),
|
||||
};
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiApprovalRequired {
|
||||
id: draft.id.clone(),
|
||||
name: draft.name.clone(),
|
||||
args: args.clone(),
|
||||
reason,
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
});
|
||||
audit_tool_call(&audit_repo, conv_id, &draft.id, &draft.name, &draft.args, "pending", risk_level, None, None).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Low 风险并行执行:execute + 即时 emit 在闭包内(不持 session 锁),
|
||||
// push tool_result / audit 在 join_all 后串行回填(持锁,与 Med/High 占位拼接)。
|
||||
// join_all 保序——结果顺序 = low_risk 输入顺序 = tc_list 原始 index 顺序,不额外 sort
|
||||
if !low_risk.is_empty() {
|
||||
let results: Vec<(ToolCallDraft, Result<String, String>)> =
|
||||
futures::future::join_all(low_risk.into_iter().map(|(draft, args)| {
|
||||
let tools = tools_arc.clone();
|
||||
let app_clone = app_handle.clone();
|
||||
let conv_clone = conv_id.to_string();
|
||||
async move {
|
||||
let result = tools.execute(&draft.name, args).await;
|
||||
match result {
|
||||
Ok(val) => {
|
||||
let _ = app_clone.emit("ai-chat-event", AiChatEvent::AiToolCallCompleted {
|
||||
id: draft.id.clone(),
|
||||
result: val.clone(),
|
||||
conversation_id: Some(conv_clone),
|
||||
});
|
||||
(draft, Ok(val.to_string()))
|
||||
}
|
||||
Err(e) => {
|
||||
let _ = app_clone.emit("ai-chat-event", AiChatEvent::AiError {
|
||||
error: format!("工具 {} 执行失败: {}", draft.name, e),
|
||||
conversation_id: Some(conv_clone),
|
||||
});
|
||||
(draft, Err(format!("错误: {}", e)))
|
||||
}
|
||||
}
|
||||
}
|
||||
})).await;
|
||||
|
||||
// 串行回填 tool_result + 审计(持 session 锁)
|
||||
for (draft, outcome) in results {
|
||||
let (status, content) = match outcome {
|
||||
Ok(c) => ("completed", c),
|
||||
Err(c) => ("failed", c),
|
||||
};
|
||||
session.messages.push(ChatMessage::tool_result(&draft.id, content.clone()));
|
||||
audit_tool_call(&audit_repo, conv_id, &draft.id, &draft.name, &draft.args, status, RiskLevel::Low, Some(content), Some("auto")).await;
|
||||
}
|
||||
}
|
||||
|
||||
pending_count
|
||||
}
|
||||
531
src-tauri/src/commands/ai/commands.rs
Normal file
531
src-tauri/src/commands/ai/commands.rs
Normal file
@@ -0,0 +1,531 @@
|
||||
//! 所有 `#[tauri::command]` IPC 函数 — 由 mod.rs 重导出供 invoke_handler 引用
|
||||
|
||||
use std::sync::atomic::Ordering;
|
||||
|
||||
use serde::Serialize;
|
||||
use tauri::{AppHandle, Emitter, State};
|
||||
|
||||
use df_ai::provider::ChatMessage;
|
||||
use df_core::types::new_id;
|
||||
use df_storage::models::AiProviderRecord;
|
||||
|
||||
use crate::state::AppState;
|
||||
use crate::commands::now_millis;
|
||||
|
||||
use super::agentic::{run_agentic_loop, try_continue_agent_loop};
|
||||
use super::audit::audit_finalize;
|
||||
use super::conversation::save_conversation;
|
||||
use super::knowledge_inject::build_knowledge_context;
|
||||
use super::prompt::build_system_prompt;
|
||||
use super::skills::{read_skill_content, SkillInfo, skills_cached};
|
||||
|
||||
use super::AiChatEvent;
|
||||
|
||||
// ============================================================
|
||||
// 发送 / 审批 / 控制
|
||||
// ============================================================
|
||||
|
||||
/// 发送消息并获取流式 AI 响应
|
||||
///
|
||||
/// 非阻塞:立即返回 "ok",通过 ai-chat-event 事件流式推送
|
||||
#[tauri::command]
|
||||
pub async fn ai_chat_send(
|
||||
app: AppHandle,
|
||||
state: State<'_, AppState>,
|
||||
message: String,
|
||||
language: Option<String>,
|
||||
skill: Option<String>,
|
||||
) -> Result<String, String> {
|
||||
// 获取活跃提供商(只读,失败可直接返回,不影响生成标志)
|
||||
let provider_config = super::prompt::get_active_provider(&state).await?;
|
||||
|
||||
// 原子检查并占用生成标志,防止并发双发;同步追加用户消息,按需自动创建对话
|
||||
{
|
||||
let mut session = state.ai_session.lock().await;
|
||||
if session.generating {
|
||||
return Err("AI 正在生成中,请等待完成".to_string());
|
||||
}
|
||||
session.generating = true;
|
||||
session.stop_flag.store(false, Ordering::SeqCst);
|
||||
session.agent_language = language.clone();
|
||||
session.messages.push(ChatMessage::user(&message));
|
||||
|
||||
// 首次发送时生成对话 id(懒创建:不立即落库,避免空对话残留;
|
||||
// 实际记录由 save_conversation 在生成内容后 upsert 写入)
|
||||
if session.active_conversation_id.is_none() {
|
||||
let conv_id = new_id();
|
||||
session.active_conversation_id = Some(conv_id);
|
||||
session.active_conv_created_at = Some(now_millis());
|
||||
}
|
||||
}
|
||||
|
||||
// 获取工具定义(预取仅用于触发注册表初始化,实际 tool_defs 在 agentic loop 内部按需获取)
|
||||
let _tool_defs = state.ai_tools.tool_definitions();
|
||||
let lang = language.unwrap_or_else(|| "zh-CN".to_string());
|
||||
let mut system_prompt = build_system_prompt(&state, &lang).await;
|
||||
// 技能注入:读 SKILL.md 全文拼到 system prompt 前作为指令
|
||||
if let Some(ref skill_name) = skill {
|
||||
if let Some(content) = read_skill_content(skill_name) {
|
||||
system_prompt = format!("# 技能指令: {}\n\n{}\n\n---\n{}", skill_name, content, system_prompt);
|
||||
}
|
||||
}
|
||||
// 快照当前对话 ID,供知识注入溯源 + spawn 后台 loop(不受切换影响)
|
||||
let conv_id = {
|
||||
let session = state.ai_session.lock().await;
|
||||
session.active_conversation_id.clone().unwrap_or_default()
|
||||
};
|
||||
|
||||
// 知识注入:检索相关知识拼到 system prompt 前(可配置开关 auto_inject,默认开)
|
||||
// 最终顺序:[知识库上下文] --- [技能指令] --- [原始 system prompt]
|
||||
{
|
||||
let config = state.knowledge_config.lock().await.clone();
|
||||
let knowledge_context = build_knowledge_context(&state, &conv_id, &message, &config).await;
|
||||
if !knowledge_context.is_empty() {
|
||||
system_prompt = format!("{}\n\n---\n{}", knowledge_context, system_prompt);
|
||||
}
|
||||
}
|
||||
|
||||
// 在后台任务中执行流式调用
|
||||
let session_arc = state.ai_session.clone();
|
||||
let tools_arc = state.ai_tools.clone();
|
||||
let db = state.db.clone();
|
||||
let app_handle = app.clone();
|
||||
let knowledge_config = state.knowledge_config.lock().await.clone();
|
||||
let llm_concurrency = state.llm_concurrency.clone();
|
||||
|
||||
tauri::async_runtime::spawn(async move {
|
||||
run_agentic_loop(session_arc, tools_arc, db, app_handle, provider_config, system_prompt, conv_id, knowledge_config, llm_concurrency).await;
|
||||
});
|
||||
|
||||
Ok("ok".to_string())
|
||||
}
|
||||
|
||||
/// 批准/拒绝挂起的工具调用
|
||||
#[tauri::command]
|
||||
pub async fn ai_approve(
|
||||
app: AppHandle,
|
||||
state: State<'_, AppState>,
|
||||
tool_call_id: String,
|
||||
approved: bool,
|
||||
) -> Result<String, String> {
|
||||
let mut session = state.ai_session.lock().await;
|
||||
|
||||
let approval = session
|
||||
.pending_approvals
|
||||
.remove(&tool_call_id)
|
||||
.ok_or_else(|| format!("未找到挂起的审批: {}", tool_call_id))?;
|
||||
// recovered 字段保留读取(标记重启恢复来源,未来扩展用),本次修复移除 if !recovered 落库守卫。
|
||||
let _recovered = approval.recovered;
|
||||
|
||||
if !approved {
|
||||
// 替换占位 tool_result 为拒绝结果
|
||||
session.messages.replace_tool_result_content(&tool_call_id, "用户拒绝了此操作");
|
||||
let conv_id = approval.conversation_id.clone();
|
||||
let _ = app.emit("ai-chat-event", AiChatEvent::AiApprovalResult {
|
||||
id: tool_call_id.clone(),
|
||||
approved: false,
|
||||
conversation_id: conv_id.clone(),
|
||||
});
|
||||
drop(session);
|
||||
// 拒绝结果立即落库(含 recovered 积压审批)——switch 时已 restore_from_messages 载完整历史,
|
||||
// messages 非空,save 不会污染老对话;原 if !recovered 守卫前提不成立已移除。
|
||||
if let Some(ref cid) = conv_id {
|
||||
save_conversation(&state.ai_session, &state.db, cid, None, None).await;
|
||||
}
|
||||
// 审计:拒绝(决策者=human)
|
||||
audit_finalize(&state, &tool_call_id, "rejected", None).await;
|
||||
// 所有待审批处理完毕后恢复 agentic 循环
|
||||
try_continue_agent_loop(&app, &state).await;
|
||||
return Ok("rejected".to_string());
|
||||
}
|
||||
|
||||
// 执行工具(通过真实 repo 调用)
|
||||
let args = approval.arguments.clone();
|
||||
let id = tool_call_id.clone();
|
||||
let conv_id = approval.conversation_id.clone();
|
||||
drop(session); // 释放锁后再执行
|
||||
|
||||
let exec_result = state.ai_tools.execute(&approval.tool_name, args.clone()).await;
|
||||
// 工具失败不 return Err:把错误包成 tool_result,落库 + emit completed + 续循环全走通。
|
||||
// 否则前端 approveToolCall 的 catch 会回滚 pending_approval,审批按钮卡死无法消除。
|
||||
let (audit_status, result_val) = match &exec_result {
|
||||
Ok(val) => ("executed", val.clone()),
|
||||
Err(e) => ("failed", serde_json::Value::String(e.to_string())),
|
||||
};
|
||||
// 审计:人工审批后无论成败回填(决策者=human)
|
||||
audit_finalize(&state, &tool_call_id, audit_status, Some(result_val.to_string())).await;
|
||||
|
||||
// 重新获取锁,替换占位 tool_result 为真实结果(失败时为错误信息,LLM 据此决定下一步)
|
||||
let mut session = state.ai_session.lock().await;
|
||||
session.messages.replace_tool_result_content(&id, &result_val.to_string());
|
||||
|
||||
let _ = app.emit("ai-chat-event", AiChatEvent::AiToolCallCompleted {
|
||||
id: id.clone(),
|
||||
result: result_val.clone(),
|
||||
conversation_id: conv_id.clone(),
|
||||
});
|
||||
let _ = app.emit("ai-chat-event", AiChatEvent::AiApprovalResult {
|
||||
id,
|
||||
approved: true,
|
||||
conversation_id: conv_id.clone(),
|
||||
});
|
||||
drop(session);
|
||||
|
||||
// 审批执行结果立即落库,不依赖后续 agentic loop(避免 loop 异常退出时丢失真实结果)
|
||||
// 含 recovered 积压审批——switch 时已 restore_from_messages 载完整历史,messages 非空,
|
||||
// save 不污染老对话;原 if !recovered 守卫前提不成立已移除。
|
||||
if let Some(ref cid) = conv_id {
|
||||
save_conversation(&state.ai_session, &state.db, cid, None, None).await;
|
||||
}
|
||||
|
||||
// 所有待审批处理完毕后恢复 agentic 循环(recovered 无 live loop,try_continue 因 generating=false 自然不续)
|
||||
try_continue_agent_loop(&app, &state).await;
|
||||
|
||||
Ok("executed".to_string())
|
||||
}
|
||||
|
||||
/// 待审批工具调用信息(前端恢复 toolCard pending_approval 态用)
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct PendingToolCallInfo {
|
||||
pub tool_call_id: String,
|
||||
pub conversation_id: Option<String>,
|
||||
}
|
||||
|
||||
/// 查询某对话积压的待审批工具(前端 switchConversation 后恢复 toolCard 的 pending_approval 态)
|
||||
#[tauri::command]
|
||||
pub async fn ai_pending_tool_calls(
|
||||
state: State<'_, AppState>,
|
||||
conv_id: String,
|
||||
) -> Result<Vec<PendingToolCallInfo>, String> {
|
||||
let session = state.ai_session.lock().await;
|
||||
let list = session
|
||||
.pending_approvals
|
||||
.values()
|
||||
.filter(|a| a.conversation_id.as_deref() == Some(conv_id.as_str()))
|
||||
.map(|a| PendingToolCallInfo {
|
||||
tool_call_id: a.tool_call_id.clone(),
|
||||
conversation_id: a.conversation_id.clone(),
|
||||
})
|
||||
.collect();
|
||||
Ok(list)
|
||||
}
|
||||
|
||||
/// 清空对话历史
|
||||
#[tauri::command]
|
||||
pub async fn ai_chat_clear(state: State<'_, AppState>) -> Result<(), String> {
|
||||
let mut session = state.ai_session.lock().await;
|
||||
session.messages.clear();
|
||||
session.pending_approvals.clear();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 停止当前 AI 生成
|
||||
///
|
||||
/// 两种场景:
|
||||
/// - loop 正在流式生成:置 stop_flag,stream_llm / 循环检查点尽快退出并 emit AiCompleted
|
||||
/// - 有挂起审批(loop 已 return 等待中):stop_flag 无人读取,直接清审批 + 复位 generating,
|
||||
/// 否则停止按钮表面无反应、会话卡在 generating=true
|
||||
#[tauri::command]
|
||||
pub async fn ai_chat_stop(state: State<'_, AppState>, app: AppHandle) -> Result<(), String> {
|
||||
let mut session = state.ai_session.lock().await;
|
||||
if !session.generating {
|
||||
return Ok(());
|
||||
}
|
||||
if !session.pending_approvals.is_empty() {
|
||||
// 审批等待态:loop 已退出,直接清理让会话立即可用
|
||||
session.pending_approvals.clear();
|
||||
session.generating = false;
|
||||
session.stop_flag.store(true, Ordering::SeqCst); // 双保险:防 try_continue 误判重启
|
||||
let conv_id = session.active_conversation_id.clone();
|
||||
drop(session);
|
||||
let _ = app.emit("ai-chat-event", AiChatEvent::AiCompleted { total_tokens: 0, prompt_tokens: 0, completion_tokens: 0, conversation_id: conv_id });
|
||||
return Ok(());
|
||||
}
|
||||
// 流式生成中:置位让 loop 自行收尾
|
||||
session.stop_flag.store(true, Ordering::SeqCst);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// 提供商管理
|
||||
// ============================================================
|
||||
|
||||
/// 列出所有已配置的 AI 提供商(is_default 真相源为 DB,重启不丢)
|
||||
#[tauri::command]
|
||||
pub async fn ai_list_providers(state: State<'_, AppState>) -> Result<Vec<AiProviderRecord>, String> {
|
||||
state.ai_providers.list_all().await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 保存/更新 AI 提供商配置
|
||||
#[tauri::command]
|
||||
pub async fn ai_save_provider(
|
||||
state: State<'_, AppState>,
|
||||
id: Option<String>,
|
||||
name: String,
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
default_model: String,
|
||||
provider_type: String,
|
||||
) -> Result<String, String> {
|
||||
// 编辑已有提供商时保留原 created_at,避免被覆盖
|
||||
let created_at = match &id {
|
||||
Some(pid) => state.ai_providers.get_by_id(pid).await
|
||||
.map_err(|e| e.to_string())?
|
||||
.map(|p| p.created_at)
|
||||
.unwrap_or_else(now_millis),
|
||||
None => now_millis(),
|
||||
};
|
||||
// is_default:编辑保留原值;新建时若全表尚无默认则设为默认(首个自动默认,避免无默认可用)
|
||||
let is_default = match &id {
|
||||
Some(pid) => state.ai_providers.get_by_id(pid).await
|
||||
.map_err(|e| e.to_string())?
|
||||
.map(|p| p.is_default)
|
||||
.unwrap_or(false),
|
||||
None => !state.ai_providers.list_all().await
|
||||
.map_err(|e| e.to_string())?
|
||||
.iter().any(|p| p.is_default),
|
||||
};
|
||||
let record = AiProviderRecord {
|
||||
id: id.unwrap_or_else(new_id),
|
||||
name,
|
||||
provider_type: if provider_type.is_empty() { "openai_compat".to_string() } else { provider_type },
|
||||
api_key,
|
||||
base_url,
|
||||
default_model,
|
||||
models: None,
|
||||
is_default,
|
||||
config: None,
|
||||
created_at,
|
||||
updated_at: now_millis(),
|
||||
};
|
||||
let id = record.id.clone();
|
||||
state
|
||||
.ai_providers
|
||||
.insert(record)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
Ok(id)
|
||||
}
|
||||
|
||||
/// 设置活跃提供商(互斥落库:目标置默认、其余清默认,重启不丢)
|
||||
#[tauri::command]
|
||||
pub async fn ai_set_provider(
|
||||
state: State<'_, AppState>,
|
||||
provider_id: String,
|
||||
) -> Result<(), String> {
|
||||
// 验证提供商存在
|
||||
let provider = state
|
||||
.ai_providers
|
||||
.get_by_id(&provider_id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("提供商不存在: {}", provider_id))?;
|
||||
|
||||
// 互斥写 DB:目标 is_default=true,其余=false。仅写变化的记录。
|
||||
let providers = state.ai_providers.list_all().await.map_err(|e| e.to_string())?;
|
||||
for p in &providers {
|
||||
let should = p.id == provider_id;
|
||||
if p.is_default != should {
|
||||
let mut updated = p.clone();
|
||||
updated.is_default = should;
|
||||
updated.updated_at = now_millis();
|
||||
state.ai_providers.update_full(&updated).await.map_err(|e| e.to_string())?;
|
||||
}
|
||||
}
|
||||
|
||||
let mut session = state.ai_session.lock().await;
|
||||
session.active_provider_id = Some(provider.id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 删除 AI 提供商
|
||||
#[tauri::command]
|
||||
pub async fn ai_delete_provider(
|
||||
state: State<'_, AppState>,
|
||||
provider_id: String,
|
||||
) -> Result<(), String> {
|
||||
state.ai_providers.delete(&provider_id).await.map_err(|e| e.to_string())?;
|
||||
// 删除的若是当前默认,清空 active 指向,避免悬空
|
||||
let mut session = state.ai_session.lock().await;
|
||||
if session.active_provider_id.as_deref() == Some(&provider_id) {
|
||||
session.active_provider_id = None;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// 对话管理
|
||||
// ============================================================
|
||||
|
||||
/// 创建新对话
|
||||
#[tauri::command]
|
||||
pub async fn ai_conversation_create(state: State<'_, AppState>) -> Result<serde_json::Value, String> {
|
||||
// 懒创建:仅生成 id 存内存,不落库;避免新建后不发消息产生空记录。
|
||||
// 首条消息发送后由 save_conversation upsert 写入。
|
||||
let id = new_id();
|
||||
let now = now_millis();
|
||||
|
||||
let mut session = state.ai_session.lock().await;
|
||||
session.active_conversation_id = Some(id.clone());
|
||||
session.active_conv_created_at = Some(now);
|
||||
session.messages.clear();
|
||||
session.pending_approvals.clear();
|
||||
|
||||
Ok(serde_json::json!({ "id": id }))
|
||||
}
|
||||
|
||||
/// 列出对话(仅摘要,不含 messages 全文)
|
||||
///
|
||||
/// limit 默认 50 防数据膨胀;include_archived 默认 false(归档对话默认隐藏)。
|
||||
#[tauri::command]
|
||||
pub async fn ai_conversation_list(
|
||||
state: State<'_, AppState>,
|
||||
limit: Option<usize>,
|
||||
include_archived: Option<bool>,
|
||||
) -> Result<Vec<serde_json::Value>, String> {
|
||||
let limit = limit.unwrap_or(50);
|
||||
let include_archived = include_archived.unwrap_or(false);
|
||||
let records = state.ai_conversations.list_all().await.map_err(|e| e.to_string())?;
|
||||
// list_all 已按 created_at DESC(最新在前);默认排除归档 + 截断 limit
|
||||
let summaries: Vec<serde_json::Value> = records.iter()
|
||||
.filter(|r| include_archived || !r.archived)
|
||||
.take(limit)
|
||||
.map(|r| {
|
||||
// 修复 models 字段类型 bug:r.models 是 JSON 字符串,前端期望数组
|
||||
let models: Vec<String> = r.models.as_deref()
|
||||
.and_then(|s| serde_json::from_str(s).ok())
|
||||
.unwrap_or_default();
|
||||
serde_json::json!({
|
||||
"id": r.id,
|
||||
"title": r.title,
|
||||
"provider_id": r.provider_id,
|
||||
"model": r.model,
|
||||
"models": models,
|
||||
"archived": r.archived,
|
||||
"prompt_tokens": r.prompt_tokens,
|
||||
"completion_tokens": r.completion_tokens,
|
||||
"created_at": r.created_at,
|
||||
"updated_at": r.updated_at,
|
||||
})
|
||||
}).collect();
|
||||
Ok(summaries)
|
||||
}
|
||||
|
||||
/// 切换到指定对话(从 DB 加载 messages 到内存 + 返回 messages 给前端)
|
||||
#[tauri::command]
|
||||
pub async fn ai_conversation_switch(
|
||||
state: State<'_, AppState>,
|
||||
conversation_id: String,
|
||||
) -> Result<serde_json::Value, String> {
|
||||
let record = state.ai_conversations.get_by_id(&conversation_id).await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("对话不存在: {}", conversation_id))?;
|
||||
|
||||
let messages: Vec<ChatMessage> = serde_json::from_str(&record.messages)
|
||||
.map_err(|e| format!("解析消息失败: {}", e))?;
|
||||
|
||||
let messages_json = record.messages.clone();
|
||||
let title = record.title.clone();
|
||||
|
||||
let mut session = state.ai_session.lock().await;
|
||||
// 生成中允许只读切换:返回目标对话的 messages 供前端展示,但不修改 session 状态
|
||||
// 后台 loop 持有快照的 conv_id,不受 active_conversation_id 变更影响
|
||||
if session.generating {
|
||||
return Ok(serde_json::json!({
|
||||
"id": record.id,
|
||||
"title": title,
|
||||
"messages": messages_json,
|
||||
"readonly": true,
|
||||
}));
|
||||
}
|
||||
session.active_conversation_id = Some(conversation_id.clone());
|
||||
session.messages.restore_from_messages(messages);
|
||||
// 仅清空目标对话自身的 pending_approvals,保留其他对话的(防 init 重建的内存 HashMap 被清空,
|
||||
// 重启恢复链路:restore_pending_approvals(init 重建) → switchConversation(此处不清目标对话的)
|
||||
// → ai_pending_tool_calls 查询 → ai_approve 落库)
|
||||
session.pending_approvals.retain(|_, a| a.conversation_id.as_deref() != Some(&conversation_id));
|
||||
|
||||
Ok(serde_json::json!({
|
||||
"id": record.id,
|
||||
"title": title,
|
||||
"messages": messages_json,
|
||||
}))
|
||||
}
|
||||
|
||||
/// 删除对话
|
||||
#[tauri::command]
|
||||
pub async fn ai_conversation_delete(
|
||||
state: State<'_, AppState>,
|
||||
conversation_id: String,
|
||||
) -> Result<(), String> {
|
||||
state.ai_conversations.delete(&conversation_id).await.map_err(|e| e.to_string())?;
|
||||
|
||||
let mut session = state.ai_session.lock().await;
|
||||
if session.active_conversation_id.as_deref() == Some(&conversation_id) {
|
||||
session.active_conversation_id = None;
|
||||
session.messages.clear();
|
||||
session.pending_approvals.clear();
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 重命名对话标题
|
||||
#[tauri::command]
|
||||
pub async fn ai_conversation_rename(
|
||||
state: State<'_, AppState>,
|
||||
conversation_id: String,
|
||||
title: String,
|
||||
) -> Result<(), String> {
|
||||
let title = title.trim().to_string();
|
||||
if title.is_empty() {
|
||||
return Err("标题不能为空".to_string());
|
||||
}
|
||||
state.ai_conversations.update_field(&conversation_id, "title", &title)
|
||||
.await.map_err(|e| e.to_string())?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 归档/取消归档对话(归档后在侧栏折叠分组展示)
|
||||
#[tauri::command]
|
||||
pub async fn ai_conversation_archive(
|
||||
state: State<'_, AppState>,
|
||||
conversation_id: String,
|
||||
archived: bool,
|
||||
) -> Result<(), String> {
|
||||
state.ai_conversations
|
||||
.set_archived(&conversation_id, archived)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 列出本机 Claude 技能(skills + commands + plugins 三类),供前端 `/` 联想
|
||||
#[tauri::command]
|
||||
pub async fn ai_list_skills() -> Result<Vec<SkillInfo>, String> {
|
||||
// 命中进程内缓存,命中后仅 clone,不重复扫盘
|
||||
Ok(skills_cached().clone())
|
||||
}
|
||||
|
||||
/// 设置 LLM 调用并发上限(运行时调整,立即生效)
|
||||
///
|
||||
/// 软收敛:缩并发时已持有旧 permit 的任务继续执行不受影响,待其释放后新限制完全生效。
|
||||
/// None 表示该层不变(前端可单独调一层)。值下限为 1。
|
||||
#[tauri::command]
|
||||
pub async fn ai_set_concurrency_config(
|
||||
state: State<'_, AppState>,
|
||||
global_limit: Option<u32>,
|
||||
per_conv_limit: Option<u32>,
|
||||
) -> Result<(), String> {
|
||||
// 下限 1,无上限;同时给 global 时约束 per-conv 不超过 global
|
||||
if let Some(g) = global_limit {
|
||||
state.llm_concurrency.set_global(g.max(1) as usize).await;
|
||||
}
|
||||
if let Some(p) = per_conv_limit {
|
||||
let mut p = p.max(1);
|
||||
if let Some(g) = global_limit {
|
||||
p = p.min(g.max(1));
|
||||
}
|
||||
state.llm_concurrency.set_per_conv(p as usize).await;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
269
src-tauri/src/commands/ai/conversation.rs
Normal file
269
src-tauri/src/commands/ai/conversation.rs
Normal file
@@ -0,0 +1,269 @@
|
||||
//! 对话持久化 + Token 累加器
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use df_storage::crud::AiConversationRepo;
|
||||
use df_storage::db::Database;
|
||||
|
||||
use crate::commands::now_millis;
|
||||
|
||||
use super::AiSession;
|
||||
|
||||
/// Token 用量累加器(agent loop 生命周期内各轮叠加)
|
||||
///
|
||||
/// 纯结构 + 方法:抽自 run_agentic_loop 的 `total_prompt`/`total_completion` 双计数器,
|
||||
/// 保证多轮累加、None 起始、跨 loop 实例叠加语义一致且可单测。
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub(crate) struct TokenAccumulator {
|
||||
prompt: u32,
|
||||
completion: u32,
|
||||
}
|
||||
|
||||
impl TokenAccumulator {
|
||||
/// 叠加一轮用量(round_usage 为本轮流式末 chunk 的累计用量)
|
||||
pub(crate) fn add(&mut self, prompt: u32, completion: u32) {
|
||||
self.prompt += prompt;
|
||||
self.completion += completion;
|
||||
}
|
||||
|
||||
pub(crate) fn prompt(&self) -> u32 {
|
||||
self.prompt
|
||||
}
|
||||
|
||||
pub(crate) fn completion(&self) -> u32 {
|
||||
self.completion
|
||||
}
|
||||
|
||||
pub(crate) fn total(&self) -> u32 {
|
||||
self.prompt + self.completion
|
||||
}
|
||||
}
|
||||
|
||||
/// 把单轮增量叠加到 DB 的 Option<i64> 字段(读旧值+增量,跨 loop 实例防覆盖)
|
||||
///
|
||||
/// 纯函数:抽自 save_conversation 的 token 累加逻辑,None 起始当作 0。
|
||||
pub(crate) fn accumulate_tokens(old: Option<i64>, add: u32) -> Option<i64> {
|
||||
Some(old.unwrap_or(0) + add as i64)
|
||||
}
|
||||
|
||||
/// 持久化截断阈值:超过此长度的消息 content 落库前截断头尾各保 HEAD/TAIL 字符。
|
||||
///
|
||||
/// 防 read_file 1MB 洞 / list_directory 大体量结果落库后每轮重发累积致 token 暴增
|
||||
/// (Sprint 19 实测单对话 in=115万 / 消息体 1.6MB)。仅作用于持久化视图,不污染内存真相源。
|
||||
pub(crate) const TRUNCATE_THRESHOLD: usize = 50_000;
|
||||
const TRUNCATE_HEAD: usize = 20_000;
|
||||
const TRUNCATE_TAIL: usize = 20_000;
|
||||
|
||||
/// 落库前对超长 content 做截断(保留头尾各 ~20KB + 中段标注省略字符数)。
|
||||
/// 50KB 阈值以下原样返回(零开销);按字符而非字节切避免 UTF-8 切坏中文。
|
||||
pub(crate) fn truncate_for_persist(content: &str) -> String {
|
||||
let chars: Vec<char> = content.chars().collect();
|
||||
if chars.len() <= TRUNCATE_THRESHOLD {
|
||||
return content.to_string();
|
||||
}
|
||||
let head: String = chars.iter().take(TRUNCATE_HEAD).collect();
|
||||
let tail: String = chars[chars.len() - TRUNCATE_TAIL..].iter().collect();
|
||||
let omitted = chars.len() - TRUNCATE_HEAD - TRUNCATE_TAIL;
|
||||
format!(
|
||||
"{}\n\n[...省略 {} 字符(已截断,完整内容仅在内存态可读)...]\n\n{}",
|
||||
head, omitted, tail
|
||||
)
|
||||
}
|
||||
|
||||
/// 保存对话到数据库(按 conv_id 写库,不受 active_conversation_id 切换影响)
|
||||
///
|
||||
/// 写 messages + updated_at + 累加 token 用量 + 首次落库的 model;标题由 ensure_conversation_title 单独生成。
|
||||
/// token 走累加模式:upsert 读旧值叠加,保证审批暂停→恢复跨 loop 实例不覆盖丢失。
|
||||
/// model 仅首次落库写入 + 旧记录缺值时补填(不覆盖历史已存值,兼容本次改造前的老对话)。
|
||||
pub(crate) async fn save_conversation(
|
||||
session_arc: &Arc<Mutex<AiSession>>,
|
||||
db: &Arc<Database>,
|
||||
conv_id: &str,
|
||||
usage: Option<&df_ai::provider::TokenUsage>,
|
||||
model: Option<&str>,
|
||||
) {
|
||||
// 取 messages + 懒创建首次落库所需的 provider_id/created_at
|
||||
// 工具结果(content)超 50KB 时截断头尾各 ~20KB + 中段标注,防大体量结果(read_file 1MB 洞 /
|
||||
// list_directory 13782 项)落库后每轮重发累积致 token 暴增。仅影响持久化视图,不污染
|
||||
// 内存真相源(ContextManager)——build_for_request 仍读全量 messages。
|
||||
let (messages_json, provider_id, created_at) = {
|
||||
let session = session_arc.lock().await;
|
||||
let mut msgs = session.messages.all_messages_clone();
|
||||
for m in &mut msgs {
|
||||
m.content = truncate_for_persist(&m.content);
|
||||
}
|
||||
(
|
||||
serde_json::to_string(&msgs).unwrap_or_else(|_| "[]".to_string()),
|
||||
session.active_provider_id.clone(),
|
||||
session.active_conv_created_at.clone(),
|
||||
)
|
||||
};
|
||||
|
||||
let conv_repo = AiConversationRepo::new(db);
|
||||
match conv_repo.get_by_id(conv_id).await {
|
||||
Ok(Some(mut rec)) => {
|
||||
// 已落库:更新 messages + updated_at;token 累加(读旧值+新值,跨 loop 实例防覆盖)
|
||||
rec.messages = messages_json;
|
||||
rec.updated_at = now_millis();
|
||||
if let Some(u) = usage {
|
||||
rec.prompt_tokens = accumulate_tokens(rec.prompt_tokens, u.prompt_tokens);
|
||||
rec.completion_tokens = accumulate_tokens(rec.completion_tokens, u.completion_tokens);
|
||||
}
|
||||
// model: 旧记录缺值时补填(不覆盖已有);models: 去重追加用过的所有 model(JSON 数组)
|
||||
if let Some(m) = model {
|
||||
if rec.model.is_none() { rec.model = Some(m.to_string()); }
|
||||
let mut list: Vec<String> = rec.models
|
||||
.as_deref()
|
||||
.and_then(|s| serde_json::from_str(s).ok())
|
||||
.unwrap_or_default();
|
||||
if !list.iter().any(|x| x == m) { list.push(m.to_string()); }
|
||||
rec.models = Some(serde_json::to_string(&list).unwrap_or_else(|_| "[]".to_string()));
|
||||
}
|
||||
if let Err(e) = conv_repo.update_full(&rec).await {
|
||||
tracing::warn!("更新对话失败 {conv_id}: {e}");
|
||||
}
|
||||
}
|
||||
Ok(None) => {
|
||||
// 懒创建首次落库(此为空对话不落库的落库点:走到这里 messages 必非空)
|
||||
let now = now_millis();
|
||||
let rec = df_storage::models::AiConversationRecord {
|
||||
id: conv_id.to_string(),
|
||||
title: None,
|
||||
messages: messages_json,
|
||||
provider_id,
|
||||
model: model.map(|m| m.to_string()),
|
||||
models: model.map(|m| serde_json::to_string(&[m]).unwrap_or_else(|_| "[]".to_string())),
|
||||
archived: false,
|
||||
prompt_tokens: usage.map(|u| u.prompt_tokens as i64),
|
||||
completion_tokens: usage.map(|u| u.completion_tokens as i64),
|
||||
created_at: created_at.unwrap_or_else(|| now.clone()),
|
||||
updated_at: now,
|
||||
};
|
||||
if let Err(e) = conv_repo.insert(rec).await {
|
||||
tracing::warn!("落库对话失败 {conv_id}: {e}");
|
||||
}
|
||||
}
|
||||
Err(e) => tracing::warn!("读取对话 {conv_id} 失败: {e}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
// ---------- TokenAccumulator + accumulate_tokens ----------
|
||||
|
||||
#[test]
|
||||
fn accumulator_starts_zero() {
|
||||
let acc = TokenAccumulator::default();
|
||||
assert_eq!(acc.prompt(), 0);
|
||||
assert_eq!(acc.completion(), 0);
|
||||
assert_eq!(acc.total(), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accumulator_single_add() {
|
||||
let mut acc = TokenAccumulator::default();
|
||||
acc.add(100, 50);
|
||||
assert_eq!(acc.prompt(), 100);
|
||||
assert_eq!(acc.completion(), 50);
|
||||
assert_eq!(acc.total(), 150);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accumulator_multi_round_accumulation() {
|
||||
// 多轮累加(模拟 agent loop 多次迭代)
|
||||
let mut acc = TokenAccumulator::default();
|
||||
acc.add(100, 20); // 轮1
|
||||
acc.add(200, 40); // 轮2
|
||||
acc.add(50, 10); // 轮3
|
||||
assert_eq!(acc.prompt(), 350);
|
||||
assert_eq!(acc.completion(), 70);
|
||||
assert_eq!(acc.total(), 420);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accumulator_add_zero_is_noop() {
|
||||
let mut acc = TokenAccumulator::default();
|
||||
acc.add(10, 5);
|
||||
acc.add(0, 0);
|
||||
assert_eq!(acc.total(), 15);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accumulate_tokens_from_none() {
|
||||
// 新记录(None 起始)落库
|
||||
assert_eq!(accumulate_tokens(None, 100), Some(100));
|
||||
assert_eq!(accumulate_tokens(None, 0), Some(0));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accumulate_tokens_adds_to_existing() {
|
||||
// 跨 loop 实例叠加:旧值 + 新增不覆盖
|
||||
assert_eq!(accumulate_tokens(Some(500), 100), Some(600));
|
||||
assert_eq!(accumulate_tokens(Some(0), 42), Some(42));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accumulate_tokens_multi_round_db_simulation() {
|
||||
// 模拟 save_conversation 多次落库累加(审批暂停→恢复跨 loop)
|
||||
let mut field: Option<i64> = None;
|
||||
field = accumulate_tokens(field, 100); // 首次
|
||||
field = accumulate_tokens(field, 200); // 二次
|
||||
field = accumulate_tokens(field, 50); // 三次
|
||||
assert_eq!(field, Some(350));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accumulator_and_db_accumulate_are_consistent() {
|
||||
// loop 内 TokenAccumulator 与落库 accumulate_tokens 总量语义一致
|
||||
let mut acc = TokenAccumulator::default();
|
||||
let mut db_prompt: Option<i64> = None;
|
||||
let mut db_completion: Option<i64> = None;
|
||||
for (p, c) in [(100u32, 20u32), (200, 40), (50, 10)] {
|
||||
acc.add(p, c);
|
||||
db_prompt = accumulate_tokens(db_prompt, p);
|
||||
db_completion = accumulate_tokens(db_completion, c);
|
||||
}
|
||||
assert_eq!(acc.prompt() as i64, db_prompt.unwrap());
|
||||
assert_eq!(acc.completion() as i64, db_completion.unwrap());
|
||||
}
|
||||
|
||||
// ---------- truncate_for_persist ----------
|
||||
|
||||
#[test]
|
||||
fn truncate_short_content_unchanged() {
|
||||
// 阈值以下原样返回
|
||||
assert_eq!(truncate_for_persist("hello"), "hello");
|
||||
assert_eq!(truncate_for_persist(""), "");
|
||||
let near_limit: String = "a".repeat(TRUNCATE_THRESHOLD);
|
||||
assert_eq!(truncate_for_persist(&near_limit).len(), TRUNCATE_THRESHOLD);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncate_long_content_keeps_head_and_tail() {
|
||||
// 超阈值:保留头尾各 TRUNCATE_HEAD/TAIL 字符 + 中段标注
|
||||
let long: String = "x".repeat(TRUNCATE_THRESHOLD + 1000);
|
||||
let result = truncate_for_persist(&long);
|
||||
// 头尾各 20k 字符应在结果中
|
||||
assert!(result.starts_with(&"x".repeat(TRUNCATE_HEAD)));
|
||||
assert!(result.ends_with(&"x".repeat(TRUNCATE_TAIL)));
|
||||
// 中段标注存在 + 标注省略字符数(中段 = 总长 - 头 - 尾 = 51000 - 20000 - 20000 = 11000)
|
||||
assert!(result.contains("已截断"));
|
||||
assert!(result.contains("省略 11000 字符"));
|
||||
// 结果总长应远小于原长(20k 头 + 20k 尾 + 标注)
|
||||
assert!(result.chars().count() < TRUNCATE_THRESHOLD + 1000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncate_preserves_utf8_chinese() {
|
||||
// 按字符切不切坏 UTF-8 中文
|
||||
let chinese: String = "中".repeat(TRUNCATE_THRESHOLD + 500);
|
||||
let result = truncate_for_persist(&chinese);
|
||||
assert!(result.starts_with('中'));
|
||||
assert!(result.ends_with('中'));
|
||||
assert!(result.contains("已截断"));
|
||||
}
|
||||
}
|
||||
557
src-tauri/src/commands/ai/knowledge_inject.rs
Normal file
557
src-tauri/src/commands/ai/knowledge_inject.rs
Normal file
@@ -0,0 +1,557 @@
|
||||
//! 知识库集成 — 注入 + 提炼(嵌入生成 / 混合检索 / 上下文构建 / 对话提炼)
|
||||
|
||||
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 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;
|
||||
}
|
||||
};
|
||||
Some((
|
||||
df_ai::build_provider(&rec.provider_type, &rec.base_url, &rec.api_key, &rec.default_model),
|
||||
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(|e| e.to_string())?;
|
||||
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> = df_ai::build_provider(
|
||||
&provider_config.provider_type,
|
||||
&provider_config.base_url,
|
||||
&provider_config.api_key,
|
||||
&provider_config.default_model,
|
||||
);
|
||||
// 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");
|
||||
}
|
||||
}
|
||||
146
src-tauri/src/commands/ai/mod.rs
Normal file
146
src-tauri/src/commands/ai/mod.rs
Normal file
@@ -0,0 +1,146 @@
|
||||
//! AI 聊天命令模块 — 流式对话、工具调用、审批门控、提供商管理
|
||||
//!
|
||||
//! 由原单文件 ai.rs(2663 行)按职责拆分为 11 个子模块。
|
||||
//!
|
||||
//! 模块布局:
|
||||
//! - [`commands`] — 所有 `#[tauri::command]` IPC 函数
|
||||
//! - [`agentic`] — run_agentic_loop / try_continue_agent_loop / MAX_AGENT_ITERATIONS
|
||||
//! - [`stream_recv`] — stream_llm 流式接收
|
||||
//! - [`conversation`] — save_conversation / TokenAccumulator / accumulate_tokens
|
||||
//! - [`title`] — 对话标题生成
|
||||
//! - [`audit`] — 工具调用审计 + pending 审批恢复 + 工具调用处理
|
||||
//! - [`skills`] — 本机 Claude 技能扫描
|
||||
//! - [`prompt`] — 系统提示词构建
|
||||
//! - [`tool_registry`] — AI 工具注册表构建 + 文件路径校验
|
||||
//! - [`knowledge_inject`] — 知识库注入 + 提炼
|
||||
//!
|
||||
//! 事件协议:通过 app.emit("ai-chat-event", payload) 流式推送到前端
|
||||
//! - AiTextDelta: 流式文本片段
|
||||
//! - AiToolCallStarted/Completed: 工具调用生命周期
|
||||
//! - AiApprovalRequired: 需要人工审批
|
||||
//! - AiCompleted/AiError: 完成/错误
|
||||
|
||||
pub mod agentic;
|
||||
pub mod audit;
|
||||
pub mod commands;
|
||||
pub mod conversation;
|
||||
pub mod knowledge_inject;
|
||||
pub mod prompt;
|
||||
pub mod skills;
|
||||
pub mod stream_recv;
|
||||
pub mod title;
|
||||
pub mod tool_registry;
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::AtomicBool;
|
||||
|
||||
use serde::Serialize;
|
||||
|
||||
use df_ai::ai_tools::RiskLevel;
|
||||
use df_ai::context::ContextManager;
|
||||
use df_ai::context::ContextConfig;
|
||||
|
||||
// ============================================================
|
||||
// 重导出 — 保路径不变(state.rs / lib.rs / knowledge.rs 在用)
|
||||
// ============================================================
|
||||
|
||||
// commands 子模块(glob 重导出) — `#[tauri::command]` 宏生成的命令函数 +
|
||||
// 内部符号(__cmd__xxx / __tauri_command_name_xxx)同源同模块定义,
|
||||
// glob 把它们全部拉到 commands::ai 路径,使 generate_handler! 能解析到。
|
||||
// 静默 unused_imports:glob 重导出用于跨模块路径解析,本文件不引用这些符号。
|
||||
#[allow(unused_imports)]
|
||||
pub use self::commands::*;
|
||||
|
||||
// 子模块的非命令 pub 项,逐个保路径
|
||||
#[allow(unused_imports)]
|
||||
pub use self::audit::restore_pending_approvals;
|
||||
#[allow(unused_imports)]
|
||||
pub use self::knowledge_inject::{spawn_embedding_for_knowledge, trigger_extraction_now};
|
||||
#[allow(unused_imports)]
|
||||
pub use self::tool_registry::build_ai_tool_registry;
|
||||
|
||||
// ============================================================
|
||||
// 事件载荷类型(放 mod.rs,各子文件经 use super::* 拿到)
|
||||
// ============================================================
|
||||
|
||||
/// AI 聊天事件(推送到前端)
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum AiChatEvent {
|
||||
/// 流式文本片段
|
||||
AiTextDelta { delta: String, conversation_id: Option<String> },
|
||||
/// 工具调用开始
|
||||
AiToolCallStarted { id: String, name: String, args: serde_json::Value, conversation_id: Option<String> },
|
||||
/// 工具调用完成
|
||||
AiToolCallCompleted { id: String, result: serde_json::Value, conversation_id: Option<String> },
|
||||
/// 需要人工审批
|
||||
AiApprovalRequired { id: String, name: String, args: serde_json::Value, reason: String, conversation_id: Option<String> },
|
||||
/// 审批结果
|
||||
AiApprovalResult { id: String, approved: bool, conversation_id: Option<String> },
|
||||
/// AI 响应完成
|
||||
AiCompleted { total_tokens: u32, prompt_tokens: u32, completion_tokens: u32, conversation_id: Option<String> },
|
||||
/// 错误
|
||||
AiError { error: String, conversation_id: Option<String> },
|
||||
/// Agent 循环新一轮(前端需新建 assistant 消息)
|
||||
AiAgentRound { round: u32, conversation_id: Option<String> },
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// 会话状态(放 mod.rs,各子文件经 use super::* 拿到)
|
||||
// ============================================================
|
||||
|
||||
/// AI 会话内状态(Mutex 保护)
|
||||
pub struct AiSession {
|
||||
/// 对话历史(ContextManager:唯一消息真相源,裁剪仅影响发送视图,不影响持久化)
|
||||
pub messages: ContextManager,
|
||||
/// 当前提供商 ID
|
||||
pub active_provider_id: Option<String>,
|
||||
/// 当前活跃对话 ID
|
||||
pub active_conversation_id: Option<String>,
|
||||
/// 活跃对话创建时间(懒创建:首条消息落库前仅存内存,upsert 时用作 created_at)
|
||||
pub active_conv_created_at: Option<String>,
|
||||
/// 挂起的审批(tool_call_id → 审批信息)
|
||||
pub pending_approvals: HashMap<String, PendingApproval>,
|
||||
/// 是否正在生成
|
||||
pub generating: bool,
|
||||
/// 当前 agent 循环的语言设置(用于审批后恢复循环)
|
||||
pub agent_language: Option<String>,
|
||||
/// 停止信号:ai_chat_stop 置位,agentic loop / stream_llm 检测后尽快退出
|
||||
pub stop_flag: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
impl AiSession {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
messages: ContextManager::new(ContextConfig::default()),
|
||||
active_provider_id: None,
|
||||
active_conversation_id: None,
|
||||
active_conv_created_at: None,
|
||||
pending_approvals: HashMap::new(),
|
||||
generating: false,
|
||||
agent_language: None,
|
||||
stop_flag: Arc::new(AtomicBool::new(false)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 待审批的工具调用
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct PendingApproval {
|
||||
pub tool_call_id: String,
|
||||
pub tool_name: String,
|
||||
pub arguments: serde_json::Value,
|
||||
pub risk_level: RiskLevel,
|
||||
pub conversation_id: Option<String>,
|
||||
/// 重启恢复的积压审批:无 live loop 持有 session.messages,审批后不 save(防空 messages 污染老对话)、不续跑
|
||||
pub recovered: bool,
|
||||
}
|
||||
|
||||
/// 工具调用草稿(流式收集时的临时结构)
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub(crate) struct ToolCallDraft {
|
||||
pub(crate) id: String,
|
||||
pub(crate) name: String,
|
||||
pub(crate) args: String,
|
||||
}
|
||||
86
src-tauri/src/commands/ai/prompt.rs
Normal file
86
src-tauri/src/commands/ai/prompt.rs
Normal file
@@ -0,0 +1,86 @@
|
||||
//! 系统提示词构建 + 活跃提供商获取
|
||||
|
||||
use df_storage::models::AiProviderRecord;
|
||||
|
||||
use crate::state::AppState;
|
||||
|
||||
/// 获取当前活跃提供商配置
|
||||
pub(crate) async fn get_active_provider(state: &AppState) -> Result<AiProviderRecord, String> {
|
||||
let session = state.ai_session.lock().await;
|
||||
if let Some(ref pid) = session.active_provider_id {
|
||||
let provider = state
|
||||
.ai_providers
|
||||
.get_by_id(pid)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("活跃提供商不存在: {}", pid))?;
|
||||
Ok(provider)
|
||||
} else {
|
||||
// 查找默认提供商
|
||||
drop(session);
|
||||
let providers = state
|
||||
.ai_providers
|
||||
.list_all()
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
let default = providers.iter().find(|p| p.is_default).cloned();
|
||||
default
|
||||
.or_else(|| providers.into_iter().next())
|
||||
.ok_or_else(|| "未配置 AI 提供商,请先在设置中添加".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// 按语言返回系统提示词的 (固定前缀, 项目上下文标题)
|
||||
fn system_prompt_parts(lang: &str) -> (&'static str, &'static str) {
|
||||
match lang {
|
||||
"en" => (
|
||||
"You are DevFlow's AI assistant. You help users manage projects, tasks, ideas, and workflows.\n\
|
||||
Please respond in English.\n\n\
|
||||
## Capabilities\n\
|
||||
You can perform the following actions via tool calls:\n\
|
||||
- Create/query projects, tasks, and ideas\n\
|
||||
- Run workflows\n\
|
||||
- Read file contents, list directories, create/write files\n\n\
|
||||
## Guidelines\n\
|
||||
- Briefly explain your intent before executing actions\n\
|
||||
- Ask for clarification if the user's intent is unclear\n\
|
||||
- Prefer using tools to complete actions rather than just describing steps\n\
|
||||
- When a tool call fails, clearly tell the user it failed and why. Never disguise a fallback action as the original intent's success (e.g. don't write to description to fake a directory binding), and never falsely report success\n",
|
||||
"\n## Current Projects\n",
|
||||
),
|
||||
_ => (
|
||||
"你是 DevFlow 的 AI 助手。你帮助用户管理项目、任务、想法和工作流。\n\
|
||||
必须使用简体中文回复,禁止使用繁体中文字符。\n\n\
|
||||
## 当前能力\n\
|
||||
你可以通过工具调用执行以下操作:\n\
|
||||
- 创建/查询项目、任务、想法\n\
|
||||
- 运行工作流\n\
|
||||
- 读取文件内容、列出目录、创建/写入文件\n\n\
|
||||
## 行为准则\n\
|
||||
- 执行操作前简要说明你的意图\n\
|
||||
- 如果不确定用户意图,先提问\n\
|
||||
- 优先使用工具完成操作,而不是只描述步骤\n\
|
||||
- 工具调用失败时必须明确告知用户失败原因,严禁用替代操作冒充原意图成功(如绑定目录失败不得改写描述冒充已绑定),也绝不谎报成功\n",
|
||||
"\n## 当前项目\n",
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
/// 构建系统提示词(固定前缀 + 当前项目上下文)
|
||||
pub(crate) async fn build_system_prompt(state: &AppState, lang: &str) -> String {
|
||||
let (prefix, ctx_label) = system_prompt_parts(lang);
|
||||
let mut prompt = String::from(prefix);
|
||||
|
||||
// 附加当前数据上下文
|
||||
if let Ok(projects) = state.projects.list_active().await {
|
||||
if !projects.is_empty() {
|
||||
prompt.push_str(ctx_label);
|
||||
// system prompt 前缀克制:仅最近 20 个项目,防 context 膨胀
|
||||
for p in projects.iter().take(20) {
|
||||
prompt.push_str(&format!("- {} ({}): {}\n", p.name, p.status, p.description));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
prompt
|
||||
}
|
||||
169
src-tauri/src/commands/ai/skills.rs
Normal file
169
src-tauri/src/commands/ai/skills.rs
Normal file
@@ -0,0 +1,169 @@
|
||||
//! 本机 Claude 技能扫描(skills / commands / plugins 三类)
|
||||
|
||||
use std::collections::HashSet;
|
||||
use std::fs;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::OnceLock;
|
||||
|
||||
use serde::Serialize;
|
||||
|
||||
/// 技能元信息(前端 `/` 联想 + 后端注入用)
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub struct SkillInfo {
|
||||
pub name: String,
|
||||
pub description: String,
|
||||
pub argument_hint: Option<String>,
|
||||
/// skill | command | plugin
|
||||
pub source: String,
|
||||
/// SKILL.md 绝对路径(注入时读全文)
|
||||
pub path: String,
|
||||
}
|
||||
|
||||
/// ~/.claude 目录(跨平台:USERPROFILE / HOME)
|
||||
fn claude_home() -> Option<PathBuf> {
|
||||
std::env::var_os("USERPROFILE")
|
||||
.or_else(|| std::env::var_os("HOME"))
|
||||
.map(PathBuf::from)
|
||||
.map(|h| h.join(".claude"))
|
||||
}
|
||||
|
||||
/// 剥离 YAML 标量值两侧的引号(`"..."` / `'...'`),简易 frontmatter 解析用
|
||||
fn unquote(s: &str) -> &str {
|
||||
s.strip_prefix('"')
|
||||
.and_then(|x| x.strip_suffix('"'))
|
||||
.or_else(|| s.strip_prefix('\'').and_then(|x| x.strip_suffix('\'')))
|
||||
.unwrap_or(s)
|
||||
}
|
||||
|
||||
/// 解析 markdown frontmatter 的 name / description / argument-hint / user_invocable
|
||||
/// (简易,按行匹配,容错缩进与 CRLF;仅扫描 frontmatter 区段)
|
||||
fn parse_frontmatter(md: &str) -> Option<(String, String, Option<String>, bool)> {
|
||||
let mut lines = md.lines();
|
||||
if lines.next()?.trim() != "---" {
|
||||
return None;
|
||||
}
|
||||
let mut name = None;
|
||||
let mut desc = None;
|
||||
let mut hint = None;
|
||||
let mut invocable = true;
|
||||
for line in lines {
|
||||
if line.trim() == "---" {
|
||||
break;
|
||||
}
|
||||
let l = line.trim_start();
|
||||
if let Some(v) = l.strip_prefix("name:") {
|
||||
name = Some(unquote(v.trim()).to_string());
|
||||
} else if let Some(v) = l.strip_prefix("description:") {
|
||||
desc = Some(unquote(v.trim()).to_string());
|
||||
} else if let Some(v) = l.strip_prefix("argument-hint:") {
|
||||
hint = Some(unquote(v.trim()).to_string());
|
||||
} else if let Some(v) = l.strip_prefix("user_invocable:") {
|
||||
invocable = v.trim() != "false";
|
||||
}
|
||||
}
|
||||
name.map(|n| (n, desc.unwrap_or_default(), hint, invocable))
|
||||
}
|
||||
|
||||
/// 解析单个 SKILL.md / command md 为 SkillInfo(排除 user_invocable: false)
|
||||
fn parse_skill_file(path: &Path, source: &str) -> Option<SkillInfo> {
|
||||
let md = fs::read_to_string(path).ok()?;
|
||||
let (name, description, argument_hint, invocable) = parse_frontmatter(&md).unwrap_or_else(|| {
|
||||
// 无 frontmatter(部分 commands):用文件名兜底,默认可调用
|
||||
let stem = path
|
||||
.file_stem()
|
||||
.map(|s| s.to_string_lossy().to_string())
|
||||
.unwrap_or_default();
|
||||
(stem, String::new(), None, true)
|
||||
});
|
||||
if !invocable {
|
||||
return None;
|
||||
}
|
||||
Some(SkillInfo {
|
||||
name,
|
||||
description,
|
||||
argument_hint,
|
||||
source: source.to_string(),
|
||||
path: path.to_string_lossy().to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
/// 递归收集目录下所有 SKILL.md(用于 plugins/marketplaces 多层嵌套)
|
||||
fn collect_skill_files(dir: &Path, out: &mut Vec<PathBuf>) {
|
||||
if let Ok(entries) = fs::read_dir(dir) {
|
||||
for entry in entries.flatten() {
|
||||
let p = entry.path();
|
||||
if p.is_dir() {
|
||||
// 跳过依赖/版本目录,避免递归爆炸
|
||||
let name = p.file_name().and_then(|n| n.to_str()).unwrap_or("");
|
||||
if name == "node_modules" || name == ".git" {
|
||||
continue;
|
||||
}
|
||||
collect_skill_files(&p, out);
|
||||
} else if p.file_name().and_then(|n| n.to_str()) == Some("SKILL.md") {
|
||||
out.push(p);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 扫描三类来源,按 name 去重(skills 优先 > commands > plugins)
|
||||
fn scan_skills() -> Vec<SkillInfo> {
|
||||
let home = match claude_home() {
|
||||
Some(h) => h,
|
||||
None => return Vec::new(),
|
||||
};
|
||||
let mut skills = Vec::new();
|
||||
let mut seen: HashSet<String> = HashSet::new();
|
||||
|
||||
// 1. ~/.claude/skills/*/SKILL.md
|
||||
if let Ok(entries) = fs::read_dir(home.join("skills")) {
|
||||
for entry in entries.flatten() {
|
||||
if let Some(info) = parse_skill_file(&entry.path().join("SKILL.md"), "skill") {
|
||||
if seen.insert(info.name.clone()) {
|
||||
skills.push(info);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2. ~/.claude/commands/*.md
|
||||
if let Ok(entries) = fs::read_dir(home.join("commands")) {
|
||||
for entry in entries.flatten() {
|
||||
let p = entry.path();
|
||||
if p.extension().and_then(|e| e.to_str()) == Some("md") {
|
||||
if let Some(info) = parse_skill_file(&p, "command") {
|
||||
if seen.insert(info.name.clone()) {
|
||||
skills.push(info);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3. ~/.claude/plugins/marketplaces/**/skills/*/SKILL.md(递归;cache 不在此路径下)
|
||||
let mut files = Vec::new();
|
||||
collect_skill_files(&home.join("plugins").join("marketplaces"), &mut files);
|
||||
for f in files {
|
||||
if let Some(info) = parse_skill_file(&f, "plugin") {
|
||||
if seen.insert(info.name.clone()) {
|
||||
skills.push(info);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
skills
|
||||
}
|
||||
|
||||
/// 技能扫描结果缓存(进程内;新增/改动技能需重启生效)
|
||||
pub(crate) fn skills_cached() -> &'static Vec<SkillInfo> {
|
||||
static SKILLS: OnceLock<Vec<SkillInfo>> = OnceLock::new();
|
||||
SKILLS.get_or_init(scan_skills)
|
||||
}
|
||||
|
||||
/// 按 name 读取技能全文(注入 system prompt);扫描走缓存,仅读单个 SKILL.md
|
||||
pub(crate) fn read_skill_content(name: &str) -> Option<String> {
|
||||
skills_cached()
|
||||
.iter()
|
||||
.find(|s| s.name == name)
|
||||
.and_then(|s| fs::read_to_string(&s.path).ok())
|
||||
}
|
||||
115
src-tauri/src/commands/ai/stream_recv.rs
Normal file
115
src-tauri/src/commands/ai/stream_recv.rs
Normal file
@@ -0,0 +1,115 @@
|
||||
//! 流式接收 LLM 响应
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::atomic::AtomicBool;
|
||||
use std::time::Duration;
|
||||
|
||||
use tauri::{AppHandle, Emitter};
|
||||
use futures::StreamExt;
|
||||
|
||||
use df_ai::provider::{CompletionRequest, LlmProvider};
|
||||
|
||||
use super::{AiChatEvent, ToolCallDraft};
|
||||
|
||||
/// 流式接收 LLM 响应,返回 (完整文本, 工具调用草稿)
|
||||
///
|
||||
/// 三类异常处理:
|
||||
/// - idle timeout(120s 无 chunk):判定连接静默断,emit AiError 返回 None
|
||||
/// - 流尽但从未收到 finished 信号:判定异常中断,emit AiError 返回 None(丢弃残缺,不当完整入库)
|
||||
/// - 用户停止(stop_flag):break 返回 Some(已收文本),由调用方入库展示后退出
|
||||
pub(crate) async fn stream_llm(
|
||||
provider: &dyn LlmProvider,
|
||||
request: CompletionRequest,
|
||||
app_handle: &AppHandle,
|
||||
stop_flag: &AtomicBool,
|
||||
conv_id: &str,
|
||||
) -> Option<(String, HashMap<u32, ToolCallDraft>, df_ai::provider::TokenUsage)> {
|
||||
/// 流式读取空闲超时:超过此时长无任何 chunk 即判定连接已断
|
||||
const STREAM_IDLE_TIMEOUT: Duration = Duration::from_secs(120);
|
||||
|
||||
match provider.stream(request).await {
|
||||
Ok(mut stream) => {
|
||||
let mut full_text = String::new();
|
||||
let mut tool_calls_acc: HashMap<u32, ToolCallDraft> = HashMap::new();
|
||||
let mut finished_received = false;
|
||||
let mut stopped = false;
|
||||
let mut final_usage: Option<df_ai::provider::TokenUsage> = None;
|
||||
|
||||
loop {
|
||||
// 用户主动停止:保留已收文本退出
|
||||
if stop_flag.load(std::sync::atomic::Ordering::SeqCst) {
|
||||
stopped = true;
|
||||
break;
|
||||
}
|
||||
|
||||
// idle timeout 防"连接存活但中途静默"无限 hang
|
||||
match tokio::time::timeout(STREAM_IDLE_TIMEOUT, stream.next()).await {
|
||||
Err(_elapsed) => {
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
|
||||
error: "流式响应超时(120 秒无数据,连接可能已断开)".to_string(),
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
});
|
||||
return None;
|
||||
}
|
||||
Ok(None) => break, // 流正常结束
|
||||
Ok(Some(chunk_result)) => match chunk_result {
|
||||
Ok(chunk) => {
|
||||
if !chunk.delta.is_empty() {
|
||||
full_text.push_str(&chunk.delta);
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiTextDelta {
|
||||
delta: chunk.delta,
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
});
|
||||
}
|
||||
if let Some(tc_deltas) = &chunk.tool_calls {
|
||||
for tc_delta in tc_deltas {
|
||||
let draft = tool_calls_acc.entry(tc_delta.index).or_default();
|
||||
if let Some(id) = &tc_delta.id { draft.id = id.clone(); }
|
||||
if let Some(name) = &tc_delta.function_name { draft.name.push_str(name); }
|
||||
if let Some(args) = &tc_delta.function_arguments { draft.args.push_str(args); }
|
||||
}
|
||||
}
|
||||
if let Some(u) = &chunk.usage {
|
||||
final_usage = Some(u.clone());
|
||||
}
|
||||
if chunk.finished {
|
||||
finished_received = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
|
||||
error: e.to_string(),
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
});
|
||||
return None;
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// 用户停止:已生成文本(可能残缺)交调用方入库展示
|
||||
if stopped {
|
||||
return Some((full_text, tool_calls_acc, final_usage.unwrap_or_default()));
|
||||
}
|
||||
|
||||
// 断连检测:流尽但从未收到 finished 信号 = 异常中断,丢弃残缺不当完整入库
|
||||
if !finished_received && (!full_text.is_empty() || !tool_calls_acc.is_empty()) {
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
|
||||
error: "流式响应意外中断(未收到完成信号,已丢弃残缺响应)".to_string(),
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
});
|
||||
return None;
|
||||
}
|
||||
|
||||
Some((full_text, tool_calls_acc, final_usage.unwrap_or_default()))
|
||||
}
|
||||
Err(e) => {
|
||||
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
|
||||
error: format!("AI 调用失败: {}", e),
|
||||
conversation_id: Some(conv_id.to_string()),
|
||||
});
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
140
src-tauri/src/commands/ai/title.rs
Normal file
140
src-tauri/src/commands/ai/title.rs
Normal file
@@ -0,0 +1,140 @@
|
||||
//! 对话标题生成
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use tauri::{AppHandle, Emitter};
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
use df_ai::provider::{ChatMessage, CompletionRequest, LlmProvider, MessageRole};
|
||||
use df_storage::crud::AiConversationRepo;
|
||||
use df_storage::db::Database;
|
||||
use df_storage::models::AiProviderRecord;
|
||||
|
||||
use crate::state::LlmConcurrency;
|
||||
|
||||
use super::AiSession;
|
||||
|
||||
/// 对话完成后按需生成智能标题(仅 title 为空时触发一次,不覆盖用户改名)
|
||||
///
|
||||
/// - 已有 title(用户改名或已生成)→ 跳过
|
||||
/// - 否则调 LLM 非流式总结生成 ≤15 字标题;LLM 失败回退 extract_title 截断
|
||||
/// - 完成后 emit ai-conversation-changed 通知前端侧栏刷新标题
|
||||
pub(crate) async fn ensure_conversation_title(
|
||||
provider_config: &AiProviderRecord,
|
||||
db: &Arc<Database>,
|
||||
conv_id: &str,
|
||||
app_handle: &AppHandle,
|
||||
session_arc: &Arc<Mutex<AiSession>>,
|
||||
llm_concurrency: LlmConcurrency,
|
||||
) {
|
||||
let conv_repo = AiConversationRepo::new(db);
|
||||
|
||||
// 已有标题(用户改名或已生成)→ 不覆盖
|
||||
if let Ok(Some(rec)) = conv_repo.get_by_id(conv_id).await {
|
||||
if rec.title.is_some() {
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
// 取对话文本(仅 user/assistant,跳过 tool 噪音),取前 6 条供 LLM 总结
|
||||
let (summary_msgs, all_msgs) = {
|
||||
let session = session_arc.lock().await;
|
||||
let summary: Vec<ChatMessage> = session.messages.iter()
|
||||
.filter(|m| matches!(m.role, MessageRole::User | MessageRole::Assistant))
|
||||
.take(6)
|
||||
.map(|m| ChatMessage {
|
||||
role: m.role.clone(),
|
||||
content: m.content.clone(),
|
||||
tool_call_id: None,
|
||||
tool_calls: None,
|
||||
model: None,
|
||||
})
|
||||
.collect();
|
||||
(summary, session.messages.all_messages_clone())
|
||||
};
|
||||
if summary_msgs.is_empty() {
|
||||
return;
|
||||
}
|
||||
|
||||
// 标题生成是独立一次 LLM 调用,自建 provider(便于后台 spawn,不借主 loop 的 &dyn LlmProvider)
|
||||
let provider: Box<dyn LlmProvider> = df_ai::build_provider(
|
||||
&provider_config.provider_type,
|
||||
&provider_config.base_url,
|
||||
&provider_config.api_key,
|
||||
&provider_config.default_model,
|
||||
);
|
||||
let title = match generate_title_via_llm(&*provider, &provider_config.default_model, summary_msgs, &llm_concurrency).await {
|
||||
Some(t) => t,
|
||||
None => extract_title(&all_msgs).unwrap_or_else(|| "新对话".to_string()),
|
||||
};
|
||||
|
||||
let _ = conv_repo.update_field(conv_id, "title", &title).await;
|
||||
let _ = app_handle.emit("ai-conversation-changed", ());
|
||||
}
|
||||
|
||||
/// 后台生成对话标题(不阻塞主流程;失败有 extract_title 兜底)
|
||||
pub(crate) fn spawn_ensure_title(
|
||||
provider_config: &AiProviderRecord,
|
||||
db: &Arc<Database>,
|
||||
conv_id: &str,
|
||||
app_handle: &AppHandle,
|
||||
session_arc: &Arc<Mutex<AiSession>>,
|
||||
llm_concurrency: &LlmConcurrency,
|
||||
) {
|
||||
let provider_config = provider_config.clone();
|
||||
let db = db.clone();
|
||||
let conv_id = conv_id.to_string();
|
||||
let app_handle = app_handle.clone();
|
||||
let session_arc = session_arc.clone();
|
||||
let llm_concurrency = llm_concurrency.clone();
|
||||
tauri::async_runtime::spawn(async move {
|
||||
ensure_conversation_title(&provider_config, &db, &conv_id, &app_handle, &session_arc, llm_concurrency).await;
|
||||
});
|
||||
}
|
||||
|
||||
/// 调 LLM 非流式生成对话标题
|
||||
async fn generate_title_via_llm(
|
||||
provider: &dyn LlmProvider,
|
||||
model: &str,
|
||||
msgs: Vec<ChatMessage>,
|
||||
llm_concurrency: &LlmConcurrency,
|
||||
) -> Option<String> {
|
||||
let mut prompt = vec![ChatMessage::system(
|
||||
"你是标题生成器。根据用户与助手的对话,生成一个简短中文标题。要求:不超过15字,纯文本,不加引号不加书名号不加标点不加 emoji,只输出标题本身,不要任何前缀或解释。"
|
||||
)];
|
||||
prompt.extend(msgs);
|
||||
let request = CompletionRequest {
|
||||
model: model.to_string(),
|
||||
messages: prompt,
|
||||
temperature: Some(0.3),
|
||||
max_tokens: Some(30),
|
||||
stream: false,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
};
|
||||
// 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.ok()?;
|
||||
Some(clean_title(&resp.text))
|
||||
}
|
||||
|
||||
/// 清理 LLM 返回的标题:去首尾引号/书名号/空白/末尾标点,截 15 字,空则兜底"新对话"
|
||||
fn clean_title(raw: &str) -> String {
|
||||
let t = raw.trim()
|
||||
.trim_matches(|c: char| matches!(c, '"' | '\'' | '「' | '」' | '《' | '》' | ' ' | '\n' | '\r'));
|
||||
let t = t.trim_end_matches(|c: char| matches!(c, '。' | '.' | ',' | ',' | '!' | '!' | '?' | '?' | ':' | ':'));
|
||||
let cleaned: String = t.chars().take(15).collect();
|
||||
if cleaned.is_empty() { "新对话".to_string() } else { cleaned }
|
||||
}
|
||||
|
||||
/// 从消息历史中提取对话标题(取第一条用户消息前 30 字)
|
||||
pub(crate) fn extract_title(messages: &[ChatMessage]) -> Option<String> {
|
||||
messages.iter()
|
||||
.find(|m| matches!(m.role, MessageRole::User))
|
||||
.map(|m| {
|
||||
let mut chars = m.content.chars();
|
||||
let t: String = chars.by_ref().take(30).collect();
|
||||
if chars.next().is_some() { format!("{}...", t) } else { t }
|
||||
})
|
||||
}
|
||||
394
src-tauri/src/commands/ai/tool_registry.rs
Normal file
394
src-tauri/src/commands/ai/tool_registry.rs
Normal file
@@ -0,0 +1,394 @@
|
||||
//! AI 工具注册表构建 + 文件路径校验
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::Arc;
|
||||
|
||||
use df_ai::ai_tools::{AiToolRegistry, RiskLevel};
|
||||
use df_storage::db::Database;
|
||||
use df_storage::models::{ProjectRecord, TaskRecord, IdeaRecord};
|
||||
|
||||
use df_core::types::new_id;
|
||||
|
||||
use crate::commands::now_millis;
|
||||
|
||||
/// 验证文件路径:禁止访问系统敏感目录
|
||||
fn validate_path(path: &str) -> anyhow::Result<()> {
|
||||
// 规范化为反斜杠:LLM 可能传正斜杠绕过黑名单(Windows tokio::fs 两种分隔符都吃)
|
||||
let normalized = path.replace('/', "\\");
|
||||
let lower = normalized.to_lowercase();
|
||||
if lower.contains("..") {
|
||||
anyhow::bail!("禁止路径遍历 (..)");
|
||||
}
|
||||
if lower.contains("\\.ssh")
|
||||
|| lower.contains("\\.aws")
|
||||
|| lower.contains("\\.gnupg")
|
||||
|| lower.contains("\\appdata\\")
|
||||
|| lower.contains("\\programdata\\")
|
||||
|| lower.contains("\\windows\\")
|
||||
|| lower.contains("\\system32\\")
|
||||
{
|
||||
anyhow::bail!("禁止访问敏感系统目录");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// workspace 根目录(项目根 = src-tauri 上两级,编译期固定)
|
||||
fn workspace_root() -> PathBuf {
|
||||
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
|
||||
.parent()
|
||||
.and_then(|p| p.parent())
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| PathBuf::from("."))
|
||||
}
|
||||
|
||||
/// 解析文件工具路径:相对路径锚定 workspace_root,禁止越出项目目录
|
||||
///
|
||||
/// 双层校验:
|
||||
/// 1. 词法层 starts_with(root)——对不存在路径(write_file 新建文件)兜底防越界
|
||||
/// 2. canonicalize 层——对存在路径解析符号链接,防 workspace 内 symlink 指向外部的逃逸
|
||||
/// 仅校验,返回词法 resolved(不含 \\?\ 前缀),保证 read_file 返回的 path 对前端友好
|
||||
fn resolve_workspace_path(path: &str) -> anyhow::Result<PathBuf> {
|
||||
validate_path(path)?;
|
||||
let root = workspace_root();
|
||||
let resolved = if Path::new(path).is_absolute() {
|
||||
PathBuf::from(path)
|
||||
} else {
|
||||
root.join(path)
|
||||
};
|
||||
// 词法层:防明显越界(不存在路径的兜底)
|
||||
if !resolved.starts_with(&root) {
|
||||
anyhow::bail!("禁止访问项目目录之外: {}", path);
|
||||
}
|
||||
// canonicalize 层:存在路径解析 symlink,防经符号链接逃逸出 workspace
|
||||
if resolved.exists() {
|
||||
let canon_root = root.canonicalize()?;
|
||||
let canon_resolved = resolved.canonicalize()?;
|
||||
if !canon_resolved.starts_with(&canon_root) {
|
||||
anyhow::bail!("禁止访问项目目录之外(符号链接逃逸): {}", path);
|
||||
}
|
||||
}
|
||||
Ok(resolved)
|
||||
}
|
||||
|
||||
/// 构建 AI 工具注册表 — handler 即唯一执行路径(schema+risk+实现同源,消除双轨)
|
||||
///
|
||||
/// CRUD 工具闭包捕获 `db` Arc 重建 Repo;文件系统工具复用 resolve_workspace_path /
|
||||
/// list_dir_recursive。新增工具只改这里一处,定义与实现同源,编译期保证一致。
|
||||
pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
|
||||
let mut registry = AiToolRegistry::new();
|
||||
|
||||
// ── 只读 (Low) ──
|
||||
registry.register(
|
||||
"list_projects", "列出所有项目,返回项目列表(ID、名称、状态、描述)",
|
||||
df_ai::ai_tools::object_schema(vec![]), RiskLevel::Low,
|
||||
{ let db = db.clone(); Box::new(move |_args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let repo = df_storage::crud::ProjectRepo::new(&db);
|
||||
let mut items = repo.list_all().await?;
|
||||
items.truncate(50); // 防 LLM context 膨胀
|
||||
Ok(serde_json::to_value(items)?)
|
||||
})
|
||||
})},
|
||||
);
|
||||
registry.register(
|
||||
"list_tasks", "列出任务,可按 project_id 筛选",
|
||||
df_ai::ai_tools::object_schema(vec![("project_id", "string", false)]), RiskLevel::Low,
|
||||
{ let db = db.clone(); Box::new(move |args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let repo = df_storage::crud::TaskRepo::new(&db);
|
||||
let mut tasks = if let Some(pid) = args.get("project_id").and_then(|v| v.as_str()) {
|
||||
repo.query("project_id", pid).await?
|
||||
} else {
|
||||
repo.list_all().await?
|
||||
};
|
||||
tasks.truncate(50); // 防 LLM context 膨胀
|
||||
Ok(serde_json::to_value(tasks)?)
|
||||
})
|
||||
})},
|
||||
);
|
||||
registry.register(
|
||||
"list_ideas", "列出所有想法",
|
||||
df_ai::ai_tools::object_schema(vec![]), RiskLevel::Low,
|
||||
{ let db = db.clone(); Box::new(move |_args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let repo = df_storage::crud::IdeaRepo::new(&db);
|
||||
let mut items = repo.list_all().await?;
|
||||
items.truncate(50); // 防 LLM context 膨胀
|
||||
Ok(serde_json::to_value(items)?)
|
||||
})
|
||||
})},
|
||||
);
|
||||
|
||||
// ── 创建 (Medium) ──
|
||||
registry.register(
|
||||
"update_project", "更新项目的指定字段(name/status/description/path/stack),需要提供项目 ID、字段名和新值。绑定代码目录推荐改用 bind_directory",
|
||||
df_ai::ai_tools::object_schema(vec![("id", "string", true), ("field", "string", true), ("value", "string", true)]),
|
||||
RiskLevel::Medium,
|
||||
{ let db = db.clone(); Box::new(move |args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let id = args["id"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 id"))?;
|
||||
let field = args["field"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 field"))?;
|
||||
let value = args["value"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 value"))?;
|
||||
match field { "name" | "status" | "description" | "path" | "stack" => {}, _ => anyhow::bail!("不允许更新字段 '{}'", field) }
|
||||
let repo = df_storage::crud::ProjectRepo::new(&db);
|
||||
repo.update_field(id, field, value).await?;
|
||||
Ok(serde_json::json!({ "id": id, "field": field, "updated": true }))
|
||||
})
|
||||
})},
|
||||
);
|
||||
registry.register(
|
||||
"create_project", "创建新项目",
|
||||
df_ai::ai_tools::object_schema(vec![("name", "string", true), ("description", "string", false)]),
|
||||
RiskLevel::Medium,
|
||||
{ let db = db.clone(); Box::new(move |args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let name = args["name"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 name 参数"))?;
|
||||
let description = args["description"].as_str().unwrap_or("");
|
||||
let repo = df_storage::crud::ProjectRepo::new(&db);
|
||||
let record = ProjectRecord {
|
||||
id: new_id(), name: name.to_string(), description: description.to_string(),
|
||||
status: "planning".to_string(), idea_id: None,
|
||||
path: None, stack: None,
|
||||
created_at: now_millis(), updated_at: now_millis(),
|
||||
};
|
||||
let id = record.id.clone();
|
||||
repo.insert(record).await?;
|
||||
Ok(serde_json::json!({ "id": id, "name": name, "status": "planning" }))
|
||||
})
|
||||
})},
|
||||
);
|
||||
registry.register(
|
||||
"bind_directory", "为项目绑定代码目录(自动探测技术栈,防重复绑定)",
|
||||
df_ai::ai_tools::object_schema(vec![("id", "string", true), ("path", "string", true)]),
|
||||
RiskLevel::Medium,
|
||||
{ let db = db.clone(); Box::new(move |args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let id = args["id"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 id 参数"))?;
|
||||
let path = args["path"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 path 参数"))?;
|
||||
let dir = std::path::Path::new(path);
|
||||
if !dir.is_dir() {
|
||||
anyhow::bail!("目录不存在: {path}");
|
||||
}
|
||||
let repo = df_storage::crud::ProjectRepo::new(&db);
|
||||
// 防重复:canonicalize 规范化比较,防路径写法差异绕过
|
||||
let normalize = |s: &str| -> String {
|
||||
std::path::Path::new(s)
|
||||
.canonicalize()
|
||||
.map(|a| a.to_string_lossy().replace('\\', "/").to_lowercase())
|
||||
.unwrap_or_else(|_| s.trim_end_matches(['\\', '/']).replace('\\', "/").to_lowercase())
|
||||
};
|
||||
let target = normalize(path);
|
||||
let projects = repo.list_active().await?;
|
||||
for proj in &projects {
|
||||
if proj.id != id {
|
||||
if let Some(pp) = &proj.path {
|
||||
if normalize(pp) == target {
|
||||
anyhow::bail!("目录已被项目「{}」绑定", proj.name);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// 探测技术栈
|
||||
let stack = df_project::scan::detect_stack(dir)?;
|
||||
let stack_json = serde_json::to_string(&stack)?;
|
||||
repo.update_field(id, "path", path).await?;
|
||||
repo.update_field(id, "stack", &stack_json).await?;
|
||||
Ok(serde_json::json!({ "id": id, "path": path, "stack": stack, "bound": true }))
|
||||
})
|
||||
})},
|
||||
);
|
||||
registry.register(
|
||||
"create_task", "在指定项目下创建新任务",
|
||||
df_ai::ai_tools::object_schema(vec![("project_id", "string", true), ("title", "string", true), ("description", "string", false), ("priority", "integer", false)]),
|
||||
RiskLevel::Medium,
|
||||
{ let db = db.clone(); Box::new(move |args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let project_id = args["project_id"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 project_id"))?;
|
||||
let title = args["title"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 title"))?;
|
||||
let repo = df_storage::crud::TaskRepo::new(&db);
|
||||
let record = TaskRecord {
|
||||
id: new_id(), project_id: project_id.to_string(), title: title.to_string(),
|
||||
description: args["description"].as_str().unwrap_or("").to_string(),
|
||||
status: "todo".to_string(), priority: args["priority"].as_i64().unwrap_or(2) as i32,
|
||||
branch_name: None, assignee: None, workflow_def_id: None, base_branch: None,
|
||||
created_at: now_millis(), updated_at: now_millis(),
|
||||
};
|
||||
let id = record.id.clone();
|
||||
repo.insert(record).await?;
|
||||
Ok(serde_json::json!({ "id": id, "title": title, "status": "todo" }))
|
||||
})
|
||||
})},
|
||||
);
|
||||
registry.register(
|
||||
"create_idea", "捕获一个新想法",
|
||||
df_ai::ai_tools::object_schema(vec![("title", "string", true), ("description", "string", false), ("tags", "string", false), ("source", "string", false)]),
|
||||
RiskLevel::Medium,
|
||||
{ let db = db.clone(); Box::new(move |args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let title = args["title"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 title"))?;
|
||||
let repo = df_storage::crud::IdeaRepo::new(&db);
|
||||
let record = IdeaRecord {
|
||||
id: new_id(), title: title.to_string(),
|
||||
description: args["description"].as_str().unwrap_or("").to_string(),
|
||||
status: "draft".to_string(), priority: args["priority"].as_i64().unwrap_or(1) as i32,
|
||||
score: None, tags: args["tags"].as_str().map(|s| s.to_string()),
|
||||
source: args["source"].as_str().map(|s| s.to_string()),
|
||||
promoted_to: None, ai_analysis: None, scores: None,
|
||||
created_at: now_millis(), updated_at: now_millis(),
|
||||
};
|
||||
let id = record.id.clone();
|
||||
repo.insert(record).await?;
|
||||
Ok(serde_json::json!({ "id": id, "title": title, "status": "draft" }))
|
||||
})
|
||||
})},
|
||||
);
|
||||
|
||||
// ── 高风险 (High) ──
|
||||
registry.register(
|
||||
"delete_project", "删除项目及其所有关联数据",
|
||||
df_ai::ai_tools::object_schema(vec![("id", "string", true)]), RiskLevel::High,
|
||||
{ let db = db.clone(); Box::new(move |args: serde_json::Value| {
|
||||
let db = db.clone();
|
||||
Box::pin(async move {
|
||||
let id = args["id"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 id"))?;
|
||||
let repo = df_storage::crud::ProjectRepo::new(&db);
|
||||
let deleted = repo.delete(id).await?;
|
||||
Ok(serde_json::json!({ "deleted": deleted, "id": id }))
|
||||
})
|
||||
})},
|
||||
);
|
||||
registry.register(
|
||||
"run_workflow", "运行指定的工作流 DAG",
|
||||
df_ai::ai_tools::object_schema(vec![("name", "string", true), ("dag", "object", true)]), RiskLevel::High,
|
||||
Box::new(|_args: serde_json::Value| Box::pin(async move {
|
||||
// run_workflow 需完整 DAG 执行,返回提示由前端触发
|
||||
Ok(serde_json::json!({ "note": "请通过工作流页面运行工作流", "tool": "run_workflow" }))
|
||||
})),
|
||||
);
|
||||
|
||||
// ── 文件系统 ──
|
||||
registry.register(
|
||||
"read_file", "读取文件内容,返回文本内容。支持 offset 和 limit 参数分页读取大文件",
|
||||
df_ai::ai_tools::object_schema(vec![("path", "string", true), ("offset", "integer", false), ("limit", "integer", false)]),
|
||||
RiskLevel::Low,
|
||||
Box::new(|args: serde_json::Value| Box::pin(async move {
|
||||
let resolved = resolve_workspace_path(
|
||||
args["path"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 path 参数"))?,
|
||||
)?;
|
||||
let path = resolved.to_str().ok_or_else(|| anyhow::anyhow!("路径含非法字符"))?;
|
||||
let metadata = tokio::fs::metadata(path).await
|
||||
.map_err(|e| anyhow::anyhow!("无法访问文件 {}: {}", path, e))?;
|
||||
if metadata.len() > 1_048_576 {
|
||||
anyhow::bail!("文件超过 1MB 限制 ({} 字节)", metadata.len());
|
||||
}
|
||||
let content = tokio::fs::read_to_string(path).await
|
||||
.map_err(|e| anyhow::anyhow!("读取文件失败: {}", e))?;
|
||||
let result = if let Some(offset) = args["offset"].as_u64() {
|
||||
let lines: Vec<&str> = content.lines().collect();
|
||||
let skip = offset as usize;
|
||||
let limit = args["limit"].as_u64().unwrap_or(200) as usize;
|
||||
lines.into_iter().skip(skip).take(limit).collect::<Vec<&str>>().join("\n")
|
||||
} else {
|
||||
content.clone()
|
||||
};
|
||||
let line_count = content.lines().count();
|
||||
Ok(serde_json::json!({ "path": path, "content": result, "size": metadata.len(), "lines": line_count }))
|
||||
})),
|
||||
);
|
||||
registry.register(
|
||||
"list_directory", "列出目录内容,返回文件和子目录列表(名称、类型、大小)",
|
||||
df_ai::ai_tools::object_schema(vec![("path", "string", true), ("recursive", "boolean", false), ("skip_noise_dirs", "boolean", false)]),
|
||||
RiskLevel::Low,
|
||||
Box::new(|args: serde_json::Value| Box::pin(async move {
|
||||
let resolved = resolve_workspace_path(
|
||||
args["path"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 path 参数"))?,
|
||||
)?;
|
||||
let path = resolved.to_str().ok_or_else(|| anyhow::anyhow!("路径含非法字符"))?;
|
||||
let recursive = args["recursive"].as_bool().unwrap_or(false);
|
||||
let skip_noise = args["skip_noise_dirs"].as_bool().unwrap_or(true);
|
||||
let mut entries = Vec::new();
|
||||
let truncated = list_dir_recursive(path, recursive, 0, 2, 1000, skip_noise, &mut entries).await?;
|
||||
Ok(serde_json::json!({ "path": path, "entries": entries, "truncated": truncated }))
|
||||
})),
|
||||
);
|
||||
registry.register(
|
||||
"write_file", "写入或创建文件,自动创建不存在的父目录",
|
||||
df_ai::ai_tools::object_schema(vec![("path", "string", true), ("content", "string", true)]),
|
||||
RiskLevel::Medium,
|
||||
Box::new(|args: serde_json::Value| Box::pin(async move {
|
||||
let resolved = resolve_workspace_path(
|
||||
args["path"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 path 参数"))?,
|
||||
)?;
|
||||
let path = resolved.to_str().ok_or_else(|| anyhow::anyhow!("路径含非法字符"))?;
|
||||
let content = args["content"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 content 参数"))?;
|
||||
if let Some(parent) = std::path::Path::new(path).parent() {
|
||||
tokio::fs::create_dir_all(parent).await
|
||||
.map_err(|e| anyhow::anyhow!("创建目录失败: {}", e))?;
|
||||
}
|
||||
tokio::fs::write(path, content).await
|
||||
.map_err(|e| anyhow::anyhow!("写入文件失败: {}", e))?;
|
||||
Ok(serde_json::json!({ "path": path, "bytes_written": content.len() }))
|
||||
})),
|
||||
);
|
||||
|
||||
registry
|
||||
}
|
||||
|
||||
/// 递归列出目录内容(最多 max_depth 层,最多 max_entries 条)
|
||||
///
|
||||
/// - 噪音目录(`.git`/`node_modules`/`target` 等)会被列出(显示存在),但不深入其内部
|
||||
/// - 达 `max_entries` 即停止,返回 `Ok(true)` 表示被截断
|
||||
fn list_dir_recursive<'a>(
|
||||
path: &'a str,
|
||||
recursive: bool,
|
||||
depth: usize,
|
||||
max_depth: usize,
|
||||
max_entries: usize,
|
||||
skip_noise: bool,
|
||||
result: &'a mut Vec<serde_json::Value>,
|
||||
) -> std::pin::Pin<Box<dyn std::future::Future<Output = anyhow::Result<bool>> + Send + 'a>> {
|
||||
Box::pin(async move {
|
||||
let mut dir = tokio::fs::read_dir(path).await
|
||||
.map_err(|e| anyhow::anyhow!("无法读取目录 {}: {}", path, e))?;
|
||||
while let Some(entry) = dir.next_entry().await? {
|
||||
if result.len() >= max_entries {
|
||||
return Ok(true); // 达上限截断
|
||||
}
|
||||
let name = entry.file_name().to_string_lossy().to_string();
|
||||
let metadata = entry.metadata().await?;
|
||||
let is_dir = metadata.is_dir();
|
||||
result.push(serde_json::json!({
|
||||
"name": name,
|
||||
"type": if is_dir { "directory" } else { "file" },
|
||||
"size": metadata.len(),
|
||||
"depth": depth,
|
||||
}));
|
||||
// 仅递归非噪音目录(skip_noise=true 时跳过 .git/node_modules/target 等)
|
||||
if recursive && is_dir && depth < max_depth && !(skip_noise && is_noise_dir(&name)) {
|
||||
let child_path = std::path::Path::new(path).join(&name).to_string_lossy().into_owned();
|
||||
let truncated = list_dir_recursive(&child_path, true, depth + 1, max_depth, max_entries, skip_noise, result).await?;
|
||||
if truncated {
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(false)
|
||||
})
|
||||
}
|
||||
|
||||
/// 判断是否为不应深入递归的噪音目录(构建产物/依赖/缓存等)
|
||||
fn is_noise_dir(name: &str) -> bool {
|
||||
const NOISE_DIRS: &[&str] = &[
|
||||
".git", "node_modules", "target", "dist", "build",
|
||||
".next", ".cache", "__pycache__", ".venv", "venv", ".idea",
|
||||
];
|
||||
NOISE_DIRS.contains(&name)
|
||||
}
|
||||
Reference in New Issue
Block a user