新增: 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:
@@ -14,12 +14,15 @@ tauri-build = { version = "2", features = [] }
|
||||
|
||||
[dependencies]
|
||||
tauri = { version = "2", features = [] }
|
||||
tauri-plugin-dialog = "2"
|
||||
tauri-plugin-opener = "2"
|
||||
tauri-plugin-window-state = "2"
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
anyhow.workspace = true
|
||||
tracing.workspace = true
|
||||
chrono.workspace = true
|
||||
|
||||
# 后端 crate
|
||||
df-core = { path = "../crates/df-core" }
|
||||
@@ -28,4 +31,6 @@ df-workflow = { path = "../crates/df-workflow" }
|
||||
df-nodes = { path = "../crates/df-nodes" }
|
||||
df-execute = { path = "../crates/df-execute" }
|
||||
df-ai = { path = "../crates/df-ai" }
|
||||
df-ideas = { path = "../crates/df-ideas" }
|
||||
df-project = { path = "../crates/df-project" }
|
||||
futures = "0.3"
|
||||
|
||||
@@ -15,6 +15,8 @@
|
||||
"core:window:allow-set-size",
|
||||
"core:window:allow-outer-position",
|
||||
"core:window:allow-inner-size",
|
||||
"core:webview:allow-create-webview-window"
|
||||
"core:webview:allow-create-webview-window",
|
||||
"dialog:default",
|
||||
"window-state:default"
|
||||
]
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
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)
|
||||
}
|
||||
@@ -3,8 +3,9 @@
|
||||
use serde::Deserialize;
|
||||
use tauri::State;
|
||||
|
||||
use df_core::types::new_id;
|
||||
use df_storage::models::IdeaRecord;
|
||||
use df_core::types::{new_id, Priority};
|
||||
use df_ideas::capture::Idea;
|
||||
use df_storage::models::{IdeaRecord, ProjectRecord};
|
||||
|
||||
use crate::state::AppState;
|
||||
|
||||
@@ -27,10 +28,16 @@ fn default_priority() -> i32 {
|
||||
1
|
||||
}
|
||||
|
||||
/// 列出全部想法
|
||||
/// 列出想法,可选按 status 过滤(指定状态走 query 走白名单索引列,否则全量)
|
||||
#[tauri::command]
|
||||
pub async fn list_ideas(state: State<'_, AppState>) -> Result<Vec<IdeaRecord>, String> {
|
||||
state.ideas.list_all().await.map_err(|e| e.to_string())
|
||||
pub async fn list_ideas(
|
||||
state: State<'_, AppState>,
|
||||
status: Option<String>,
|
||||
) -> Result<Vec<IdeaRecord>, String> {
|
||||
match status {
|
||||
Some(s) => state.ideas.query("status", &s).await.map_err(|e| e.to_string()),
|
||||
None => state.ideas.list_all().await.map_err(|e| e.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建想法,返回完整记录
|
||||
@@ -83,3 +90,216 @@ pub async fn update_idea(
|
||||
pub async fn delete_idea(state: State<'_, AppState>, id: String) -> Result<bool, String> {
|
||||
state.ideas.delete(&id).await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 将想法晋升为项目 — 复用 df-project 领域逻辑创建项目,回写想法 status=promoted/promoted_to
|
||||
#[tauri::command]
|
||||
pub async fn promote_idea(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
) -> Result<df_ideas::promotion::PromotionResult, String> {
|
||||
let record = state
|
||||
.ideas
|
||||
.get_by_id(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("想法不存在: {id}"))?;
|
||||
|
||||
if record.promoted_to.is_some() {
|
||||
return Err(format!("想法已立项: {}", record.promoted_to.unwrap()));
|
||||
}
|
||||
|
||||
// 复用 df-project 领域逻辑构造项目实体(create_from_idea)
|
||||
let project = df_project::manager::ProjectManager::create_from_idea(
|
||||
record.title.clone(),
|
||||
record.description.clone(),
|
||||
id.clone(),
|
||||
);
|
||||
let project_id = project.id.clone();
|
||||
let now = now_millis();
|
||||
let project_record = ProjectRecord {
|
||||
id: project_id.clone(),
|
||||
name: project.name,
|
||||
description: project.description,
|
||||
status: "planning".to_string(),
|
||||
idea_id: Some(id.clone()),
|
||||
path: None,
|
||||
stack: None,
|
||||
created_at: now.clone(),
|
||||
updated_at: now.clone(),
|
||||
};
|
||||
state
|
||||
.projects
|
||||
.insert(project_record)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// 回写想法:status=promoted + promoted_to(update_full 单事务覆盖可变字段)
|
||||
// 补偿删除:第二步失败时回滚第一步已建的 project,保证最终一致性(非原子,但防项目存留而
|
||||
// 想法状态未变的数据不一致)。Repository 方法各自持锁不支持跨 repo 共享事务对象,故选补偿
|
||||
// 删除而非真事务(改动最小,工程投入产出比最高)。
|
||||
let updated = IdeaRecord {
|
||||
status: "promoted".to_string(),
|
||||
promoted_to: Some(project_id.clone()),
|
||||
updated_at: now,
|
||||
..record
|
||||
};
|
||||
if let Err(e) = state.ideas.update_full(&updated).await {
|
||||
// 回写失败:补偿删除已建项目,避免悬空项目(idea.promoted_to 仍空,可重试立项)
|
||||
tracing::error!("想法 {id} 回写失败,补偿删除已建项目 {project_id}: {e}");
|
||||
if let Err(del_err) = state.projects.delete(&project_id).await {
|
||||
tracing::error!("补偿删除项目 {project_id} 也失败(需人工清理): {del_err}");
|
||||
}
|
||||
return Err(format!("想法立项回写失败(已回滚项目创建): {}", e));
|
||||
}
|
||||
|
||||
Ok(df_ideas::promotion::PromotionResult {
|
||||
idea_id: id,
|
||||
project_id: project_id,
|
||||
promoted: true,
|
||||
reason: "手动立项".to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// 想法评估 — 多维评分 + 对抗式评估
|
||||
// ============================================================
|
||||
|
||||
/// 评估想法:多维评分 + 对抗式评估,结果写回 scores/score/ai_analysis,状态置 pending_review,返回更新后的记录
|
||||
#[tauri::command]
|
||||
pub async fn evaluate_idea(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
) -> Result<IdeaRecord, String> {
|
||||
// 取出想法
|
||||
let record = state
|
||||
.ideas
|
||||
.get_by_id(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("想法不存在: {id}"))?;
|
||||
|
||||
let idea = record_to_idea(&record);
|
||||
|
||||
// 多维评分(0-10,IPC 层 *10 缩放为 0-100)
|
||||
let scores = df_ideas::scoring::ScoringEngine::compute_default(&idea);
|
||||
|
||||
// 对抗式评估
|
||||
let eval = df_ideas::adversarial::AdversarialEngine::evaluate(&idea)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// 组装前端扁平结构(与 Ideas.vue 的 AdversarialEval interface 对齐)
|
||||
let positive_strength = eval.positive.confidence;
|
||||
let negative_strength = eval.negative.confidence;
|
||||
let net_sentiment = positive_strength - negative_strength;
|
||||
let recommendation = recommendation_str(&eval.recommendation).to_string();
|
||||
let final_score = eval.final_score;
|
||||
let analyst_summary = eval.analyst.summary.clone();
|
||||
let action_items = action_items_for(&eval.recommendation);
|
||||
let positive = serde_json::json!({
|
||||
"thesis": eval.positive.thesis,
|
||||
"evidence": eval.positive.evidence,
|
||||
});
|
||||
let negative = serde_json::json!({
|
||||
"thesis": eval.negative.thesis,
|
||||
"evidence": eval.negative.evidence,
|
||||
});
|
||||
|
||||
let ai_analysis = serde_json::json!({
|
||||
"positive_strength": positive_strength,
|
||||
"negative_strength": negative_strength,
|
||||
"net_sentiment": net_sentiment,
|
||||
"recommendation": recommendation,
|
||||
"final_score": final_score,
|
||||
"summary": analyst_summary,
|
||||
"action_items": action_items,
|
||||
"positive": positive,
|
||||
"negative": negative,
|
||||
"analyst": { "summary": analyst_summary },
|
||||
})
|
||||
.to_string();
|
||||
|
||||
// scores JSON:中文维度 key + 0-100 值(前端雷达图直接当百分比用)
|
||||
let scores_json = serde_json::json!({
|
||||
"可行性": (scores.feasibility * 10.0).round() as i64,
|
||||
"影响力": (scores.impact * 10.0).round() as i64,
|
||||
"紧急度": (scores.urgency * 10.0).round() as i64,
|
||||
"综合": (scores.overall * 10.0).round() as i64,
|
||||
})
|
||||
.to_string();
|
||||
|
||||
let score_value = (scores.overall * 10.0).round() as i64;
|
||||
|
||||
// 构造完整记录后单次原子写回(update_full 保留 id 与 created_at)
|
||||
let updated = IdeaRecord {
|
||||
scores: Some(scores_json),
|
||||
ai_analysis: Some(ai_analysis),
|
||||
score: Some(score_value as f64),
|
||||
status: "pending_review".to_string(),
|
||||
updated_at: now_millis(),
|
||||
..record
|
||||
};
|
||||
state
|
||||
.ideas
|
||||
.update_full(&updated)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
/// IdeaRecord → df_ideas::Idea(评估用,status/time 不影响评分)
|
||||
fn record_to_idea(record: &IdeaRecord) -> Idea {
|
||||
let tags: Vec<String> = record
|
||||
.tags
|
||||
.as_deref()
|
||||
.and_then(|t| serde_json::from_str(t).ok())
|
||||
.unwrap_or_default();
|
||||
Idea {
|
||||
id: record.id.clone(),
|
||||
title: record.title.clone(),
|
||||
description: record.description.clone(),
|
||||
status: df_core::types::IdeaStatus::Draft,
|
||||
priority: priority_from_i32(record.priority),
|
||||
scores: None,
|
||||
tags,
|
||||
source: record.source.clone(),
|
||||
related_ids: Vec::new(),
|
||||
created_at: chrono::Utc::now(),
|
||||
updated_at: chrono::Utc::now(),
|
||||
}
|
||||
}
|
||||
|
||||
/// i32 优先级 → Priority 枚举(与 df-core 枚举值一致:Low=0/Medium=1/High=2/Critical=3)
|
||||
fn priority_from_i32(p: i32) -> Priority {
|
||||
match p {
|
||||
0 => Priority::Low,
|
||||
2 => Priority::High,
|
||||
x if x >= 3 => Priority::Critical,
|
||||
_ => Priority::Medium,
|
||||
}
|
||||
}
|
||||
|
||||
/// Recommendation → 前端 assessmentLabel 期望的全小写空格分隔(匹配 map key)
|
||||
fn recommendation_str(r: &df_ideas::adversarial::Recommendation) -> &'static str {
|
||||
use df_ideas::adversarial::Recommendation::*;
|
||||
match r {
|
||||
ImmediateAction => "immediate action",
|
||||
Soon => "soon",
|
||||
WithResources => "with resources",
|
||||
ResearchMore => "research more",
|
||||
Monitor => "monitor",
|
||||
}
|
||||
}
|
||||
|
||||
/// 行动建议 — 按推荐等级返回
|
||||
fn action_items_for(r: &df_ideas::adversarial::Recommendation) -> Vec<String> {
|
||||
use df_ideas::adversarial::Recommendation::*;
|
||||
match r {
|
||||
ImmediateAction => vec!["立即组建项目团队".into(), "制定详细执行计划".into(), "分配必要资源".into()],
|
||||
Soon => vec!["下周启动项目".into(), "准备资源需求".into(), "制定时间表".into()],
|
||||
WithResources => vec!["确认资源预算".into(), "评估 ROI".into(), "制定风险预案".into()],
|
||||
ResearchMore => vec!["进行市场调研".into(), "收集用户反馈".into(), "验证技术可行性".into()],
|
||||
Monitor => vec!["持续跟踪相关指标".into(), "定期评估进展".into(), "等待更好时机".into()],
|
||||
}
|
||||
}
|
||||
|
||||
450
src-tauri/src/commands/knowledge.rs
Normal file
450
src-tauri/src/commands/knowledge.rs
Normal file
@@ -0,0 +1,450 @@
|
||||
//! 知识库相关命令 — 共享记忆层(沉淀 / 检索注入 / 审核收件箱 / 配置)
|
||||
//!
|
||||
//! 11 个 command 对齐 MCP 语义(search/list/get/create/update_status/record_reuse 可对外暴露,
|
||||
//! archive/get_config/save_config/extract_now/list_candidates 为内部便利方法)。
|
||||
|
||||
use serde::Deserialize;
|
||||
use tauri::State;
|
||||
|
||||
use df_core::types::new_id;
|
||||
use df_storage::models::{KnowledgeEventRecord, KnowledgeRecord};
|
||||
|
||||
use crate::state::{AppState, KnowledgeConfig};
|
||||
|
||||
use super::knowledge_timeline::KnowledgeTimeline;
|
||||
use super::now_millis;
|
||||
use serde::Serialize;
|
||||
|
||||
/// 创建知识入参
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct CreateKnowledgeInput {
|
||||
/// 7 种 KnowledgeKind snake_case 之一
|
||||
pub kind: String,
|
||||
pub title: String,
|
||||
pub content: String,
|
||||
/// 标签 JSON 数组字符串
|
||||
#[serde(default)]
|
||||
pub tags: Option<String>,
|
||||
#[serde(default)]
|
||||
pub source_project: Option<String>,
|
||||
#[serde(default)]
|
||||
pub source_ref: Option<String>,
|
||||
/// high | medium | low
|
||||
#[serde(default)]
|
||||
pub confidence: Option<String>,
|
||||
}
|
||||
|
||||
/// 检索入参
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct KnowledgeSearchInput {
|
||||
pub query: String,
|
||||
#[serde(default)]
|
||||
pub kind: Option<String>,
|
||||
#[serde(default)]
|
||||
pub limit: Option<usize>,
|
||||
}
|
||||
|
||||
/// 状态转换合法矩阵校验
|
||||
///
|
||||
/// candidate → pending_review | published | archived
|
||||
/// pending_review → published | archived
|
||||
/// published → archived
|
||||
/// (其他组合非法)
|
||||
fn validate_transition(from: &str, to: &str) -> Result<(), String> {
|
||||
let legal = match from {
|
||||
"candidate" => matches!(to, "pending_review" | "published" | "archived"),
|
||||
"pending_review" => matches!(to, "published" | "archived"),
|
||||
"published" => matches!(to, "archived"),
|
||||
_ => false,
|
||||
};
|
||||
if legal {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(format!("非法状态转换: {from} → {to}"))
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// CRUD
|
||||
// ============================================================
|
||||
|
||||
/// 列出知识 — 可按 status 筛选,status=None 时默认排除 archived
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_list(
|
||||
state: State<'_, AppState>,
|
||||
status: Option<String>,
|
||||
) -> Result<Vec<KnowledgeRecord>, String> {
|
||||
match status {
|
||||
// 显式查 archived 时原样返回(含归档项)
|
||||
Some(s) if s == "archived" => state.knowledge.list_by_status("archived").await.map_err(|e| e.to_string()),
|
||||
Some(s) => state.knowledge.list_by_status(&s).await.map_err(|e| e.to_string()),
|
||||
// 默认:列出非 archived 的全部(单查询 status != 'archived')
|
||||
None => state.knowledge.list_non_archived().await.map_err(|e| e.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 单条查询
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_get(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
) -> Result<KnowledgeRecord, String> {
|
||||
state
|
||||
.knowledge
|
||||
.get_by_id(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("知识不存在: {id}"))
|
||||
}
|
||||
|
||||
/// 检索知识(top-N,默认 3) — MCP search tool
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_search(
|
||||
state: State<'_, AppState>,
|
||||
input: KnowledgeSearchInput,
|
||||
) -> Result<Vec<KnowledgeRecord>, String> {
|
||||
let limit = input.limit.unwrap_or(3).min(3);
|
||||
state
|
||||
.knowledge
|
||||
.search(&input.query, input.kind.as_deref(), limit)
|
||||
.await
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 创建知识(始终 candidate 状态) — MCP create tool
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_create(
|
||||
state: State<'_, AppState>,
|
||||
input: CreateKnowledgeInput,
|
||||
) -> Result<KnowledgeRecord, String> {
|
||||
let now = now_millis();
|
||||
let record = KnowledgeRecord {
|
||||
id: new_id(),
|
||||
kind: input.kind,
|
||||
title: input.title,
|
||||
content: input.content,
|
||||
tags: input.tags,
|
||||
status: "candidate".to_string(),
|
||||
confidence: input.confidence,
|
||||
reuse_count: 0,
|
||||
verified: false,
|
||||
source_project: input.source_project,
|
||||
source_ref: input.source_ref,
|
||||
reasoning: None,
|
||||
created_at: now.clone(),
|
||||
updated_at: now,
|
||||
};
|
||||
state
|
||||
.knowledge
|
||||
.insert(record.clone())
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
// 生命线:手动录入产生(fire-and-forget)
|
||||
KnowledgeTimeline::new(&state.db)
|
||||
.record_created(&record.id, "manual")
|
||||
.await;
|
||||
Ok(record)
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// 状态机 + 计数
|
||||
// ============================================================
|
||||
|
||||
/// 状态转换(含合法矩阵校验) — MCP update_status tool
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_update_status(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
status: String,
|
||||
) -> Result<bool, String> {
|
||||
let record = state
|
||||
.knowledge
|
||||
.get_by_id(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("知识不存在: {id}"))?;
|
||||
|
||||
let old_status = record.status.clone();
|
||||
validate_transition(&record.status, &status)?;
|
||||
|
||||
// published 时一次性标 verified=true(发布审核标)
|
||||
let now = now_millis();
|
||||
let updated = KnowledgeRecord {
|
||||
status: status.clone(),
|
||||
verified: if status == "published" { true } else { record.verified },
|
||||
updated_at: now,
|
||||
..record
|
||||
};
|
||||
let result = state
|
||||
.knowledge
|
||||
.update_full(&updated)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
|
||||
// 生命线:状态变更(fire-and-forget;to=published=审核通过,to=archived=归档)
|
||||
KnowledgeTimeline::new(&state.db)
|
||||
.record_status_change(&id, &old_status, &status)
|
||||
.await;
|
||||
|
||||
// 发布时后台生成嵌入(只有 published 参与检索,candidate 不浪费 embed 调用;
|
||||
// vector_enabled 关闭或 embed 失败时该条走 LIKE 降级,非阻断)
|
||||
if status == "published" {
|
||||
crate::commands::ai::spawn_embedding_for_knowledge(&state, &updated).await;
|
||||
}
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
/// 复用计数 +1(检索命中时调用) — MCP record_reuse tool
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_record_reuse(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
) -> Result<bool, String> {
|
||||
state
|
||||
.knowledge
|
||||
.increment_reuse_count(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 审核收件箱 — 列出 candidate(按 confidence 语义排序)
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_list_candidates(
|
||||
state: State<'_, AppState>,
|
||||
) -> Result<Vec<KnowledgeRecord>, String> {
|
||||
state
|
||||
.knowledge
|
||||
.list_by_status("candidate")
|
||||
.await
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 归档(软删除) — UPDATE status='archived'
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_archive(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
) -> Result<bool, String> {
|
||||
knowledge_update_status(state, id, "archived".to_string()).await
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// 配置 + 手动提炼
|
||||
// ============================================================
|
||||
|
||||
/// 读取知识库行为配置(提取 + 注入)
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_get_config(
|
||||
state: State<'_, AppState>,
|
||||
) -> Result<KnowledgeConfig, String> {
|
||||
let cfg = state.knowledge_config.lock().await;
|
||||
Ok(cfg.clone())
|
||||
}
|
||||
|
||||
/// 保存知识库行为配置
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_save_config(
|
||||
state: State<'_, AppState>,
|
||||
config: KnowledgeConfig,
|
||||
) -> Result<bool, String> {
|
||||
let mut cfg = state.knowledge_config.lock().await;
|
||||
*cfg = config;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
/// 手动触发提炼(ManualOnly 模式 / 用户主动点按钮)
|
||||
///
|
||||
/// 实际提炼逻辑在 commands::ai 模块(需访问 active provider + conversation messages)。
|
||||
/// 此 command 仅作前端入口,委托给 ai 模块的提炼函数。
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_extract_now(
|
||||
state: State<'_, AppState>,
|
||||
) -> Result<bool, String> {
|
||||
crate::commands::ai::trigger_extraction_now(&state).await
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// 生命线:详情 / 编辑 / 事件查询
|
||||
// ============================================================
|
||||
|
||||
/// 编辑知识入参(部分更新,仅传需改字段)
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct UpdateKnowledgeInput {
|
||||
#[serde(default)]
|
||||
pub title: Option<String>,
|
||||
#[serde(default)]
|
||||
pub content: Option<String>,
|
||||
#[serde(default)]
|
||||
pub tags: Option<String>,
|
||||
#[serde(default)]
|
||||
pub confidence: Option<String>,
|
||||
#[serde(default)]
|
||||
pub reasoning: Option<String>,
|
||||
}
|
||||
|
||||
/// 知识详情聚合负载(基本信息 + 全部生命线事件)
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct KnowledgeDetailPayload {
|
||||
pub knowledge: KnowledgeRecord,
|
||||
pub events: Vec<KnowledgeEventRecord>,
|
||||
}
|
||||
|
||||
/// 知识详情(基本信息 + 生命线事件) — 详情页一次拉全
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_get_detail(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
) -> Result<KnowledgeDetailPayload, String> {
|
||||
let knowledge = state
|
||||
.knowledge
|
||||
.get_by_id(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("知识不存在: {id}"))?;
|
||||
let events = state
|
||||
.knowledge_events
|
||||
.list_by_knowledge(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
Ok(KnowledgeDetailPayload { knowledge, events })
|
||||
}
|
||||
|
||||
/// 编辑知识(部分更新 title/content/tags/confidence/reasoning)
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_update(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
input: UpdateKnowledgeInput,
|
||||
) -> Result<KnowledgeRecord, String> {
|
||||
let mut record = state
|
||||
.knowledge
|
||||
.get_by_id(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("知识不存在: {id}"))?;
|
||||
if let Some(v) = input.title {
|
||||
record.title = v;
|
||||
}
|
||||
if let Some(v) = input.content {
|
||||
record.content = v;
|
||||
}
|
||||
if let Some(v) = input.tags {
|
||||
record.tags = Some(v);
|
||||
}
|
||||
// 空串 = 清空(前端编辑器清空 confidence/reasoning 输入框时传 "")
|
||||
if let Some(v) = input.confidence {
|
||||
record.confidence = if v.is_empty() { None } else { Some(v) };
|
||||
}
|
||||
if let Some(v) = input.reasoning {
|
||||
record.reasoning = if v.is_empty() { None } else { Some(v) };
|
||||
}
|
||||
record.updated_at = now_millis();
|
||||
state
|
||||
.knowledge
|
||||
.update_full(&record)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
Ok(record)
|
||||
}
|
||||
|
||||
/// 查询生命线事件(可按 event_type 过滤 + limit)
|
||||
#[tauri::command]
|
||||
pub async fn knowledge_events(
|
||||
state: State<'_, AppState>,
|
||||
knowledge_id: String,
|
||||
event_type: Option<String>,
|
||||
limit: Option<usize>,
|
||||
) -> Result<Vec<KnowledgeEventRecord>, String> {
|
||||
match event_type {
|
||||
Some(et) => {
|
||||
let limit = limit.unwrap_or(50);
|
||||
state
|
||||
.knowledge_events
|
||||
.list_by_knowledge_type(&knowledge_id, &et, limit)
|
||||
.await
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
None => state
|
||||
.knowledge_events
|
||||
.list_by_knowledge(&knowledge_id)
|
||||
.await
|
||||
.map_err(|e| e.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::validate_transition;
|
||||
|
||||
#[test]
|
||||
fn candidate_legal_transitions() {
|
||||
assert!(validate_transition("candidate", "pending_review").is_ok());
|
||||
assert!(validate_transition("candidate", "published").is_ok());
|
||||
assert!(validate_transition("candidate", "archived").is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn candidate_illegal_transitions() {
|
||||
// 禁止自转、回退
|
||||
assert!(validate_transition("candidate", "candidate").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pending_review_legal_transitions() {
|
||||
assert!(validate_transition("pending_review", "published").is_ok());
|
||||
assert!(validate_transition("pending_review", "archived").is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pending_review_illegal_transitions() {
|
||||
// 已进入审核态不能再退回 candidate
|
||||
assert!(validate_transition("pending_review", "candidate").is_err());
|
||||
assert!(validate_transition("pending_review", "pending_review").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn published_legal_transitions() {
|
||||
// 已发布只能归档
|
||||
assert!(validate_transition("published", "archived").is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn published_illegal_transitions() {
|
||||
// 发布态不可再进审核、不可回 candidate、不可自转、不可重复 published
|
||||
assert!(validate_transition("published", "pending_review").is_err());
|
||||
assert!(validate_transition("published", "candidate").is_err());
|
||||
assert!(validate_transition("published", "published").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn archived_is_terminal() {
|
||||
// 归档为终态,任何转换都非法
|
||||
for to in ["candidate", "pending_review", "published", "archived"] {
|
||||
assert!(validate_transition("archived", to).is_err(),
|
||||
"archived → {to} 应为非法");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_from_is_illegal() {
|
||||
// 未知源状态一律拒绝
|
||||
assert!(validate_transition("unknown", "published").is_err());
|
||||
assert!(validate_transition("", "archived").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn illegal_target_is_rejected() {
|
||||
// 合法源 → 未知/空目标一律拒绝
|
||||
assert!(validate_transition("candidate", "unknown").is_err());
|
||||
assert!(validate_transition("candidate", "").is_err());
|
||||
assert!(validate_transition("pending_review", "draft").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_message_contains_transition() {
|
||||
let err = validate_transition("archived", "candidate").unwrap_err();
|
||||
assert!(err.contains("archived"), "错误信息应包含源状态: {err}");
|
||||
assert!(err.contains("candidate"), "错误信息应包含目标状态: {err}");
|
||||
}
|
||||
}
|
||||
123
src-tauri/src/commands/knowledge_timeline.rs
Normal file
123
src-tauri/src/commands/knowledge_timeline.rs
Normal file
@@ -0,0 +1,123 @@
|
||||
//! 知识生命线记录器 — 统一入口,各业务流程只调一行
|
||||
//!
|
||||
//! 设计目标:
|
||||
//! - 事件记录逻辑不散落在各业务流程内部(提取/注入/状态变更),集中于此模块
|
||||
//! - fire-and-forget 友好:便捷方法内部吞掉错误(只 warn),调用方无需处理 Result
|
||||
//! - 易扩展:新事件类型只加一个便捷方法 + 一行 context 构造,不改主流程
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use df_core::types::new_id;
|
||||
use df_storage::crud::KnowledgeEventsRepo;
|
||||
use df_storage::db::Database;
|
||||
use df_storage::models::KnowledgeEventRecord;
|
||||
|
||||
use super::now_millis;
|
||||
|
||||
/// 事件类型常量(生命线节点;归档复用 status_changed 的 to=archived,不单列)
|
||||
pub const EVENT_CREATED: &str = "created";
|
||||
pub const EVENT_EXTRACTED: &str = "extracted";
|
||||
pub const EVENT_STATUS_CHANGED: &str = "status_changed";
|
||||
pub const EVENT_REFERENCED: &str = "referenced";
|
||||
|
||||
/// 知识生命线记录器 — 持有 DB 句柄,提供便捷事件写入方法
|
||||
///
|
||||
/// 用法:`KnowledgeTimeline::new(&state.db).record_xxx(...).await`
|
||||
///
|
||||
/// 异步策略(按路径热度,非 bug,刻意不一致):
|
||||
/// - 热路径(对话注入 build_knowledge_context / 提炼 extract):整块已在 spawn 任务内,
|
||||
/// 此处直接 await = 对用户 fire-and-forget,不阻塞 AI 对话
|
||||
/// - 低频命令(状态变更 / 创建):事件写入 ms 级,直接 await 比 detached spawn 更简单
|
||||
/// 便捷方法内部 `fire()` 已吞错误(warn 不阻断),调用方无需处理 Result。
|
||||
pub struct KnowledgeTimeline {
|
||||
db: Arc<Database>,
|
||||
}
|
||||
|
||||
impl KnowledgeTimeline {
|
||||
pub fn new(db: &Arc<Database>) -> Self {
|
||||
Self { db: db.clone() }
|
||||
}
|
||||
|
||||
/// 通用写入(底层入口,便捷方法均经此)
|
||||
pub async fn record(
|
||||
&self,
|
||||
knowledge_id: &str,
|
||||
event_type: &str,
|
||||
source_ref: Option<String>,
|
||||
context_json: Option<String>,
|
||||
) -> Result<(), String> {
|
||||
let record = KnowledgeEventRecord {
|
||||
id: new_id(),
|
||||
knowledge_id: knowledge_id.to_string(),
|
||||
event_type: event_type.to_string(),
|
||||
source_ref,
|
||||
context_json,
|
||||
timestamp: now_millis(),
|
||||
};
|
||||
KnowledgeEventsRepo::new(&self.db)
|
||||
.insert(record)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 手动录入产生(context: method)
|
||||
pub async fn record_created(&self, knowledge_id: &str, method: &str) {
|
||||
let ctx = serde_json::json!({ "method": method }).to_string();
|
||||
self.fire(knowledge_id, EVENT_CREATED, Some("manual".into()), Some(ctx)).await;
|
||||
}
|
||||
|
||||
/// AI 对话提炼产生(context: conv_id / conv_title / reasoning)
|
||||
pub async fn record_extracted(
|
||||
&self,
|
||||
knowledge_id: &str,
|
||||
conv_id: &str,
|
||||
conv_title: &str,
|
||||
reasoning: &str,
|
||||
) {
|
||||
let ctx = serde_json::json!({
|
||||
"conv_id": conv_id,
|
||||
"conv_title": conv_title,
|
||||
"reasoning": reasoning,
|
||||
})
|
||||
.to_string();
|
||||
self.fire(
|
||||
knowledge_id,
|
||||
EVENT_EXTRACTED,
|
||||
Some(format!("conv:{}", conv_id)),
|
||||
Some(ctx),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
/// 被检索命中/注入(context: conv_id / query)
|
||||
pub async fn record_referenced(&self, knowledge_id: &str, conv_id: &str, query: &str) {
|
||||
let ctx = serde_json::json!({ "conv_id": conv_id, "query": query }).to_string();
|
||||
self.fire(
|
||||
knowledge_id,
|
||||
EVENT_REFERENCED,
|
||||
Some(format!("conv:{}", conv_id)),
|
||||
Some(ctx),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
/// 状态变更(context: from / to;to=published=审核通过,to=archived=归档)
|
||||
pub async fn record_status_change(&self, knowledge_id: &str, from: &str, to: &str) {
|
||||
let ctx = serde_json::json!({ "from": from, "to": to }).to_string();
|
||||
self.fire(knowledge_id, EVENT_STATUS_CHANGED, None, Some(ctx)).await;
|
||||
}
|
||||
|
||||
/// 内部:写入并吞错误(fire-and-forget,失败只 warn 不阻断主流程)
|
||||
async fn fire(
|
||||
&self,
|
||||
knowledge_id: &str,
|
||||
event_type: &str,
|
||||
source_ref: Option<String>,
|
||||
context_json: Option<String>,
|
||||
) {
|
||||
if let Err(e) = self.record(knowledge_id, event_type, source_ref, context_json).await {
|
||||
tracing::warn!("知识生命线事件写入失败(非阻断) [{}]: {}", event_type, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -4,7 +4,10 @@
|
||||
|
||||
pub mod ai;
|
||||
pub mod idea;
|
||||
pub mod knowledge;
|
||||
pub mod knowledge_timeline;
|
||||
pub mod project;
|
||||
pub mod settings;
|
||||
pub mod task;
|
||||
pub mod workflow;
|
||||
|
||||
|
||||
@@ -1,9 +1,14 @@
|
||||
//! 项目相关命令
|
||||
|
||||
use serde::Deserialize;
|
||||
use std::path::Path;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tauri::State;
|
||||
|
||||
use df_ai::build_provider;
|
||||
use df_ai::provider::{ChatMessage, CompletionRequest};
|
||||
use df_core::types::new_id;
|
||||
use df_project::scan::{collect_sample, detect_stack};
|
||||
use df_storage::models::ProjectRecord;
|
||||
|
||||
use crate::state::AppState;
|
||||
@@ -17,20 +22,50 @@ pub struct CreateProjectInput {
|
||||
#[serde(default)]
|
||||
pub description: String,
|
||||
pub idea_id: Option<String>,
|
||||
/// 绑定的本地代码目录(可选,空=不绑定)
|
||||
#[serde(default)]
|
||||
pub path: Option<String>,
|
||||
/// 技术栈 JSON 数组字符串(可选,空则自动探测)
|
||||
#[serde(default)]
|
||||
pub stack: Option<String>,
|
||||
}
|
||||
|
||||
/// 列出全部项目
|
||||
/// 列出未删除项目(过滤回收站)
|
||||
#[tauri::command]
|
||||
pub async fn list_projects(state: State<'_, AppState>) -> Result<Vec<ProjectRecord>, String> {
|
||||
state.projects.list_all().await.map_err(|e| e.to_string())
|
||||
state.projects.list_active().await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 创建项目,返回完整记录
|
||||
///
|
||||
/// 绑定目录时(path 非空):校验目录存在 + 防重复绑定 + 自动探测技术栈(stack 为空时)。
|
||||
#[tauri::command]
|
||||
pub async fn create_project(
|
||||
state: State<'_, AppState>,
|
||||
input: CreateProjectInput,
|
||||
) -> Result<ProjectRecord, String> {
|
||||
// 绑定目录:校验存在 + 防重复 + 自动探测技术栈
|
||||
let (path, stack) = match input.path.as_deref().map(str::trim).filter(|p| !p.is_empty()) {
|
||||
Some(p) => {
|
||||
if !Path::new(p).is_dir() {
|
||||
return Err(format!("目录不存在: {p}"));
|
||||
}
|
||||
if let Some(conflict) = find_binding_conflict(&state, p, None).await? {
|
||||
return Err(format!("目录已被项目「{}」绑定", conflict.name));
|
||||
}
|
||||
// stack 优先用入参,否则自动探测
|
||||
let stack_json = match input.stack.as_deref().map(str::trim).filter(|s| !s.is_empty()) {
|
||||
Some(s) => s.to_string(),
|
||||
None => {
|
||||
let detected = detect_stack(Path::new(p)).map_err(|e| e.to_string())?;
|
||||
serde_json::to_string(&detected).map_err(|e| e.to_string())?
|
||||
}
|
||||
};
|
||||
(Some(p.to_string()), Some(stack_json))
|
||||
}
|
||||
None => (None, None),
|
||||
};
|
||||
|
||||
let now = now_millis();
|
||||
let record = ProjectRecord {
|
||||
id: new_id(),
|
||||
@@ -38,6 +73,8 @@ pub async fn create_project(
|
||||
description: input.description,
|
||||
status: "planning".to_string(),
|
||||
idea_id: input.idea_id,
|
||||
path,
|
||||
stack,
|
||||
created_at: now.clone(),
|
||||
updated_at: now,
|
||||
};
|
||||
@@ -77,8 +114,262 @@ pub async fn update_project(
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 删除项目
|
||||
/// 删除项目(软删 → 回收站,可恢复)
|
||||
#[tauri::command]
|
||||
pub async fn delete_project(state: State<'_, AppState>, id: String) -> Result<bool, String> {
|
||||
state.projects.delete(&id).await.map_err(|e| e.to_string())
|
||||
state.projects.soft_delete(&id).await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 列出回收站项目(deleted_at IS NOT NULL)
|
||||
#[tauri::command]
|
||||
pub async fn list_deleted_projects(
|
||||
state: State<'_, AppState>,
|
||||
) -> Result<Vec<ProjectRecord>, String> {
|
||||
state.projects.list_deleted().await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 恢复项目(从回收站还原,清 deleted_at)
|
||||
#[tauri::command]
|
||||
pub async fn restore_project(state: State<'_, AppState>, id: String) -> Result<bool, String> {
|
||||
state.projects.restore(&id).await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 彻底删除项目(级联物理删 branches/releases/tasks,不可恢复)
|
||||
#[tauri::command]
|
||||
pub async fn purge_project(state: State<'_, AppState>, id: String) -> Result<bool, String> {
|
||||
state
|
||||
.projects
|
||||
.purge_with_descendants(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// 项目目录绑定 — 探测 / 防重复 / 重定位 / 有效性检查
|
||||
// ============================================================
|
||||
|
||||
/// 规范化路径用于比较:canonicalize 解析绝对规范路径(失败降级),
|
||||
/// 统一正斜杠 + 小写。防 `C:\a\b` vs `C:/a/b/` 绕过重复检查。
|
||||
/// 注:仅用于比较,存库保留用户输入的原始可读路径。
|
||||
fn normalize_path(p: &str) -> String {
|
||||
match Path::new(p).canonicalize() {
|
||||
Ok(abs) => abs.to_string_lossy().replace('\\', "/").to_lowercase(),
|
||||
Err(_) => p
|
||||
.trim_end_matches(['\\', '/'])
|
||||
.replace('\\', "/")
|
||||
.to_lowercase(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 查找已绑定该目录的项目(排除 exclude_id 自身)。无冲突返回 None。
|
||||
async fn find_binding_conflict(
|
||||
state: &AppState,
|
||||
path: &str,
|
||||
exclude_id: Option<&str>,
|
||||
) -> Result<Option<ProjectRecord>, String> {
|
||||
let norm = normalize_path(path);
|
||||
let projects = state.projects.list_active().await.map_err(|e| e.to_string())?;
|
||||
Ok(projects.into_iter().find(|p| {
|
||||
let excluded = exclude_id.is_some_and(|eid| p.id == eid);
|
||||
!excluded && p.path.as_ref().is_some_and(|pp| normalize_path(pp) == norm)
|
||||
}))
|
||||
}
|
||||
|
||||
/// 探测目录技术栈(前端选目录后实时预览)
|
||||
#[tauri::command]
|
||||
pub async fn scan_project_stack(path: String) -> Result<Vec<String>, String> {
|
||||
let root = std::path::PathBuf::from(&path);
|
||||
tokio::task::spawn_blocking(move || detect_stack(&root).map_err(|e| e.to_string()))
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
}
|
||||
|
||||
/// 检查目录是否已被其他项目绑定(防重复绑定)。返回占用项目(若有)。
|
||||
/// exclude_id 用于编辑/重定位时排除自身。
|
||||
#[tauri::command]
|
||||
pub async fn check_path_binding(
|
||||
state: State<'_, AppState>,
|
||||
path: String,
|
||||
exclude_id: Option<String>,
|
||||
) -> Result<Option<ProjectRecord>, String> {
|
||||
find_binding_conflict(&state, &path, exclude_id.as_deref()).await
|
||||
}
|
||||
|
||||
/// 重定位项目目录(目录移动后重新指向)。校验存在 + 防重复 + 重探测 stack,返回最新记录。
|
||||
#[tauri::command]
|
||||
pub async fn relocate_project_path(
|
||||
state: State<'_, AppState>,
|
||||
id: String,
|
||||
new_path: String,
|
||||
) -> Result<ProjectRecord, String> {
|
||||
if !Path::new(&new_path).is_dir() {
|
||||
return Err(format!("目录不存在: {new_path}"));
|
||||
}
|
||||
if let Some(conflict) = find_binding_conflict(&state, &new_path, Some(&id)).await? {
|
||||
return Err(format!("目录已被项目「{}」绑定", conflict.name));
|
||||
}
|
||||
// 重探测技术栈
|
||||
let stack = detect_stack(Path::new(&new_path)).map_err(|e| e.to_string())?;
|
||||
let stack_json = serde_json::to_string(&stack).map_err(|e| e.to_string())?;
|
||||
state
|
||||
.projects
|
||||
.update_field(&id, "path", &new_path)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
state
|
||||
.projects
|
||||
.update_field(&id, "stack", &stack_json)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
state
|
||||
.projects
|
||||
.get_by_id(&id)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| "项目不存在".to_string())
|
||||
}
|
||||
|
||||
/// 检查目录是否存在(详情页「目录是否还在」用)
|
||||
#[tauri::command]
|
||||
pub async fn check_path_exists(path: String) -> Result<bool, String> {
|
||||
Ok(Path::new(&path).is_dir())
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// AI 扫描项目 — LLM 分析采样自动填基础信息
|
||||
// ============================================================
|
||||
|
||||
/// AI 扫描项目结果(预览用,用户确认后填入 ProjectRecord)
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct AiScanResult {
|
||||
/// LLM 产出的项目摘要(空=LLM 未得出,前端提示手填)
|
||||
pub description: String,
|
||||
/// 技术栈(规则探测 ∪ LLM 推断,去重小写)
|
||||
pub stack: Vec<String>,
|
||||
/// 项目类型(web/api/cli/library/desktop/mobile/monorepo/other)
|
||||
pub project_type: Option<String>,
|
||||
/// LLM 原始返回(降级时含错误信息,前端可展示)
|
||||
pub raw: Option<String>,
|
||||
}
|
||||
|
||||
/// AI 扫描项目目录,自动分析基础信息(description/stack/project_type)
|
||||
///
|
||||
/// 规则探测(detect_stack)兜底 + LLM 分析采样产出摘要。LLM 失败/解析失败降级纯规则。
|
||||
/// 需已配置默认 AI provider(无则报错提示去设置)。
|
||||
#[tauri::command]
|
||||
pub async fn scan_project_with_ai(
|
||||
state: State<'_, AppState>,
|
||||
path: String,
|
||||
) -> Result<AiScanResult, String> {
|
||||
let root = Path::new(&path);
|
||||
if !root.is_dir() {
|
||||
return Err(format!("目录不存在: {path}"));
|
||||
}
|
||||
|
||||
// 1. 规则探测(兜底)+ 采样(纯 IO 轻量,直接调)
|
||||
let rule_stack = detect_stack(root).map_err(|e| e.to_string())?;
|
||||
let sample = collect_sample(root).map_err(|e| e.to_string())?;
|
||||
|
||||
// 2. 取默认 provider(优先 is_default,否则首个)
|
||||
let providers = state.ai_providers.list_all().await.map_err(|e| e.to_string())?;
|
||||
let pc = providers
|
||||
.iter()
|
||||
.find(|p| p.is_default)
|
||||
.cloned()
|
||||
.or_else(|| providers.into_iter().next())
|
||||
.ok_or_else(|| "未配置 AI 提供商,请先在设置中添加".to_string())?;
|
||||
|
||||
// 3. 构造 provider + LLM 调用(非流式)
|
||||
let provider = build_provider(&pc.provider_type, &pc.base_url, &pc.api_key, &pc.default_model);
|
||||
let request = CompletionRequest {
|
||||
model: pc.default_model.clone(),
|
||||
messages: build_scan_prompt(&sample, &rule_stack),
|
||||
temperature: Some(0.2),
|
||||
max_tokens: Some(400),
|
||||
stream: false,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
};
|
||||
|
||||
// 4. 双层限流 + complete
|
||||
let _global_permit = state.llm_concurrency.acquire_global().await;
|
||||
let _per_conv_permit = state.llm_concurrency.acquire_per_conv().await;
|
||||
let llm_result = provider.complete(request).await;
|
||||
|
||||
// 5. 解析 + 合并 stack(LLM 失败降级纯规则)
|
||||
match llm_result {
|
||||
Ok(resp) => {
|
||||
let raw = resp.text.clone();
|
||||
match parse_scan_result(&resp.text) {
|
||||
Some(p) => {
|
||||
let mut stack = rule_stack;
|
||||
for s in p.stack {
|
||||
let s = s.trim().to_lowercase();
|
||||
if !s.is_empty() && !stack.contains(&s) {
|
||||
stack.push(s);
|
||||
}
|
||||
}
|
||||
Ok(AiScanResult { description: p.description, stack, project_type: p.project_type, raw: Some(raw) })
|
||||
}
|
||||
None => Ok(AiScanResult { description: String::new(), stack: rule_stack, project_type: None, raw: Some(raw) }),
|
||||
}
|
||||
}
|
||||
Err(e) => Ok(AiScanResult {
|
||||
description: String::new(),
|
||||
stack: rule_stack,
|
||||
project_type: None,
|
||||
raw: Some(format!("LLM 调用失败: {e}")),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
struct ParsedScan {
|
||||
description: String,
|
||||
stack: Vec<String>,
|
||||
project_type: Option<String>,
|
||||
}
|
||||
|
||||
/// 拼 LLM 扫描 prompt(system + 采样信息)
|
||||
fn build_scan_prompt(sample: &df_project::scan::ProjectSample, rule_stack: &[String]) -> Vec<ChatMessage> {
|
||||
let system = "你是项目分析助手。根据给定的项目采样信息,分析并输出项目基础信息。\n\
|
||||
严格只输出一个 JSON 对象,不要任何解释、markdown 代码块或额外文字。格式:\n\
|
||||
{\"description\":\"一句话中文项目摘要,描述项目做什么,30-60字\",\"stack\":[\"技术栈标签(小写英文,如 vue/rust/go)\"],\"project_type\":\"web|api|cli|library|desktop|mobile|monorepo|other\"}\n\
|
||||
规则:stack 用小写英文标签且去重;description 中文;project_type 从给定枚举选最接近的。若无足够信息,description 填空字符串。";
|
||||
let rule = if rule_stack.is_empty() { "(无)".to_string() } else { rule_stack.join(", ") };
|
||||
let tree = if sample.tree.is_empty() { "(无)".to_string() } else { sample.tree.join("\n") };
|
||||
let readme = sample.readme.clone().unwrap_or_else(|| "(无)".to_string());
|
||||
let manifests = if sample.manifests.is_empty() {
|
||||
"(无)".to_string()
|
||||
} else {
|
||||
sample.manifests.iter().map(|(n, c)| format!("### {n}\n{c}")).collect::<Vec<_>>().join("\n\n")
|
||||
};
|
||||
let user = format!(
|
||||
"## 已探测技术栈(规则)\n{rule}\n\n## 目录结构(2层)\n{tree}\n\n## README\n{readme}\n\n## 清单文件\n{manifests}"
|
||||
);
|
||||
vec![ChatMessage::system(system), ChatMessage::user(user)]
|
||||
}
|
||||
|
||||
/// 解析 LLM 返回的 JSON(容错:直接解析失败则提取首个 {...} 再解析)
|
||||
fn parse_scan_result(text: &str) -> Option<ParsedScan> {
|
||||
let extract = |v: &serde_json::Value| -> Option<ParsedScan> {
|
||||
let description = v.get("description").and_then(|x| x.as_str()).unwrap_or("").to_string();
|
||||
let stack = v
|
||||
.get("stack")
|
||||
.and_then(|x| x.as_array())
|
||||
.map(|arr| arr.iter().filter_map(|x| x.as_str().map(String::from)).collect())
|
||||
.unwrap_or_default();
|
||||
let project_type = v.get("project_type").and_then(|x| x.as_str()).map(String::from);
|
||||
Some(ParsedScan { description, stack, project_type })
|
||||
};
|
||||
if let Ok(v) = serde_json::from_str::<serde_json::Value>(text) {
|
||||
return extract(&v);
|
||||
}
|
||||
// 提取首个 {...}(LLM 可能裹 markdown 代码块或前后文字)
|
||||
let start = text.find('{')?;
|
||||
let end = text.rfind('}')?;
|
||||
if end <= start {
|
||||
return None;
|
||||
}
|
||||
let v = serde_json::from_str::<serde_json::Value>(&text[start..=end]).ok()?;
|
||||
extract(&v)
|
||||
}
|
||||
|
||||
47
src-tauri/src/commands/settings.rs
Normal file
47
src-tauri/src/commands/settings.rs
Normal file
@@ -0,0 +1,47 @@
|
||||
//! 通用应用设置 KV IPC — 前端 localStorage 迁移目标
|
||||
//!
|
||||
//! 4 个 command 读写 `app_settings` 表(key/value JSON 字符串),承载主题、面板折叠态、
|
||||
//! 最近使用项等前端持久化偏好。返回统一 `Result<T, String>`。
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use tauri::State;
|
||||
|
||||
use crate::state::AppState;
|
||||
|
||||
/// 取单个 key 的值(JSON 字符串),不存在返回 None
|
||||
#[tauri::command]
|
||||
pub async fn settings_get(
|
||||
state: State<'_, AppState>,
|
||||
key: String,
|
||||
) -> Result<Option<String>, String> {
|
||||
state.settings.get(&key).await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 写 key/value(`INSERT OR REPLACE`),刷新 updated_at
|
||||
#[tauri::command]
|
||||
pub async fn settings_set(
|
||||
state: State<'_, AppState>,
|
||||
key: String,
|
||||
value: String,
|
||||
) -> Result<bool, String> {
|
||||
state.settings.set(&key, &value).await.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 取全部 key/value(前端启动时一次性拉回恢复偏好)
|
||||
#[tauri::command]
|
||||
pub async fn settings_get_all(
|
||||
state: State<'_, AppState>,
|
||||
) -> Result<HashMap<String, String>, String> {
|
||||
let rows = state.settings.get_all().await.map_err(|e| e.to_string())?;
|
||||
Ok(rows.into_iter().collect())
|
||||
}
|
||||
|
||||
/// 删除单个 key
|
||||
#[tauri::command]
|
||||
pub async fn settings_delete(
|
||||
state: State<'_, AppState>,
|
||||
key: String,
|
||||
) -> Result<bool, String> {
|
||||
state.settings.delete(&key).await.map_err(|e| e.to_string())
|
||||
}
|
||||
@@ -24,7 +24,7 @@ pub struct CreateTaskInput {
|
||||
}
|
||||
|
||||
fn default_priority() -> i32 {
|
||||
1
|
||||
2 // medium — 新任务默认中优先级(非 high),符合常识
|
||||
}
|
||||
|
||||
/// 列出任务,可按 project_id 过滤
|
||||
|
||||
@@ -7,20 +7,19 @@ use tauri::Manager;
|
||||
|
||||
use state::AppState;
|
||||
|
||||
#[tauri::command]
|
||||
fn greet(name: &str) -> String {
|
||||
format!("Hello, {}! Welcome to DevFlow.", name)
|
||||
}
|
||||
|
||||
#[cfg_attr(mobile, tauri::mobile_entry_point)]
|
||||
pub fn run() {
|
||||
tauri::Builder::default()
|
||||
.plugin(tauri_plugin_opener::init())
|
||||
.plugin(tauri_plugin_dialog::init())
|
||||
.plugin(tauri_plugin_window_state::Builder::default().build())
|
||||
.setup(|app| {
|
||||
// 数据库放在系统应用数据目录下:<app_data_dir>/devflow.db
|
||||
// 数据库放在系统应用数据目录 <app_data_dir> 下(dev/release 拆分):
|
||||
// Dev 模式 devflow-dev.db(可随意改动/清空),Build 模式 devflow.db(长期保留真实运行数据)
|
||||
let db_name = if cfg!(debug_assertions) { "devflow-dev.db" } else { "devflow.db" };
|
||||
let data_dir = app.path().app_data_dir()?;
|
||||
std::fs::create_dir_all(&data_dir)?;
|
||||
let db_path = data_dir.join("devflow.db");
|
||||
let db_path = data_dir.join(db_name);
|
||||
|
||||
// 初始化全局状态(打开数据库 + 执行迁移)
|
||||
let app_state = tauri::async_runtime::block_on(AppState::init(&db_path))?;
|
||||
@@ -28,13 +27,20 @@ pub fn run() {
|
||||
Ok(())
|
||||
})
|
||||
.invoke_handler(tauri::generate_handler![
|
||||
greet,
|
||||
// 项目
|
||||
commands::project::list_projects,
|
||||
commands::project::create_project,
|
||||
commands::project::get_project,
|
||||
commands::project::update_project,
|
||||
commands::project::delete_project,
|
||||
commands::project::list_deleted_projects,
|
||||
commands::project::restore_project,
|
||||
commands::project::purge_project,
|
||||
commands::project::scan_project_stack,
|
||||
commands::project::check_path_binding,
|
||||
commands::project::relocate_project_path,
|
||||
commands::project::check_path_exists,
|
||||
commands::project::scan_project_with_ai,
|
||||
// 任务
|
||||
commands::task::list_tasks,
|
||||
commands::task::create_task,
|
||||
@@ -45,6 +51,8 @@ pub fn run() {
|
||||
commands::idea::create_idea,
|
||||
commands::idea::update_idea,
|
||||
commands::idea::delete_idea,
|
||||
commands::idea::evaluate_idea,
|
||||
commands::idea::promote_idea,
|
||||
// 工作流
|
||||
commands::workflow::run_workflow,
|
||||
commands::workflow::list_workflow_executions,
|
||||
@@ -54,16 +62,41 @@ pub fn run() {
|
||||
commands::ai::ai_chat_send,
|
||||
commands::ai::ai_chat_stop,
|
||||
commands::ai::ai_approve,
|
||||
commands::ai::ai_pending_tool_calls,
|
||||
commands::ai::ai_chat_clear,
|
||||
commands::ai::ai_list_providers,
|
||||
commands::ai::ai_save_provider,
|
||||
commands::ai::ai_set_provider,
|
||||
commands::ai::ai_delete_provider,
|
||||
// AI 对话管理
|
||||
commands::ai::ai_conversation_create,
|
||||
commands::ai::ai_conversation_list,
|
||||
commands::ai::ai_conversation_switch,
|
||||
commands::ai::ai_conversation_delete,
|
||||
commands::ai::ai_conversation_rename,
|
||||
commands::ai::ai_conversation_archive,
|
||||
commands::ai::ai_list_skills,
|
||||
commands::ai::ai_set_concurrency_config,
|
||||
// 知识库
|
||||
commands::knowledge::knowledge_list,
|
||||
commands::knowledge::knowledge_get,
|
||||
commands::knowledge::knowledge_search,
|
||||
commands::knowledge::knowledge_create,
|
||||
commands::knowledge::knowledge_update_status,
|
||||
commands::knowledge::knowledge_record_reuse,
|
||||
commands::knowledge::knowledge_list_candidates,
|
||||
commands::knowledge::knowledge_archive,
|
||||
commands::knowledge::knowledge_get_config,
|
||||
commands::knowledge::knowledge_save_config,
|
||||
commands::knowledge::knowledge_extract_now,
|
||||
commands::knowledge::knowledge_get_detail,
|
||||
commands::knowledge::knowledge_update,
|
||||
commands::knowledge::knowledge_events,
|
||||
// 通用应用设置 KV(前端 localStorage 迁移目标)
|
||||
commands::settings::settings_get,
|
||||
commands::settings::settings_set,
|
||||
commands::settings::settings_get_all,
|
||||
commands::settings::settings_delete,
|
||||
])
|
||||
.run(tauri::generate_context!())
|
||||
.expect("error while running tauri application");
|
||||
|
||||
@@ -4,12 +4,14 @@ use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
|
||||
use anyhow::Result;
|
||||
use tokio::sync::Mutex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::sync::{Mutex, Semaphore};
|
||||
|
||||
use df_ai::ai_tools::{AiToolRegistry, RiskLevel};
|
||||
use df_ai::ai_tools::AiToolRegistry;
|
||||
use df_storage::crud::{
|
||||
AiConversationRepo, AiProviderRepo, AiToolExecutionRepo, IdeaRepo, NodeExecutionRepo,
|
||||
ProjectRepo, ReleaseRepo, TaskRepo, WorkflowRepo,
|
||||
AiConversationRepo, AiProviderRepo, AiToolExecutionRepo, IdeaRepo, KnowledgeEventsRepo,
|
||||
KnowledgeRepo, NodeExecutionRepo, ProjectRepo, ReleaseRepo, SettingsRepo, TaskRepo,
|
||||
WorkflowRepo,
|
||||
};
|
||||
use df_storage::db::Database;
|
||||
use df_workflow::eventbus::EventBus;
|
||||
@@ -17,6 +19,116 @@ use df_workflow::registry::NodeRegistry;
|
||||
|
||||
use crate::commands::ai::AiSession;
|
||||
|
||||
// ============================================================
|
||||
// 知识库配置(提取 + 注入)
|
||||
// ============================================================
|
||||
|
||||
/// AI 提炼触发方式
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ExtractTrigger {
|
||||
/// 对话正常完成时(默认)
|
||||
OnComplete,
|
||||
/// 对话闲置 N 秒后
|
||||
OnIdle,
|
||||
/// 仅手动按钮触发
|
||||
ManualOnly,
|
||||
}
|
||||
|
||||
/// 知识库行为配置(存 AppState 内存,前后端通过 IPC 读写)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct KnowledgeConfig {
|
||||
/// 提炼总开关,默认 true
|
||||
pub auto_extract: bool,
|
||||
/// 提炼触发方式,默认 OnComplete
|
||||
pub trigger_mode: ExtractTrigger,
|
||||
/// 最少消息数守卫(防闲聊噪音),默认 4
|
||||
pub min_messages: u32,
|
||||
/// 闲置触发超时(ms),默认 30000
|
||||
pub idle_timeout_ms: u64,
|
||||
/// 聊天时自动注入相关知识开关,默认 true
|
||||
pub auto_inject: bool,
|
||||
/// 语义检索(向量)总开关,默认 false——关闭时纯 LIKE 零外部依赖
|
||||
#[serde(default)]
|
||||
pub vector_enabled: bool,
|
||||
/// embedding 用的 provider id(仅 openai_compat 类型,Anthropic 无 embed API)
|
||||
#[serde(default)]
|
||||
pub embedding_provider_id: Option<String>,
|
||||
/// embedding 模型名(如 embedding-3 / text-embedding-3-small)
|
||||
#[serde(default)]
|
||||
pub embedding_model: Option<String>,
|
||||
}
|
||||
|
||||
impl Default for KnowledgeConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
auto_extract: true,
|
||||
trigger_mode: ExtractTrigger::OnComplete,
|
||||
min_messages: 4,
|
||||
idle_timeout_ms: 30_000,
|
||||
auto_inject: true,
|
||||
vector_enabled: false,
|
||||
embedding_provider_id: None,
|
||||
embedding_model: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// LLM 调用并发控制(双层 Semaphore)
|
||||
// ============================================================
|
||||
|
||||
/// LLM 调用并发控制 — 全局 + 单对话双层 Semaphore
|
||||
///
|
||||
/// 限流对象:所有真实 LLM 调用(主循环 stream_llm / 标题生成 / 知识提炼)。
|
||||
/// 不限流本地工具执行(tools.execute)——本地操作无外部成本、不受 RPM 约束。
|
||||
///
|
||||
/// 运行时调整:tokio Semaphore 的 permits 数构造时固定、不可增减,
|
||||
/// 故用 `Arc<Mutex<Arc<Semaphore>>>` 双层包装——替换内层 Arc 即重建 Semaphore。
|
||||
/// 已持有旧 permit 的任务不受影响(permit 绑定旧 Semaphore,软收敛),
|
||||
/// 新请求 lock 后克隆到最新 Arc、自动走新限制。旧 Semaphore 随最后 permit 释放而 drop。
|
||||
///
|
||||
/// 注意:per_conv 当前是应用级单一信号量(非 per-conv map)。因 AiSession 为单例 +
|
||||
/// generating 互斥,同一时刻仅一个对话的 loop 在跑,per_conv 退化为"单对话内并发"
|
||||
/// (主循环 stream_llm + 标题生成 + 知识提炼)。未来若支持多对话并发,
|
||||
/// 需改为 HashMap<conv_id, Semaphore>。
|
||||
#[derive(Clone)]
|
||||
pub struct LlmConcurrency {
|
||||
global: Arc<Mutex<Arc<Semaphore>>>,
|
||||
per_conv: Arc<Mutex<Arc<Semaphore>>>,
|
||||
}
|
||||
|
||||
impl LlmConcurrency {
|
||||
pub fn new(global: usize, per_conv: usize) -> Self {
|
||||
Self {
|
||||
global: Arc::new(Mutex::new(Arc::new(Semaphore::new(global)))),
|
||||
per_conv: Arc::new(Mutex::new(Arc::new(Semaphore::new(per_conv)))),
|
||||
}
|
||||
}
|
||||
|
||||
/// 取全局并发 permit(重建后新请求自动走最新 Semaphore)
|
||||
pub async fn acquire_global(&self) -> tokio::sync::OwnedSemaphorePermit {
|
||||
let sema = self.global.lock().await.clone();
|
||||
sema.acquire_owned().await.expect("llm global semaphore closed")
|
||||
}
|
||||
|
||||
/// 取单对话并发 permit
|
||||
pub async fn acquire_per_conv(&self) -> tokio::sync::OwnedSemaphorePermit {
|
||||
let sema = self.per_conv.lock().await.clone();
|
||||
sema.acquire_owned().await.expect("llm per_conv semaphore closed")
|
||||
}
|
||||
|
||||
/// 重建全局 Semaphore(软收敛:旧 permit 不回收,待其释放后新限制完全生效)
|
||||
pub async fn set_global(&self, permits: usize) {
|
||||
*self.global.lock().await = Arc::new(Semaphore::new(permits));
|
||||
}
|
||||
|
||||
/// 重建单对话 Semaphore
|
||||
pub async fn set_per_conv(&self, permits: usize) {
|
||||
*self.per_conv.lock().await = Arc::new(Semaphore::new(permits));
|
||||
}
|
||||
}
|
||||
|
||||
/// 应用全局状态 — 通过 `app.manage()` 注入,command 中以 `State<'_, AppState>` 取用
|
||||
pub struct AppState {
|
||||
/// 数据库句柄(Arc 包装,便于在异步任务中重建 Repo)
|
||||
@@ -48,14 +160,26 @@ pub struct AppState {
|
||||
pub ai_tools: Arc<AiToolRegistry>,
|
||||
/// AI 会话状态
|
||||
pub ai_session: Arc<Mutex<AiSession>>,
|
||||
// ── 知识库 ──
|
||||
/// 知识库 Repo
|
||||
pub knowledge: KnowledgeRepo,
|
||||
/// 知识生命线事件 Repo(产生/审核/引用/归档审计)
|
||||
pub knowledge_events: KnowledgeEventsRepo,
|
||||
/// 知识库行为配置(提取 + 注入)
|
||||
pub knowledge_config: Arc<Mutex<KnowledgeConfig>>,
|
||||
/// 通用应用设置 KV Repo(前端 localStorage 迁移目标)
|
||||
pub settings: SettingsRepo,
|
||||
// ── LLM 并发控制 ──
|
||||
/// LLM 调用并发上限(全局 + 单对话双层 Semaphore,运行时可调)
|
||||
pub llm_concurrency: LlmConcurrency,
|
||||
}
|
||||
|
||||
impl AppState {
|
||||
/// 初始化应用状态:打开(或创建)数据库并执行迁移,构建各 Repo 与节点注册表
|
||||
pub async fn init(db_path: &Path) -> Result<Self> {
|
||||
let db = Arc::new(Database::open(db_path).await?);
|
||||
let ai_tools = Arc::new(build_ai_tool_registry());
|
||||
Ok(Self {
|
||||
let ai_tools = Arc::new(crate::commands::ai::build_ai_tool_registry(&db));
|
||||
let state = Self {
|
||||
ideas: IdeaRepo::new(&db),
|
||||
projects: ProjectRepo::new(&db),
|
||||
tasks: TaskRepo::new(&db),
|
||||
@@ -66,11 +190,20 @@ impl AppState {
|
||||
ai_conversations: AiConversationRepo::new(&db),
|
||||
ai_tool_executions: AiToolExecutionRepo::new(&db),
|
||||
ai_session: Arc::new(Mutex::new(AiSession::new())),
|
||||
knowledge: KnowledgeRepo::new(&db),
|
||||
knowledge_events: KnowledgeEventsRepo::new(&db),
|
||||
knowledge_config: Arc::new(Mutex::new(KnowledgeConfig::default())),
|
||||
settings: SettingsRepo::new(&db),
|
||||
llm_concurrency: LlmConcurrency::new(3, 2),
|
||||
db,
|
||||
event_bus: EventBus::new(),
|
||||
registry: Arc::new(build_registry()),
|
||||
ai_tools,
|
||||
})
|
||||
};
|
||||
// 启动恢复:重启前卡 pending 的工具审批(内存 pending_approvals 已丢)从审计表重建,
|
||||
// 使重启后待审批不丢。前端经 ai_pending_tool_calls + switchConversation 恢复 toolCard 态。
|
||||
crate::commands::ai::restore_pending_approvals(&state).await;
|
||||
Ok(state)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -92,191 +225,3 @@ fn build_registry() -> NodeRegistry {
|
||||
registry
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// AI 工具注册
|
||||
// ============================================================
|
||||
|
||||
use serde_json::json;
|
||||
|
||||
fn object_schema(fields: Vec<(&str, &str, bool)>) -> serde_json::Value {
|
||||
let mut props = serde_json::Map::new();
|
||||
let mut required = Vec::new();
|
||||
|
||||
for (name, typ, req) in &fields {
|
||||
props.insert(name.to_string(), json!({ "type": typ }));
|
||||
if *req {
|
||||
required.push(name.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
json!({
|
||||
"type": "object",
|
||||
"properties": props,
|
||||
"required": required,
|
||||
})
|
||||
}
|
||||
|
||||
/// 构建 AI 工具注册表 — 注册所有 CRUD 操作为可调用工具
|
||||
fn build_ai_tool_registry() -> AiToolRegistry {
|
||||
let mut registry = AiToolRegistry::new();
|
||||
|
||||
// ── 只读工具 (Low risk) ──
|
||||
|
||||
registry.register(
|
||||
"list_projects",
|
||||
"列出所有项目,返回项目列表(ID、名称、状态、描述)",
|
||||
object_schema(vec![]),
|
||||
RiskLevel::Low,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry.register(
|
||||
"list_tasks",
|
||||
"列出任务,可按 project_id 筛选",
|
||||
object_schema(vec![("project_id", "string", false)]),
|
||||
RiskLevel::Low,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry.register(
|
||||
"list_ideas",
|
||||
"列出所有想法",
|
||||
object_schema(vec![]),
|
||||
RiskLevel::Low,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
// ── 创建工具 (Medium risk) ──
|
||||
|
||||
registry.register(
|
||||
"update_project",
|
||||
"更新项目的指定字段(name/status/description),需要提供项目 ID、字段名和新值",
|
||||
object_schema(vec![
|
||||
("id", "string", true),
|
||||
("field", "string", true),
|
||||
("value", "string", true),
|
||||
]),
|
||||
RiskLevel::Medium,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry.register(
|
||||
"create_project",
|
||||
"创建新项目",
|
||||
object_schema(vec![
|
||||
("name", "string", true),
|
||||
("description", "string", false),
|
||||
]),
|
||||
RiskLevel::Medium,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry.register(
|
||||
"create_task",
|
||||
"在指定项目下创建新任务",
|
||||
object_schema(vec![
|
||||
("project_id", "string", true),
|
||||
("title", "string", true),
|
||||
("description", "string", false),
|
||||
("priority", "integer", false),
|
||||
]),
|
||||
RiskLevel::Medium,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry.register(
|
||||
"create_idea",
|
||||
"捕获一个新想法",
|
||||
object_schema(vec![
|
||||
("title", "string", true),
|
||||
("description", "string", false),
|
||||
("tags", "string", false),
|
||||
("source", "string", false),
|
||||
]),
|
||||
RiskLevel::Medium,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
// ── 高风险工具 (High risk) ──
|
||||
|
||||
registry.register(
|
||||
"delete_project",
|
||||
"删除项目及其所有关联数据",
|
||||
object_schema(vec![("id", "string", true)]),
|
||||
RiskLevel::High,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry.register(
|
||||
"run_workflow",
|
||||
"运行指定的工作流 DAG",
|
||||
object_schema(vec![
|
||||
("name", "string", true),
|
||||
("dag", "object", true),
|
||||
]),
|
||||
RiskLevel::High,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
// ── 文件系统工具 ──
|
||||
|
||||
registry.register(
|
||||
"read_file",
|
||||
"读取文件内容,返回文本内容。支持 offset 和 limit 参数分页读取大文件",
|
||||
object_schema(vec![
|
||||
("path", "string", true),
|
||||
("offset", "integer", false),
|
||||
("limit", "integer", false),
|
||||
]),
|
||||
RiskLevel::Low,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry.register(
|
||||
"list_directory",
|
||||
"列出目录内容,返回文件和子目录列表(名称、类型、大小)",
|
||||
object_schema(vec![
|
||||
("path", "string", true),
|
||||
("recursive", "boolean", false),
|
||||
]),
|
||||
RiskLevel::Low,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry.register(
|
||||
"write_file",
|
||||
"写入或创建文件,自动创建不存在的父目录",
|
||||
object_schema(vec![
|
||||
("path", "string", true),
|
||||
("content", "string", true),
|
||||
]),
|
||||
RiskLevel::Medium,
|
||||
Box::new(|_args| {
|
||||
Box::pin(async { Ok(json!({"note": "工具执行在 IPC 层处理"})) })
|
||||
}),
|
||||
);
|
||||
|
||||
registry
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user