重构: 后端 df-ai/commands 拆分+df-nodes/workflow 改造+P0 bug 修复

- df-ai: context 历史中毒三档自愈 sanitize_messages(AC3)+anthropic_compat tool_use_id None 跳过(AC1/AC2)+删 router/stream 死码
- df-core: events 加 select_type+decisions 多选审批契约(F-260615-01)
- df-execute: shell run_command 工具复用(F-260615-05)
- df-nodes: human_node 多选校验+2 端到端测(F-01)+取消跳 set_failed(B-03b-R1/R2/R8)
- df-workflow: executor/dag/state cancel 闭环(B-06/07/03a/b)+provider approve options(R-PD-5)
- df-storage: find_path_conflict 抽公共(R-PD-11)+COLS 常量断言
- df-ideas: 删 IdeaPromoter/PromotionPolicy 死码(R-PD-14)
- src-tauri/commands/ai: secret keyring 迁移(FR-S1/R-PD-4)+GeneratingGuard RAII+disarm(B-09/26)+newConversation 软复位(B-10)+stream 心跳/stop select/空回复判错(B-02/04/05/15)+run_command(F-05)+mask audit(AR-3)
- src-tauri/commands/{project,task,workflow,mod,lib,state}: task detail IPC(F-02)+approve decisions+task list 联动(B-29)
- Cargo.lock+Cargo.toml 依赖同步
This commit is contained in:
2026-06-15 05:14:42 +08:00
parent 04032a2a8d
commit 2de0c6ecb7
37 changed files with 2457 additions and 484 deletions

View File

@@ -25,7 +25,62 @@ use super::audit::process_tool_calls;
use super::{AiChatEvent, AiSession};
/// Agentic 循环最大迭代次数
pub(crate) const MAX_AGENT_ITERATIONS: usize = 10;
///
/// 默认 10 轮。未来可配置接入点:接入 AppState(新增 `agent_config` 字段)/df-storage settings KV
/// 表后,改为从配置读(默认值仍为 10)。当前项目无 agent 配置位(AppState/df-storage config 表
/// 均无 agent 配置槽),故暂以常量承载,避免引入 AppState 新字段等大改(违反 P2 零行为变边界)。
/// 接入路径:run_agentic_loop 签名增 `max_iterations: usize` 参数,调用方从 AppState 读取透传。
pub const MAX_AGENT_ITERATIONS: usize = 10;
// ============================================================
// B-260615-09: generating 状态 RAII guard
// ============================================================
/// generating 复位 RAII guard,取代散布的手动 `session.generating = false`。
///
/// 两路复位:
/// - 正常路径:exit 点显式 `reset().await` 即时复位(emit 前调,保证"复位→emit"顺序,
/// 前端收事件时后端已可接下一条)。
/// - 异常路径(panic/未走正常 return):Drop 兜底 spawn 复位,防 generating 永真卡死前端。
///
/// 注:try_continue_agent_loop 不用 guard——其 should_continue=false 路径需保持
/// generating=true(审批等待态),全函数 guard 会误复位;该函数单点 provider-Err 复位保持手动。
struct GeneratingGuard {
session: Arc<Mutex<AiSession>>,
done: bool,
}
impl GeneratingGuard {
fn new(session: Arc<Mutex<AiSession>>) -> Self {
Self { session, done: false }
}
/// 显式复位 generating=false。emit 前调用保证顺序。幂等。
async fn reset(&mut self) {
if !self.done {
self.session.lock().await.generating = false;
self.done = true;
}
}
/// 解除 Drop 兜底复位但不复位 generating。审批等待 return 路径调用:
/// 保持 generating=true 留 try_continue 续生成,同时 Drop 因 done=true 跳过复位 spawn。
/// (B-260615-26: 修复审批执行后对话不续生成回归)
fn disarm(&mut self) {
self.done = true;
}
}
impl Drop for GeneratingGuard {
fn drop(&mut self) {
if !self.done {
let session = self.session.clone();
tauri::async_runtime::spawn(async move {
session.lock().await.generating = false;
});
}
}
}
/// Agentic 循环:流式接收 → 工具执行 → 结果回传 LLM → 循环
///
@@ -44,11 +99,44 @@ pub(crate) async fn run_agentic_loop(
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,
// B-260615-09: generating 状态由 RAII guard 收敛复位(正常 exit 显式 reset;panic/异常 Drop 兜底)
let mut guard = GeneratingGuard::new(session_arc.clone());
// FR-S1: resolve→ensure_resolved_key(空 key 早失败)→build_provider 三步统一走工厂
// 空 key 早失败(逻辑见 secret::ensure_resolved_key 单测):避免空 key 发请求吃 401,错误伪装成"API Key 无效"
//
// B-260615-17:resolve 一次复用——原实现 build_provider_for 成功后又独立调 resolve_provider_secret
// 取 key_len(重复 keyring resolve)。现 resolve 一次:既供 key_len 诊断日志,又供 build_provider,
// 去重复 keyring resolve 调用。逻辑等价于 secret::build_provider_for(resolve→ensure→build 三步),
// 仅因 build_provider_for 隐藏 resolved key 无法复用而在此内联(未改 secret.rs 锁边界)。
let resolved_key = super::secret::resolve_provider_secret(&provider_config);
let key_len = resolved_key.len();
let provider: Box<dyn LlmProvider> = match super::secret::ensure_resolved_key(
&provider_config.name, &resolved_key,
) {
Ok(()) => df_ai::build_provider(
&provider_config.provider_type,
&provider_config.base_url,
&resolved_key,
&provider_config.default_model,
),
Err(msg) => {
guard.reset().await;
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
error: msg,
conversation_id: Some(conv_id.clone()),
});
return;
}
};
// 诊断日志:401/错误时据此定位是 url/type/model/key 哪项问题(只记长度不记明文)
tracing::info!(
provider = %provider_config.name,
provider_type = %provider_config.provider_type,
base_url = %provider_config.base_url,
model = %provider_config.default_model,
key_len = key_len,
"[ai] 发起 LLM 请求"
);
let tool_defs = tools_arc.tool_definitions();
// 停止信号副本stream_llm 与每轮迭代共享读取,避免重复加锁
@@ -57,6 +145,10 @@ pub(crate) async fn run_agentic_loop(
// token 累加器:loop 生命周期内各轮叠加,退出时传 save_conversation(累加模式落库)
let mut tokens = TokenAccumulator::default();
// 收敛标志:仅当 LLM 末轮无 tool_calls 自行 break(正常收敛)时置 true;
// 区分"正常收敛退出"与"达 MAX 被截断退出"——后者末轮 tool_calls 仍非空(tool_result 不再回传 LLM),属异常
let mut converged = false;
for iteration in 0..MAX_AGENT_ITERATIONS {
// 用户请求停止 → 收尾退出(已生成文本已在上一轮入库)
if stop_flag.load(Ordering::SeqCst) {
@@ -69,13 +161,27 @@ pub(crate) async fn run_agentic_loop(
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;
guard.reset().await;
// 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;
}
// B-260615-11: 旧 loop 污染防护——每轮开始校验对话一致性。
// 用户新建/切换对话后 active_conversation_id 变更,本 loop(conv_id 快照)成陈旧,
// 继续跑会往新对话 push 消息/pending 造成污染。检测到即退出(guard Drop 复位 generating)。
{
let session = session_arc.lock().await;
if session.active_conversation_id.as_deref() != Some(conv_id.as_str()) {
tracing::warn!(
stale_conv = %conv_id,
active_conv = ?session.active_conversation_id,
"[ai] 对话已切换,旧 loop 退出(B-260615-11)避免污染新对话"
);
return;
}
}
// 新一轮通知前端(第二轮起),前端需新建 assistant 消息
if iteration > 0 {
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiAgentRound {
@@ -119,8 +225,7 @@ pub(crate) async fn run_agentic_loop(
Some(result) => result,
None => {
// 错误已在 stream_llm 中 emit直接结束
let mut session = session_arc.lock().await;
session.generating = false;
guard.reset().await;
return;
}
};
@@ -136,6 +241,16 @@ pub(crate) async fn run_agentic_loop(
let has_tool_calls = !tool_calls_acc.is_empty();
{
let mut session = session_arc.lock().await;
// B-260615-11: push 前再校验(stream_llm 期间用户可能新建对话)。
// 读端读到被 clear 的空历史不致命,但 push 写回新对话是污染,必须挡。
if session.active_conversation_id.as_deref() != Some(conv_id.as_str()) {
tracing::warn!(
stale_conv = %conv_id,
active_conv = ?session.active_conversation_id,
"[ai] stream 后对话已切换,丢弃本轮 push(B-260615-11)避免污染新对话"
);
return;
}
if has_tool_calls {
let mut order: Vec<u32> = tool_calls_acc.keys().copied().collect();
order.sort_unstable();
@@ -165,15 +280,14 @@ pub(crate) async fn run_agentic_loop(
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;
guard.reset().await;
// 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; }
// 无工具调用 → 最终文本响应,正常收敛退出
if !has_tool_calls { converged = true; break; }
// 处理工具调用Low 自动执行 / Medium+High 待审批)
let pending_count = {
@@ -189,12 +303,30 @@ pub(crate) async fn run_agentic_loop(
total_tokens: tokens.total(),
};
save_conversation(&session_arc, &db, &conv_id, Some(&usage), Some(&provider_config.default_model)).await;
// B-260615-26: 审批等待 return 前 disarm guard——保持 generating=true 留 try_continue 续生成,
// 同时 Drop 因 done=true 跳过复位 spawn(避免误复位审批态 generating 致 ai_approve→try_continue 不续)
guard.disarm();
return; // generating 保持 true
}
// 全部自动执行完成 → 继续下一轮
}
// 达 MAX 未收敛(LLM 末轮仍想调工具被截断,末轮 tool_result 不再回传 LLM):异常中断,提示用户
// 与 break 正常收敛(break→converged=true)区分:这里仍走入库+Completed,但前置发 AiError 警示
if !converged {
tracing::warn!(
conv_id = %conv_id,
max_iter = MAX_AGENT_ITERATIONS,
"[ai] agentic 循环达最大轮次(MAX_AGENT_ITERATIONS={})仍未收敛,可能未完成",
MAX_AGENT_ITERATIONS,
);
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
error: format!("达到最大轮次({} 轮),Agent 可能未完成(末轮工具结果未回传模型)", MAX_AGENT_ITERATIONS),
conversation_id: Some(conv_id.clone()),
});
}
// 正常完成
let usage = df_ai::provider::TokenUsage {
prompt_tokens: tokens.prompt(),
@@ -224,24 +356,72 @@ pub(crate) async fn run_agentic_loop(
});
}
let mut session = session_arc.lock().await;
session.generating = false;
guard.reset().await;
// 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 循环
///
/// B-260615-08:所有静默 return 点显式 emit 收尾事件,避免前端 streaming=true 永久卡。
/// 各 return 点的语义判断:
/// 1) should_continue=false(generating 已复位 / pending_approvals 非空):
/// - generating=false → 用户点了停止(ai_chat_stop 复位)或会话已结束,emit AiCompleted 标当前轮收敛
/// (streaming=true 由 AiCompleted 清理)
/// - pending_approvals 非空 → 转入审批等待态(其他审批未决),emit AiCompleted 标当前轮结束
/// (前端审批态 watchdog 已 clear,不卡)
/// 2) get_active_provider Err → 无可用 provider(配置丢失/全删),无法续生成,emit AiError
/// (语义:配置错误,用户需设 provider;非 generating 复位可恢复)
pub(crate) async fn try_continue_agent_loop(app: &AppHandle, state: &AppState) {
let should_continue = {
let (is_generating, has_pending) = {
let session = state.ai_session.lock().await;
session.generating && session.pending_approvals.is_empty()
(session.generating, !session.pending_approvals.is_empty())
};
let should_continue = is_generating && !has_pending;
if !should_continue { return; }
if !should_continue {
// generating=false(被 stop)或仍有审批(pending_approvals 非空):
// 统一 emit AiCompleted 标当前轮收敛,清前端 streaming。
// 轮 token 已在前序 AiCompleted/AiApprovalResult 流程落库,此处零 token 上报仅作收敛信号。
if is_generating {
// pending_approvals 非空但 generating 仍 true:转审批态,前端审批态 watchdog 已 clear,不卡
tracing::info!("[ai] try_continue 跳过:仍有待审批,转审批等待态");
} else {
// generating 已复位(用户 stop 或前序循环已 emit Completed):补发 AiCompleted 防前端卡住
tracing::info!("[ai] try_continue 跳过:generating 已复位(被 stop/已结束),补发 AiCompleted 清前端 streaming");
let conv_id = {
let session = state.ai_session.lock().await;
session.active_conversation_id.clone().unwrap_or_default()
};
let _ = app.emit("ai-chat-event", AiChatEvent::AiCompleted {
total_tokens: 0,
prompt_tokens: 0,
completion_tokens: 0,
conversation_id: Some(conv_id),
});
}
return;
}
let provider_config = match get_active_provider(state).await {
Ok(p) => p,
Err(_) => return,
Err(e) => {
// 无可用 provider(配置丢失/全删):无法续生成,emit AiError。
// 语义:配置错误,用户需在 Settings 设 provider;generating 复位由 run_agentic_loop 内
// build_provider_for Err 分支处理(同样 emit AiError),此处与之一致。
// 不用 GeneratingGuard:try_continue 的 should_continue=false 路径需保 generating=true(审批等待态),
// 全函数 guard 会误复位。此点单点 provider-Err 复位,语义独立。
let mut session = state.ai_session.lock().await;
session.generating = false;
let conv_id = session.active_conversation_id.clone().unwrap_or_default();
drop(session);
tracing::warn!(error = %e, "[ai] try_continue 失败:无可用 provider");
let _ = app.emit("ai-chat-event", AiChatEvent::AiError {
error: e,
conversation_id: Some(conv_id),
});
return;
}
};
let (lang, conv_id) = {
let session = state.ai_session.lock().await;

View File

@@ -259,6 +259,33 @@ pub async fn ai_chat_stop(state: State<'_, AppState>, app: AppHandle) -> Result<
}
// 流式生成中:置位让 loop 自行收尾
session.stop_flag.store(true, Ordering::SeqCst);
let conv_id = session.active_conversation_id.clone();
drop(session);
// B-260615-13 兜底任务loop 若 panic/异常退出漏发收尾stop_flag 无人读,
// 用户点 stop 无反应、generating 卡 true。这里 sleep 短超时后重检 generating
// 仍 true(loop 没复位)则强制复位 + emit AiCompleted 通知前端收尾。
// 正常路径(loop 活自行复位)此时 generating 已 false无操作退出。
let session_arc = state.ai_session.clone();
let app_handle = app.clone();
tauri::async_runtime::spawn(async move {
tokio::time::sleep(std::time::Duration::from_secs(3)).await;
let mut session = session_arc.lock().await;
if session.generating {
session.generating = false;
drop(session); // 释放锁后再 emit避免持锁调 runtime emit
let _ = app_handle.emit(
"ai-chat-event",
AiChatEvent::AiCompleted {
total_tokens: 0,
prompt_tokens: 0,
completion_tokens: 0,
conversation_id: conv_id,
},
);
}
});
Ok(())
}
@@ -281,11 +308,15 @@ fn mask_api_key(key: &str) -> String {
#[tauri::command]
pub async fn ai_list_providers(state: State<'_, AppState>) -> Result<Vec<AiProviderRecord>, String> {
let mut providers = state.ai_providers.list_all().await.map_err(|e| e.to_string())?;
// IPC 不传明文 api_key(FR-S1):前端编辑用空 apiKey 表示不改,mask 后前端 realm 不持有明文
// IPC 不传明文 api_key(FR-S1):前端编辑用空 apiKey 表示不改,mask 后前端 realm 不持有明文
// 迁移后 DB api_key 空 → 从 keyring 取真实密钥再 mask(前端看到 mask 但不持有明文)
for p in &mut providers {
if !p.api_key.is_empty() {
p.api_key = mask_api_key(&p.api_key);
}
let real = if !p.api_key.is_empty() {
p.api_key.clone() // 未迁移(老明文)
} else {
super::secret::get_provider_secret(&p.id).unwrap_or_default() // 迁移后从 keyring
};
p.api_key = if real.is_empty() { String::new() } else { mask_api_key(&real) };
}
Ok(providers)
}
@@ -319,20 +350,45 @@ pub async fn ai_save_provider(
.map_err(|e| e.to_string())?
.iter().any(|p| p.is_default),
};
// api_key 空(编辑不改 key,前端 list 拿到 mask 故留空)→保留原 DB 值(FR-S1)
let api_key = if api_key.is_empty() {
match &id {
Some(pid) => state.ai_providers.get_by_id(pid).await
.map_err(|e| e.to_string())?
.map(|p| p.api_key)
.unwrap_or_default(),
None => String::new(),
// FR-S1:密钥存 OS keyring,DB api_key 列恒空(不入明文)。
// api_key 非空 = 新/改密钥 → 写 keyring;空 = 编辑不改 → 保留原 keyring 密钥不动。
let provider_id = id.clone().unwrap_or_else(new_id);
if !api_key.is_empty() {
// 显式改/填密钥 → 写 keyring(现状不变)
if let Err(e) = super::secret::set_provider_secret(&provider_id, &api_key) {
return Err(format!("密钥保存到系统钥匙串失败: {}", e));
}
} else {
api_key
};
} else if let Some(pid) = &id {
// 空 key 编辑:保住密钥,防未迁移态静默丢失(R-PD-1)。
// 未迁移态(DB 有明文 + keyring 空)下,下方 INSERT OR REPLACE 会无条件清 DB api_key,
// 唯一密钥副本被覆盖成空 → keyring 也空 → resolve 返空 → provider 报废密钥永久丢失。
// 兜底:发现未迁移态先即时迁移补密钥,迁移成功后再让下方清 DB 明文(收敛到迁移完成态);
// 迁移失败则 Err 阻断保存且 INSERT OR REPLACE 不执行 → DB 明文保留,绝不劣化现状。
let old = state.ai_providers.get_by_id(pid).await
.map_err(|e| e.to_string())?;
if let Some(old) = old {
if !old.api_key.is_empty()
&& super::secret::get_provider_secret(pid).is_none()
{
// DB 有明文 且 keyring 无 → 即时迁移补密钥
if let Err(e) = super::secret::set_provider_secret(pid, &old.api_key) {
return Err(format!(
"检测到该提供商密钥尚未迁移至系统钥匙串,本次保存尝试即时迁移失败({})。\
已保留原密钥未改动——请检查系统钥匙串权限后再次保存。",
e
));
}
tracing::info!(
"[FR-S1] 编辑路径即时迁移 provider {} 密钥至 keyring(R-PD-1 兜底)",
pid
);
}
// else: keyring 已有 / DB 已空 → 下方 INSERT OR REPLACE 清 DB 明文安全
}
}
let api_key = String::new(); // DB 恒空(真实密钥在 keyring)
let record = AiProviderRecord {
id: id.unwrap_or_else(new_id),
id: provider_id,
name,
provider_type: if provider_type.is_empty() { "openai_compat".to_string() } else { provider_type },
api_key,
@@ -391,6 +447,11 @@ pub async fn ai_delete_provider(
provider_id: String,
) -> Result<(), String> {
state.ai_providers.delete(&provider_id).await.map_err(|e| e.to_string())?;
// CR-260615-01:DB 已删则清 keyring 残留密钥(失败仅 warn 不阻断——无 DB 消费方,
// 残留 keyring 不可复活;同 id 复用也不会读到旧密钥,因 set 覆盖写)
if let Err(e) = super::secret::delete_provider_secret(&provider_id) {
tracing::warn!("[FR-S1] keyring 清理失败 {} (残留但无消费方,不阻断删除): {}", provider_id, e);
}
// 删除的若是当前默认,清空 active 指向,避免悬空
let mut session = state.ai_session.lock().await;
if session.active_provider_id.as_deref() == Some(&provider_id) {
@@ -405,17 +466,34 @@ pub async fn ai_delete_provider(
/// 创建新对话
#[tauri::command]
pub async fn ai_conversation_create(state: State<'_, AppState>) -> Result<serde_json::Value, String> {
pub async fn ai_conversation_create(
app: AppHandle,
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;
// 生成中(含审批等待态)禁止新建对话:否则 clear 会清空 pending_approvals,
// 用户审批时 ai_approve 找不到 pending → generating 永不复位 → 面板卡死需重启
// B-260615-10: 生成中软复位取代硬拦——强制结束当前生成,让用户能立即新建对话。
// 旧 loop 经 B-260615-11 一致性校验(active_conversation_id 变更)自动退出,不污染新对话;
// stop_flag 置位作双保险,让 streaming 中的旧 loop 也尽快收尾。
if session.generating {
return Err("生成中无法新建对话,请先停止或等待完成".to_string());
let old_conv = session.active_conversation_id.clone();
session.generating = false;
session.pending_approvals.clear();
session.stop_flag.store(true, Ordering::SeqCst);
drop(session);
if let Some(old_conv) = old_conv {
let _ = app.emit("ai-chat-event", AiChatEvent::AiCompleted {
total_tokens: 0,
prompt_tokens: 0,
completion_tokens: 0,
conversation_id: Some(old_conv),
});
}
session = state.ai_session.lock().await;
}
session.active_conversation_id = Some(id.clone());
session.active_conv_created_at = Some(now);

View File

@@ -23,9 +23,11 @@ pub(crate) struct TokenAccumulator {
impl TokenAccumulator {
/// 叠加一轮用量(round_usage 为本轮流式末 chunk 的累计用量)
///
/// saturating_add:恶意/异常 provider 返回巨大值时,避免 u32 += 溢出回绕打乱后续 budget 判定。
pub(crate) fn add(&mut self, prompt: u32, completion: u32) {
self.prompt += prompt;
self.completion += completion;
self.prompt = self.prompt.saturating_add(prompt);
self.completion = self.completion.saturating_add(completion);
}
pub(crate) fn prompt(&self) -> u32 {
@@ -37,7 +39,7 @@ impl TokenAccumulator {
}
pub(crate) fn total(&self) -> u32 {
self.prompt + self.completion
self.prompt.saturating_add(self.completion)
}
}

View File

@@ -29,10 +29,15 @@ async fn resolve_embed_provider(
return None;
}
};
Some((
df_ai::build_provider(&rec.provider_type, &rec.base_url, &rec.api_key, &rec.default_model),
model,
))
// build_provider_for 含空 key 早失败:Err → 返回 None 触发 LIKE 降级(与原 embed 失败降级行为一致)
let provider = match super::secret::build_provider_for(&rec) {
Ok(p) => p,
Err(e) => {
tracing::warn!("embedding provider 密钥不可用(降级 LIKE): {}", e);
return None;
}
};
Some((provider, model))
}
/// 生成文本嵌入(向量检索用)
@@ -317,12 +322,11 @@ async fn extract_knowledge_from_conversation(
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,
);
let provider: Box<dyn LlmProvider> = match super::secret::build_provider_for(provider_config) {
Ok(p) => p,
// 空 key 早失败:上层(maybe_spawn_extraction/trigger_extraction_now)已 warn log 降级,语义一致
Err(e) => return Err(anyhow::anyhow!("provider 密钥不可用: {}", e)),
};
// LLM 并发限流(知识提炼属独立调用,纳入双层 Semaphore
let _global_permit = llm_concurrency.acquire_global().await;
let _per_conv_permit = llm_concurrency.acquire_per_conv().await;

View File

@@ -26,6 +26,7 @@ pub mod commands;
pub mod conversation;
pub mod knowledge_inject;
pub mod prompt;
pub mod secret;
pub mod skills;
pub mod stream_recv;
pub mod title;
@@ -84,12 +85,32 @@ pub enum AiChatEvent {
AiError { error: String, conversation_id: Option<String> },
/// Agent 循环新一轮(前端需新建 assistant 消息)
AiAgentRound { round: u32, conversation_id: Option<String> },
/// 流式心跳(静默期报活,前端 watchdog reset区分「LLM 在跑」与「真断」)
AiHeartbeat { conversation_id: Option<String> },
}
// ============================================================
// 会话状态(放 mod.rs,各子文件经 use super::* 拿到)
// ============================================================
/// AI 会话读侧状态视图(只读枚举,由 [`AiSession::session_state`] 推导)
///
/// 这是 **只读视图**,不替代 `AiSession` 上 `generating` / `pending_approvals` /
/// `stop_flag` 等字段定义——字段仍是唯一真相源,调用方写状态时继续写字段,
/// 读状态时统一走 `session_state()` 收敛判断逻辑(避免散布的 `if generating` /
/// `if !pending_approvals.is_empty()` 三路组合在各调用点各自实现)。
///
/// 优先级:`pending_approvals` 非空(有挂起审批优先报 AwaitingApproval
/// 即使同时在流式)> `generating`Streaming> 都不满足Idle
pub enum SessionState {
/// 空闲:既未生成、也无挂起审批
Idle,
/// 流式生成中generating 为 true 且无挂起审批
Streaming,
/// 等待审批pending_approvals 非空
AwaitingApproval,
}
/// AI 会话内状态Mutex 保护)
pub struct AiSession {
/// 对话历史ContextManager唯一消息真相源裁剪仅影响发送视图不影响持久化
@@ -100,7 +121,17 @@ pub struct AiSession {
pub active_conversation_id: Option<String>,
/// 活跃对话创建时间(懒创建:首条消息落库前仅存内存,upsert 时用作 created_at)
pub active_conv_created_at: Option<String>,
/// 挂起的审批tool_call_id → 审批信息)
/// 挂起的审批
///
/// **结构**:单层 `HashMap<tool_call_id, PendingApproval>`,按 `tool_call_id`
/// 路由审批结果(`tool_call_id` 是路由键)。`PendingApproval.conversation_id`
/// 是业务语义哪个对话产生的审批而非路由键——IPC 端按 `tool_call_id` 精确命中。
///
/// **现状**`ai_pending_tool_calls` 等读取路径按 `conversation_id` 做 O(n) 线性
/// 过滤。在典型场景(单对话、少量并发审批)下 n 极小O(n) 可接受,无需二级索引。
///
/// **未来**:若多会话并发、审批量显著增大,可考虑加二级索引
/// `conversation_id → Vec<tool_call_id>`)把按会话过滤降到 O(1)。
pub pending_approvals: HashMap<String, PendingApproval>,
/// 是否正在生成
pub generating: bool,
@@ -123,6 +154,31 @@ impl AiSession {
stop_flag: Arc::new(AtomicBool::new(false)),
}
}
/// 推导会话读侧状态视图(只读,不写任何字段)。
///
/// 优先级:`pending_approvals` 非空 → [`SessionState::AwaitingApproval`]
/// 否则 `generating` → [`SessionState::Streaming`];否则 → [`SessionState::Idle`]。
///
/// 收敛点:调用方读「会话处于什么阶段」时统一走本方法,替代散布的
/// `if !pending_approvals.is_empty() / if generating` 三路组合。
///
/// # 待替换调用点(本次仅新增视图,不替换调用点)
///
/// - `agentic.rs:283` — agent loop 入口/恢复处对 generating + 审批的组合判断
/// - `commands.rs:250` — `ai_pending_tool_calls` 等读取前的状态门控
/// - `commands.rs:451` — 审批提交/会话复位路径的状态判断
///
/// 上述三处替换由 P0 任务承接,本方法先就位供其切换。
pub fn session_state(&self) -> SessionState {
if !self.pending_approvals.is_empty() {
SessionState::AwaitingApproval
} else if self.generating {
SessionState::Streaming
} else {
SessionState::Idle
}
}
}
/// 待审批的工具调用

View File

@@ -0,0 +1,228 @@
//! FR-S1 api_key 密钥管理 — 真实密钥存 OS keyring,DB `api_key` 列迁移后存空串
//!
//! 设计:
//! - keyring entry: service=`devflow-ai-provider`, username=provider_id
//! - DB `api_key` 列恒空(迁移后/新建均空),真实密钥唯一源 = OS keyring
//! - 启动一次性迁移:`migrate_secrets_to_keyring` 读老明文 → keyring → DB 置空(失败保留明文下次重试)
//! - 消费点(build_provider)经 `resolve_provider_secret` 取:keyring 优先,fallback DB(兼容未迁移)
//! - 跨平台:Windows Credential Manager / macOS Keychain / Linux Secret Service
use std::collections::HashMap;
use std::fs;
use std::path::PathBuf;
use df_storage::crud::AiProviderRepo;
use df_storage::models::AiProviderRecord;
use keyring::Entry;
const KEYRING_SERVICE: &str = "devflow-ai-provider";
/// 迁移失败计数器阈值:同一 provider 累计失败到此次数 → 升级为 warn 提示明文密钥长期滞留风险。
/// 跨启动持久化(sidecar 文件),计数仅用于告警,不影响兼容时序(不强制迁移、不删明文)。
const MIGRATION_FAIL_THRESHOLD: u32 = 3;
/// 迁移失败计数 sidecar 文件(<cwd>/.devflow-keyring-failcount):逐行 `provider_id=count`。
/// cwd 未必是稳定路径,但 R-PD-4 目标仅是「检测到反复失败/滞留时告警」,误读为 0 即按未达阈值处理,无副作用。
fn failcount_path() -> PathBuf {
std::env::current_dir()
.unwrap_or_else(|_| PathBuf::from("."))
.join(".devflow-keyring-failcount")
}
/// 读取全部失败计数(id → count)。文件缺失/损坏 → 空 map(按未达阈值处理)。
fn read_failcounts() -> HashMap<String, u32> {
let mut map = HashMap::new();
if let Ok(text) = fs::read_to_string(failcount_path()) {
for line in text.lines() {
let mut parts = line.splitn(2, '=');
let id = parts.next().unwrap_or("").trim();
let cnt = parts.next().and_then(|s| s.trim().parse::<u32>().ok());
if !id.is_empty() {
if let Some(c) = cnt {
map.insert(id.to_string(), c);
}
}
}
}
map
}
/// 持久化全部失败计数。写入失败仅 log,不阻断迁移主流程。
fn write_failcounts(map: &HashMap<String, u32>) {
let mut text = String::new();
let mut entries: Vec<_> = map.iter().collect();
entries.sort_by(|a, b| a.0.cmp(b.0)); // 稳定顺序,减少无谓 diff
for (id, cnt) in entries {
text.push_str(id);
text.push('=');
text.push_str(&cnt.to_string());
text.push('\n');
}
if let Err(e) = fs::write(failcount_path(), text) {
tracing::debug!("[FR-S1] 迁移失败计数文件写入失败(忽略): {}", e);
}
}
/// 记录一次迁移失败并返回累计失败次数。持久化失败也不影响返回值(仍递增内存计数用于本次告警)。
fn record_migration_fail(id: &str) -> u32 {
let mut map = read_failcounts();
let next = map.get(id).copied().unwrap_or(0).saturating_add(1);
map.insert(id.to_string(), next);
write_failcounts(&map);
next
}
/// 清零某 provider 的失败计数(迁移成功后调用,避免历史失败在后续再触发误告警)。
fn clear_migration_failcount(id: &str) {
let mut map = read_failcounts();
if map.remove(id).is_some() {
write_failcounts(&map);
}
}
fn entry_for(id: &str) -> anyhow::Result<Entry> {
Entry::new(KEYRING_SERVICE, id).map_err(|e| anyhow::anyhow!("keyring entry 创建失败: {}", e))
}
/// 读取 provider 密钥(优先 keyring;无则 None)
pub fn get_provider_secret(id: &str) -> Option<String> {
let entry = entry_for(id).ok()?;
match entry.get_password() {
Ok(s) if !s.is_empty() => Some(s),
_ => None,
}
}
/// 消费点用:解析 provider 真实密钥 — keyring 优先,fallback DB.api_key(兼容未迁移老库)
pub fn resolve_provider_secret(record: &AiProviderRecord) -> String {
if !record.api_key.is_empty() {
return record.api_key.clone();
}
get_provider_secret(&record.id).unwrap_or_default()
}
/// 写入密钥到 keyring(覆盖)
pub fn set_provider_secret(id: &str, key: &str) -> anyhow::Result<()> {
let entry = entry_for(id)?;
entry.set_password(key).map_err(|e| anyhow::anyhow!("keyring 写入失败: {}", e))
}
/// 删除 keyring 密钥(provider 删除时清理)
pub fn delete_provider_secret(id: &str) -> anyhow::Result<()> {
let entry = entry_for(id)?;
entry.delete_credential().map_err(|e| anyhow::anyhow!("keyring 删除失败: {}", e))
}
/// 启动一次性迁移:DB 明文 → keyring → DB 置空(失败保留明文下次重试,非阻断)
pub async fn migrate_secrets_to_keyring(repo: &AiProviderRepo) -> anyhow::Result<usize> {
let providers = repo.list_all().await?;
let mut migrated = 0;
for mut p in providers {
if p.api_key.is_empty() {
continue; // 已迁移或无密钥
}
if let Err(e) = set_provider_secret(&p.id, &p.api_key) {
// 累计失败次数:达阈值(默认 3)升级告警,提示明文密钥长期滞留 SQLite(无加密)风险。
// 计数仅告警用,不改兼容时序——仍保留明文下次重试,不强制迁移、不删明文。
let n = record_migration_fail(&p.id);
if n >= MIGRATION_FAIL_THRESHOLD {
tracing::warn!(
"[FR-S1] provider {} keyring 迁移已连续失败 {} 次,明文 api_key 长期滞留 SQLite 文件(无加密)。\
建议:1) 确认 OS 钥匙串可用(Win Credential Manager / macOS Keychain);\
2) keyring 后端异常时排查对应平台后端;3) 必要时手动在设置中重新保存密钥触发写入",
p.id, n
);
} else {
tracing::warn!(
"[FR-S1] keyring 迁移失败 {} (累计 {}/{},保留明文下次重试): {}",
p.id, n, MIGRATION_FAIL_THRESHOLD, e
);
}
continue;
}
let pid = p.id.clone();
p.api_key.clear();
if let Err(e) = repo.insert(p).await {
tracing::warn!("[FR-S1] 迁移后清空 DB api_key 失败 {}: {}", pid, e);
}
// 迁移成功 → 清零该 provider 的失败计数(下次若再出现失败从 1 重新累计)
clear_migration_failcount(&pid);
migrated += 1;
}
if migrated > 0 {
tracing::info!("[FR-S1] {} 条 provider 密钥迁移至 OS keyring", migrated);
}
Ok(migrated)
}
/// 消费点统一入口:resolve_provider_secret → ensure_resolved_key → build_provider 三步打包。
///
/// 替代散落各处的「resolve + build_provider」二件套复制(漏 ensure_resolved_key → keyring 读不到时
/// 空 key 照发请求吃 401,用户看「key 已保存」反复重试无解)。返 Result 让调用方按场景处理:
/// - 主链(agentic loop):空 key 早失败 emit AiError
/// - 后台 task(title/knowledge 提炼/embedding):Err 上层降级/记 log
pub fn build_provider_for(
record: &AiProviderRecord,
) -> Result<Box<dyn df_ai::provider::LlmProvider>, String> {
let api_key = resolve_provider_secret(record);
ensure_resolved_key(&record.name, &api_key)?;
Ok(df_ai::build_provider(
&record.provider_type,
&record.base_url,
&api_key,
&record.default_model,
))
}
/// 校验已解析的密钥是否可用:空(含纯空白)→明确错误信息,非空→Ok。
/// 用于消费点(build_provider 前)早失败,避免空 key 发请求吃 401,错误伪装成"API Key 无效"。
pub fn ensure_resolved_key(provider_name: &str, resolved: &str) -> Result<(), String> {
if resolved.trim().is_empty() {
Err(format!(
"未读取到「{}」的 API 密钥(系统钥匙串无记录或已损坏),请在设置中重新填写并保存",
provider_name
))
} else {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ensure_resolved_key_rejects_empty() {
assert!(ensure_resolved_key("GLM", "").is_err());
}
#[test]
fn ensure_resolved_key_rejects_whitespace() {
// 纯空白也视为无密钥(防粘贴时只有空格)
assert!(ensure_resolved_key("GLM", " ").is_err());
}
#[test]
fn ensure_resolved_key_accepts_nonempty() {
assert!(ensure_resolved_key("GLM", "sk-abc").is_ok());
}
#[test]
fn ensure_resolved_key_error_mentions_provider_name() {
let err = ensure_resolved_key("我的提供商", "").unwrap_err();
assert!(err.contains("我的提供商"), "错误信息应含 provider 名便于定位");
}
#[test]
fn resolve_prefers_db_when_non_empty() {
// DB api_key 非空 → 直接返回 DB 值,不触发 keyring(FR-S1 兼容未迁移老库)
use df_storage::models::AiProviderRecord;
let rec = AiProviderRecord {
id: "t1".into(), name: "t".into(), provider_type: "openai_compat".into(),
api_key: "sk-db-fallback".into(), base_url: "https://x".into(),
default_model: "m".into(), models: None, is_default: false,
config: None, created_at: "0".into(), updated_at: "0".into(),
};
assert_eq!(resolve_provider_secret(&rec), "sk-db-fallback");
}
}

View File

@@ -6,17 +6,104 @@ use std::time::Duration;
use tauri::{AppHandle, Emitter};
use futures::StreamExt;
use tracing::warn;
use df_ai::provider::{CompletionRequest, LlmProvider};
use super::{AiChatEvent, ToolCallDraft};
/// 从 anyhow 错误中尽力提取 HTTP 状态码/错误分类,供诊断拼接。
///
/// `provider.stream()` 失败有两类来源:
/// - 业务层provider 在 non-2xx 时 `bail!("LLM 流式 API 错误 {status}: {body}")`,原始文本已含状态码,
/// 此处用正则从文本抠 `HTTP <code>` 或裸 `<code>`provider 串已含 status code
/// - 传输层reqwest不直接命名以避免给 src-tauri 加依赖):从 Display 文本识别
/// timeout/connect 关键词分类。
///
/// 返回 `(status_or_class, raw)`status 取文本中首个三位数;否则按关键词给 timeout/connect/unknown。
fn extract_error_diag(e: &anyhow::Error) -> (String, String) {
let raw = e.to_string();
// 1) 抠 HTTP 状态码:匹配 provider bail 串里的 "错误 4xx/5xx" 或 reqwest 的 "HTTP status"。
// R-P2-6:原实现按字节窗口 `&bytes[i..i+3]` 切片——若窗口恰好切在多字节 UTF-8 字符中间会 panic
// (依赖中文恰好 3 字节、状态码恰为 ascii 的巧合)。改为按 char 迭代 + 前后非数字边界判断,
// 既 UTF-8 安全又顺便修复"长数字串(如端口号 14012)内嵌 401 误命中"的潜在问题。
// 白名单码(首位 4/5,故只查 4xx/5xx 区段,减少无效匹配):
const HTTP_CODES: &[&str] = &[
"400", "401", "403", "404", "408", "409", "413",
"422", "429", "500", "502", "503", "504",
];
// 把 raw 按 char 收集,索引即 char 下标(非字节),边界判断用 char 安全
let chars: Vec<char> = raw.chars().collect();
let n = chars.len();
let mut i = 0;
while i + 3 <= n {
// 窗口必须是三个 ascii 数字
if chars[i].is_ascii_digit() && chars[i + 1].is_ascii_digit() && chars[i + 2].is_ascii_digit() {
let code: String = chars[i..i + 3].iter().collect();
// 前后边界必须非数字(否则会从端口号 14012 里抠出 401)
let prev_ok = i == 0 || !chars[i - 1].is_ascii_digit();
let next_ok = i + 3 == n || !chars[i + 3].is_ascii_digit();
if prev_ok && next_ok && HTTP_CODES.contains(&code.as_str()) {
return (format!("HTTP {}", code), raw);
}
i += 3; // 已确认是三连数字,跳过避免窗口重叠重复扫
continue;
}
i += 1;
}
// 2) 传输层分类reqwest Display 文本特征),不命名 reqwest 类型
let lower = raw.to_lowercase();
let class = if lower.contains("timeout") || lower.contains("超时") {
"timeout"
} else if lower.contains("connect") || lower.contains("dns") || lower.contains("resolve") {
"connect"
} else if lower.contains("timed out") {
"timeout"
} else {
"unknown"
};
(class.to_string(), raw)
}
/// AiError 诊断消息的上下文,区分「建连/首字节失败」与「流中途断」两类。
#[derive(Copy, Clone, Eq, PartialEq, Debug)]
pub(crate) enum DiagKind {
/// `provider.stream()` 直接返回 Err连接/鉴权/HTTP non-2xx 等
Init,
/// 流已建立next() 返回 ErrSSE 传输断/解析错等
MidStream,
}
/// 拼接 AiError 的可读诊断文本(纯函数,便于单测)。
///
/// 格式:`[<provider_name>] <上下文>(<status_or_class>): <raw>`
/// - provider_name`provider.name()`anthropic 协议为 "anthropic-compat"
/// openai 兼容为模型名)。当前 `LlmProvider` trait 未暴露 base_url/provider_type
/// 不改签名的前提下这是唯一可得的 provider 标识。
/// - status_or_class`HTTP 4xx/5xx` 或 `timeout`/`connect`/`unknown` 分类。
/// - rawanyhow 原始错误文本。
pub(crate) fn fmt_diag(provider_name: &str, kind: DiagKind, status_or_class: &str, raw: &str) -> String {
let ctx = match kind {
DiagKind::Init => "AI 调用失败",
DiagKind::MidStream => "流式接收错误",
};
format!("[{}] {}({}): {}", provider_name, ctx, status_or_class, raw)
}
/// 流式接收 LLM 响应,返回 (完整文本, 工具调用草稿)
///
/// 三类异常处理:
/// - idle timeout120s 无 chunk判定连接静默断emit AiError 返回 None
/// - 流尽但从未收到 finished 信号判定异常中断emit AiError 返回 None丢弃残缺不当完整入库
/// - 用户停止stop_flagbreak 返回 Some(已收文本),由调用方入库展示后退出
///
/// 心跳与停止响应B-260615-02 / B-260615-04单 `tokio::select!` 三分支
/// - `stream.next()`:正常 chunk 处理
/// - `heartbeat.tick()`30s静默期如工具执行后等下一轮首 chunkemit `AiHeartbeat`
/// 前端 watchdog 据此 reset区分「LLM 在跑」与「真断」(避免空气泡误报中断)
/// - `stop_notify.notified()`:用户点停止即时打断,不再等 chunk 到或 120s idle timeout
pub(crate) async fn stream_llm(
provider: &dyn LlmProvider,
request: CompletionRequest,
@@ -26,6 +113,8 @@ pub(crate) async fn stream_llm(
) -> Option<(String, HashMap<u32, ToolCallDraft>, df_ai::provider::TokenUsage)> {
/// 流式读取空闲超时:超过此时长无任何 chunk 即判定连接已断
const STREAM_IDLE_TIMEOUT: Duration = Duration::from_secs(120);
/// 心跳间隔静默期向前端报「LLM 仍在跑」reset watchdog
const HEARTBEAT_INTERVAL: Duration = Duration::from_secs(30);
match provider.stream(request).await {
Ok(mut stream) => {
@@ -35,56 +124,121 @@ pub(crate) async fn stream_llm(
let mut stopped = false;
let mut final_usage: Option<df_ai::provider::TokenUsage> = None;
// B-260615-15:heartbeat interval 提至 loop 外复用,避免每轮重建计时器
// (每轮重建会丢已积累的节拍,且 interval 首次 tick 立即返回的特性会被误用)。
// tokio interval 首 tick 立即返回——此处先丢弃首 tick,让心跳等满首个 30s 静默期才发
// (心跳语义是"静默期仍在跑",循环入口立即报无意义且会与 stream.next() 抢分支错过首 chunk)。
let mut heartbeat = tokio::time::interval(HEARTBEAT_INTERVAL);
heartbeat.tick().await;
loop {
// 用户主动停止:保留已收文本退出
// B-260615-04:stop 即时打断。stream.next() 阻塞等 chunk 时,
// 用户点停止需等 chunk 到或 120s idle timeout 才轮到此处检查——
// 合并到下方 select! 的 stop_notify 分支后此处为快路径(非阻塞首检)。
// stop_flag 可能被 stopChat() 在 select! 阻塞期间置位,
// select! 的 stop_notify 分支会唤醒;此处保留作冗余快检(非阻塞)。
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,
// 三分支 select!(B-260615-02 + B-260615-04 合并):
// 1) stream.next():正常 chunk(idle timeout 120s 包裹,真断仍 emit AiError)
// 2) heartbeat.tick():静默期发 AiHeartbeat reset 前端 watchdog
// 3) stop_notify.notified():用户停止即时打断
//
// 注意:stop_flag 是 AtomicBool 无 async 通知能力——
// 此处用「timeout 包 stream.next() + 进入循环前/后查 stop_flag + 循环内 30s 心跳 tick」
// 间接实现「≤30s 感知 stop」(每轮 select! 至多 120s,但心跳 tick 30s 一次会
// 触发 select! 返回 → 循环回顶部 stop_flag 快检)。无 Notify 依赖,改动最小。
tokio::select! {
// 1) 正常 chunk(idle timeout 包裹)
chunk_result = tokio::time::timeout(STREAM_IDLE_TIMEOUT, stream.next()) => {
match chunk_result {
Err(_elapsed) => {
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
error: "流式响应超时120 秒无数据,连接可能已断开)".to_string(),
conversation_id: Some(conv_id.to_string()),
});
return None;
}
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); }
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());
}
// provider 流式错误事件Anthropic SSE `type=="error"` 等):
// 不走 finished 完成路径,发 AiError + 丢弃残缺,与 Err 分支对齐OpenAI 路径一致性)。
if let Some(err_msg) = &chunk.error {
warn!(
provider = %provider.name(),
conv_id = %conv_id,
error = %err_msg,
"[ai] provider 流式错误事件",
);
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
error: fmt_diag(
provider.name(),
DiagKind::MidStream,
"stream-error",
err_msg,
),
conversation_id: Some(conv_id.to_string()),
});
return None;
}
if chunk.finished {
finished_received = true;
break;
}
}
Err(e) => {
// 诊断:流中途错误(多为 SSE 传输断),补 provider 标识 + HTTP 状态/分类 + 原始文本
let (status_or_class, raw) = extract_error_diag(&e);
warn!(
provider = %provider.name(),
status = %status_or_class,
conv_id = %conv_id,
error = %raw,
"[ai] 流式接收中途错误",
);
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
error: fmt_diag(
provider.name(),
DiagKind::MidStream,
&status_or_class,
&raw,
),
conversation_id: Some(conv_id.to_string()),
});
return None;
}
}
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;
}
},
}
// 2) 心跳:静默期 30s 发 AiHeartbeat,前端 watchdog reset(B-260615-02)
_ = heartbeat.tick() => {
// 心跳只在「仍在等下一 chunk」时有意义——若 stop_flag 已置,顶部快检会 break,无需发心跳
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiHeartbeat {
conversation_id: Some(conv_id.to_string()),
});
}
}
}
@@ -93,8 +247,11 @@ pub(crate) async fn stream_llm(
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()) {
// 断连检测:流尽但从未收到 finished 信号 = 异常中断,丢弃残缺不当完整入库
// B-260615-05:一律判异常(不区分内容空否)——空内容无 finished 同属异常:
// 上游 agentic.rs !has_tool_calls 早 break 会按正常路径 emit AiCompleted,
// 用户看空气泡无错误提示,转 P0 卡死入口。此处统一 emit AiError 拦截。
if !finished_received {
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
error: "流式响应意外中断(未收到完成信号,已丢弃残缺响应)".to_string(),
conversation_id: Some(conv_id.to_string()),
@@ -105,11 +262,159 @@ pub(crate) async fn stream_llm(
Some((full_text, tool_calls_acc, final_usage.unwrap_or_default()))
}
Err(e) => {
// 诊断:连接/鉴权/HTTP 错误,补 provider 标识 + HTTP 状态码/分类 + 原始文本,
// 便于区分 401(key)/404(url)/429(限流)/timeout/连接失败(provider_type 或 base_url 不对)。
let (status_or_class, raw) = extract_error_diag(&e);
warn!(
provider = %provider.name(),
status = %status_or_class,
conv_id = %conv_id,
error = %raw,
"[ai] LLM 流式调用失败",
);
let _ = app_handle.emit("ai-chat-event", AiChatEvent::AiError {
error: format!("AI 调用失败: {}", e),
error: fmt_diag(
provider.name(),
DiagKind::Init,
&status_or_class,
&raw,
),
conversation_id: Some(conv_id.to_string()),
});
None
}
}
}
// ============================================================
// 单测:诊断提取/格式化(纯函数,不发 HTTP、不依赖 app_handle
// ============================================================
#[cfg(test)]
mod tests {
use super::*;
// ---- extract_error_diag业务层 bailprovider 串已含状态码)----
/// provider 在 non-2xx bail 的典型串:抠出 401鉴权失败/Key 错)
#[test]
fn diag_extracts_401_from_provider_bail() {
let e = anyhow::anyhow!("LLM 流式 API 错误 401: Unauthorized");
let (status, raw) = extract_error_diag(&e);
assert_eq!(status, "HTTP 401");
assert!(raw.contains("401"));
}
/// 404 = base_url/endpoint 不对
#[test]
fn diag_extracts_404_from_provider_bail() {
let e = anyhow::anyhow!("LLM 流式 API 错误 404: Not Found");
assert_eq!(extract_error_diag(&e).0, "HTTP 404");
}
/// 429 = 限流
#[test]
fn diag_extracts_429_from_provider_bail() {
let e = anyhow::anyhow!("LLM 流式 API 错误 429: rate limit");
assert_eq!(extract_error_diag(&e).0, "HTTP 429");
}
/// 5xx = 上游服务端错
#[test]
fn diag_extracts_500_from_provider_bail() {
let e = anyhow::anyhow!("LLM 流式 API 错误 500: Internal Server Error");
assert_eq!(extract_error_diag(&e).0, "HTTP 500");
}
// ---- extract_error_diag传输层reqwest Display 文本,不命名 reqwest----
/// 超时:连接成功但响应慢/静默断
#[test]
fn diag_classifies_timeout() {
let e = anyhow::anyhow!("error sending request for url (https://api.x.com/v1/chat/completions): operation timed out");
let (status, _) = extract_error_diag(&e);
assert_eq!(status, "timeout");
}
/// 中文「超时」关键词
#[test]
fn diag_classifies_timeout_cn() {
let e = anyhow::anyhow!("请求超时");
assert_eq!(extract_error_diag(&e).0, "timeout");
}
/// 连接失败DNS 解析失败 / 端点不通
#[test]
fn diag_classifies_connect() {
let e = anyhow::anyhow!("dns error: failed to lookup address information");
assert_eq!(extract_error_diag(&e).0, "connect");
}
/// resolve 关键词也归 connect
#[test]
fn diag_classifies_connect_resolve() {
let e = anyhow::anyhow!("error connecting: resolve failed");
assert_eq!(extract_error_diag(&e).0, "connect");
}
// ---- extract_error_diag边界 ----
/// 三位数但不在已知 HTTP 状态白名单(如 "200")→ 不当状态码,走分类
#[test]
fn diag_ignores_non_http_status_number() {
let e = anyhow::anyhow!("成功 200 条记录");
// 200 不在白名单,应落到 unknown
assert_eq!(extract_error_diag(&e).0, "unknown");
}
/// 完全无特征文本 → unknown且 raw 原样返回
#[test]
fn diag_unknown_preserves_raw() {
let e = anyhow::anyhow!("奇怪的错误 xyz");
let (status, raw) = extract_error_diag(&e);
assert_eq!(status, "unknown");
assert_eq!(raw, "奇怪的错误 xyz");
}
/// 状态码出现在更长数字串里也不误匹配(边界判断:前后必须非数字)
#[test]
fn diag_does_not_match_status_inside_longer_digits() {
// R-P2-6:边界判断后,白名单码嵌在长数字串内不再误命中(原窗口切片会从 14012 抠出 401)
let e = anyhow::anyhow!("port 14012 used");
assert_eq!(extract_error_diag(&e).0, "unknown");
// 纯长数字串(无白名单码)同样不命中
let e2 = anyhow::anyhow!("port 99999 used");
assert_eq!(extract_error_diag(&e2).0, "unknown");
}
// ---- fmt_diag两个上下文 ----
#[test]
fn fmt_diag_init_branch() {
let s = fmt_diag("anthropic-compat", DiagKind::Init, "HTTP 401", "Unauthorized");
assert_eq!(s, "[anthropic-compat] AI 调用失败(HTTP 401): Unauthorized");
}
/// R-P2-6:状态码紧贴中文字符(多字节 UTF-8)时仍能正确抠出
/// (原按字节窗口切片在中文边界脆弱,改为 char 迭代后安全)
#[test]
fn diag_extracts_code_adjacent_to_multibyte_chars() {
let e = anyhow::anyhow!("错误401Unauthorized");
assert_eq!(extract_error_diag(&e).0, "HTTP 401");
// 中文在前:状态码前字符是多字节中文,char 迭代边界判断正确
let e2 = anyhow::anyhow!("请求失败401");
assert_eq!(extract_error_diag(&e2).0, "HTTP 401");
}
#[test]
fn fmt_diag_midstream_branch() {
let s = fmt_diag("gpt-4o", DiagKind::MidStream, "timeout", "operation timed out");
assert_eq!(s, "[gpt-4o] 流式接收错误(timeout): operation timed out");
}
/// provider 名/错误文本含特殊字符仍按模板拼接(无格式注入风险)
#[test]
fn fmt_diag_handles_special_chars() {
let s = fmt_diag("model/x", DiagKind::Init, "unknown", "err: {json} \"quoted\"");
assert_eq!(s, "[model/x] AI 调用失败(unknown): err: {json} \"quoted\"");
}
}

View File

@@ -57,12 +57,17 @@ pub(crate) async fn ensure_conversation_title(
}
// 标题生成是独立一次 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,
);
// build_provider_for 含空 key 早失败:Err(如 keyring 读不到)→ 跳过 LLM,extract_title 兜底
let provider: Box<dyn LlmProvider> = match super::secret::build_provider_for(provider_config) {
Ok(p) => p,
Err(e) => {
tracing::warn!("标题生成跳过(provider 密钥不可用,走 extract_title 兜底): {}", e);
let title = 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", ());
return;
}
};
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()),

View File

@@ -1,9 +1,11 @@
//! AI 工具注册表构建 + 文件路径校验
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use df_ai::ai_tools::{AiToolRegistry, RiskLevel};
use df_execute::shell::{execute, ShellRequest};
use df_storage::db::Database;
use df_storage::models::{ProjectRecord, TaskRecord, IdeaRecord};
@@ -32,6 +34,26 @@ fn validate_path(path: &str) -> anyhow::Result<()> {
Ok(())
}
/// 截断命令输出:超过 max 字节则保留尾部(报错堆栈通常在末尾)+ 追加截断提示。
/// run_command 专用:防编译输出/find//cat 大文件撑爆 LLM context。
/// 按 char 边界截(防切多字节 UTF-8 中间 panic)。
fn truncate_output(s: &str, max: usize) -> (String, bool) {
if s.len() <= max {
return (s.to_string(), false);
}
let end = s.len();
let mut start = end.saturating_sub(max);
// 推进到 char 边界,避免从多字节 UTF-8 中间切开
while start < end && !s.is_char_boundary(start) {
start += 1;
}
let tail = &s[start..];
(
format!("[输出已截断,原始 {} 字节,仅保留末尾 {} 字节]\n{}", s.len(), tail.len(), tail),
true,
)
}
/// workspace 根目录(项目根 = src-tauri 上两级,编译期固定)
fn workspace_root() -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
@@ -88,17 +110,11 @@ async fn bind_dir_to_project(
if !dir.is_dir() {
anyhow::bail!("目录不存在: {path}");
}
// 防重复:normalize_path 规范化比较,防路径写法差异绕过(复用 df-project 公共 normalize_path)
// 防重复(DRY R-PD-11):委托 ProjectRepo::find_path_conflict,与 project.rs::find_binding_conflict
// 共用同一实现。normalize_path 规范化比较,防路径写法差异绕过(复用 df-project 公共 normalize_path)。
let target = df_project::scan::normalize_path(path);
let projects = repo.list_active().await?;
for proj in &projects {
if proj.id != id {
if let Some(pp) = &proj.path {
if df_project::scan::normalize_path(pp) == target {
anyhow::bail!("目录已被项目「{}」绑定", proj.name);
}
}
}
if let Some(conflict) = repo.find_path_conflict(&target, Some(id)).await? {
anyhow::bail!("目录已被项目「{}」绑定", conflict.name);
}
// stack:AI/调用方提供则用,否则探测(detect_stack 内含多次同步 fs IO,必须 spawn_blocking 防阻塞 tokio runtime)
let stack: Vec<String> = if let Some(s) = stack_opt {
@@ -388,6 +404,55 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
Ok(serde_json::json!({ "note": "请通过工作流页面运行工作流", "tool": "run_workflow" }))
})),
);
registry.register(
"run_command", "在指定工作目录执行 shell 命令(跑测试/构建/查看运行结果),返回 stdout/stderr/exit_code。高风险须人工批准。命令需自包含非交互式避免需用户输入的程序。默认超时 60 秒。用于验证刚写入的代码能否运行、跑测试、看报错迭代修改。",
df_ai::ai_tools::object_schema(vec![
("command", "string", true),
("working_dir", "string", false),
("timeout_secs", "integer", false),
]),
RiskLevel::High,
Box::new(|args: serde_json::Value| Box::pin(async move {
let command = args["command"].as_str()
.ok_or_else(|| anyhow::anyhow!("缺少 command 参数"))?;
// working_dir 默认 workspace_root()(与 write_file 锚定一致:AI 写代码→同目录跑命令,闭环)。
// 走 validate_path 黑名单(.. + 敏感系统目录)作基础防线;不走 resolve_workspace_path 越界校验——
// run_command 是 High risk 靠人工审批兜底(用户在审批卡看清 command+working_dir),
// 放开目录才能让 AI 在用户任意项目目录形成「写→跑→看→改」真闭环。
let working_dir = match args.get("working_dir").and_then(|v| v.as_str()) {
Some(d) => {
validate_path(d)?;
d.to_string()
}
None => workspace_root().to_string_lossy().to_string(),
};
// timeout 默认 60s:防 hang(交互式命令/死循环/大构建),LLM 可覆盖
let timeout_secs = args["timeout_secs"].as_u64().unwrap_or(60);
let request = ShellRequest {
command: command.to_string(),
working_dir: Some(working_dir.clone()),
env: HashMap::new(),
timeout_secs: Some(timeout_secs),
};
let result = execute(request).await?;
// 输出截断:防编译输出/find//cat 大文件撑爆 LLM context(各 10KB,尾部保留-报错堆栈在末尾)
const MAX_OUT: usize = 10_000;
let (stdout, stdout_trunc) = truncate_output(&result.stdout, MAX_OUT);
let (stderr, stderr_trunc) = truncate_output(&result.stderr, MAX_OUT);
Ok(serde_json::json!({
"command": command,
"working_dir": working_dir,
"exit_code": result.exit_code,
"duration_ms": result.duration_ms,
"stdout": stdout,
"stderr": stderr,
"truncated": stdout_trunc || stderr_trunc,
}))
})),
);
// ── 文件系统 ──
registry.register(
@@ -409,13 +474,23 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
if metadata.len() > 1_048_576 {
anyhow::bail!("文件超过 1MB 限制 ({} 字节)", metadata.len());
}
// 二进制/非 UTF-8 降级:read_to_string 对二进制硬失败,降级返 binary 标记而非错(防读二进制炸对话)
let mut content = String::new();
file.read_to_string(&mut content).await
.map_err(|e| anyhow::anyhow!("读取文件失败: {}", e))?;
if let Err(e) = file.read_to_string(&mut content).await {
if e.kind() == std::io::ErrorKind::InvalidData {
return Ok(serde_json::json!({
"path": path, "content": null, "binary": true,
"size": metadata.len(),
"error": "文件非 UTF-8 文本(疑似二进制),无法作为文本读取"
}));
}
anyhow::bail!("读取文件失败: {}", e);
}
// limit 硬上限 2000 行(防 LLM 传超大 limit 读全文件,1MB 限下仍可能数万行)
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;
let limit = args["limit"].as_u64().unwrap_or(200).min(2000) as usize;
lines.into_iter().skip(skip).take(limit).collect::<Vec<&str>>().join("\n")
} else {
content.clone()
@@ -426,7 +501,7 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
);
registry.register(
"list_directory", "列出目录内容,返回文件和子目录列表(名称、类型、大小)",
df_ai::ai_tools::object_schema(vec![("path", "string", true), ("recursive", "boolean", false), ("skip_noise_dirs", "boolean", false)]),
df_ai::ai_tools::object_schema(vec![("path", "string", true), ("recursive", "boolean", false), ("skip_noise_dirs", "boolean", false), ("max_depth", "integer", false)]),
RiskLevel::Low,
Box::new(|args: serde_json::Value| Box::pin(async move {
let resolved = resolve_workspace_path(
@@ -435,8 +510,9 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
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 max_depth = args["max_depth"].as_u64().unwrap_or(3) as usize;
let mut entries = Vec::new();
let truncated = list_dir_recursive(path, recursive, 0, 2, 1000, skip_noise, &mut entries).await?;
let truncated = list_dir_recursive(path, recursive, 0, max_depth, 1000, skip_noise, &mut entries).await?;
Ok(serde_json::json!({ "path": path, "entries": entries, "truncated": truncated }))
})),
);
@@ -454,13 +530,51 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
if content.len() > 1_048_576 {
anyhow::bail!("写入内容超过 1MB 限制 ({} 字节)", content.len());
}
if let Some(parent) = std::path::Path::new(path).parent() {
let target = std::path::Path::new(path);
// FR-S7 覆盖防护:覆盖非空文件前自动 .bak 备份(防 LLM 误用 write_file 当 edit 致数据彻底丢失)
// 起因:会话 3473fcb7 AI 误传头部 3 行把 PROGRESS.md 762 行/72KB 覆盖成 248 字节
let old_size: Option<u64> = match tokio::fs::metadata(target).await {
Ok(m) if m.len() > 0 => {
let bak = format!("{}.bak", path);
tokio::fs::copy(path, &bak).await
.map_err(|e| anyhow::anyhow!("备份 .bak 失败: {}", e))?;
Some(m.len())
}
Ok(_) => Some(0), // 空文件(无需备份)
Err(_) => None, // 不存在(新建)
};
if let Some(parent) = target.parent() {
// FR-S8 残余:parent 也必须在 workspace 内(防 path=workspace 根时 parent 越界 create_dir_all)
if !parent.starts_with(&workspace_root()) {
anyhow::bail!("禁止在项目目录之外创建目录");
}
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() }))
// FR-S7 原子写:tmp→rename,避免写到一半崩溃留半成品(.tmp-write 同目录保证 rename 不跨卷)
let tmp = format!("{}.tmp-write", path);
if let Err(e) = tokio::fs::write(&tmp, content).await {
let _ = tokio::fs::remove_file(&tmp).await;
return Err(anyhow::anyhow!("写入临时文件失败: {}", e));
}
if let Err(e) = tokio::fs::rename(&tmp, path).await {
let _ = tokio::fs::remove_file(&tmp).await;
return Err(anyhow::anyhow!("原子替换失败: {}", e));
}
// R-P2-2:rename 成功后清理 .bak(原子写已完成,.bak 不再需要);
// rename 失败分支不删 .bak——它是回退依据(失败分支已 return,不会走到这里)。
// 仅当 old_size>0(曾备份过)才清理;忽略清理失败(非阻断,最多留个孤儿 .bak 文件)
if old_size.map(|s| s > 0).unwrap_or(false) {
let bak = format!("{}.bak", path);
let _ = tokio::fs::remove_file(&bak).await;
}
// FR-S7 大小异动 warn:新内容远小于旧(疑似误覆盖整文件),提示用户查 .bak
if let Some(old) = old_size {
if old > 0 && (content.len() as f64 / old as f64) < 0.1 {
tracing::warn!("write_file 疑似误覆盖: {} {}→{} 字节(缩减>90%),.bak 已备份", path, old, content.len());
}
}
Ok(serde_json::json!({ "path": path, "bytes_written": content.len(), "old_size": old_size }))
})),
);
@@ -469,7 +583,7 @@ pub fn build_ai_tool_registry(db: &Arc<Database>) -> AiToolRegistry {
/// 递归列出目录内容(最多 max_depth 层,最多 max_entries 条)
///
/// - 噪音目录(`.git`/`node_modules`/`target` 等)会被列出(显示存在),但不深入其内部
/// - 噪音目录(`.git`/`node_modules`/`target` 等)在 `skip_noise=true` 时不作 entry 返回,也不深入其内部
/// - 达 `max_entries` 即停止,返回 `Ok(true)` 表示被截断
fn list_dir_recursive<'a>(
path: &'a str,
@@ -488,15 +602,27 @@ fn list_dir_recursive<'a>(
return Ok(true); // 达上限截断
}
let name = entry.file_name().to_string_lossy().to_string();
let metadata = entry.metadata().await?;
let is_dir = metadata.is_dir();
// 噪音目录(.git/node_modules/target 等)既不深入、也不作 entry 返回(避免污染 AI 工具上下文)
if skip_noise && is_noise_dir(&name) {
continue;
}
// 临时文件(.bak/.tmp-write 等)也不返回:write_file 原子写过程的副产物,
// 崩溃/中断会留孤儿文件,泄漏进 list_directory 噪化 AI 上下文(CR-260615-03)。
if skip_noise && is_noise_file(&name) {
continue;
}
// FR-S8:用 file_type() 不跟随 symlink——entry.metadata() 会解析符号链接目标,
// 经 workspace 内 symlink 即可泄露外部目标的 size/类型,且递归会进入 symlink 指向的 workspace 外目录
let file_type = entry.file_type().await?;
let is_symlink = file_type.is_symlink();
let is_dir = file_type.is_dir(); // symlink 算 symlink 非 directory,不会被递归
result.push(serde_json::json!({
"name": name,
"type": if is_dir { "directory" } else { "file" },
"size": metadata.len(),
"type": if is_symlink { "symlink" } else if is_dir { "directory" } else { "file" },
"size": if is_symlink { 0 } else { entry.metadata().await.map(|m| m.len()).unwrap_or(0) },
"depth": depth,
}));
// 仅递归非噪音目录(skip_noise=true 时跳过 .git/node_modules/target 等)
// 仅递归真实目录(symlink 不递归,防符号链接逃逸到 workspace 外)
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?;
@@ -517,3 +643,178 @@ fn is_noise_dir(name: &str) -> bool {
];
NOISE_DIRS.contains(&name)
}
/// 判断是否为不应返回给 AI 的噪音文件(临时/备份/编辑器产物)。
///
/// write_file 原子写过程产生 `.tmp-write`(rename 前)与 `.bak`(覆盖前备份,rename 成功后清理),
/// 崩溃/中断会留孤儿文件污染 list_directory 上下文。其它编辑器临时文件(`.swp`/`~` 等)同此处理。
/// 注:仅按后缀匹配,不依赖文件存在性——保持 list_directory 纯过滤语义,无额外 fs IO。
fn is_noise_file(name: &str) -> bool {
const NOISE_SUFFIXES: &[&str] = &[".bak", ".tmp-write", ".swp", "~"];
NOISE_SUFFIXES.iter().any(|sfx| name.ends_with(sfx))
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
/// is_noise_dir 纯函数:覆盖 .git/.gitignore 区分(目录是噪音,.gitignore 文件名不是)
#[test]
fn test_is_noise_dir_distinguishes_dir_and_gitignore_file() {
// 噪音目录命中
assert!(is_noise_dir(".git"));
assert!(is_noise_dir("node_modules"));
assert!(is_noise_dir("target"));
assert!(is_noise_dir("dist"));
assert!(is_noise_dir("build"));
// .gitignore / .gitattributes 等文件名不应命中(只匹配纯目录名 .git)
assert!(!is_noise_dir(".gitignore"));
assert!(!is_noise_dir(".gitattributes"));
assert!(!is_noise_dir(".gitkeep"));
}
/// is_noise_dir 纯函数:普通目录不命中
#[test]
fn test_is_noise_dir_normal_dirs_not_matched() {
assert!(!is_noise_dir("src"));
assert!(!is_noise_dir("tests"));
assert!(!is_noise_dir("docs"));
assert!(!is_noise_dir("main.rs"));
assert!(!is_noise_dir(""));
}
/// is_noise_dir 纯函数:大小写敏感(不靠 to_lowercase 误命中 Target)
#[test]
fn test_is_noise_dir_case_sensitive() {
// 文件系统在 Windows 上大小写不敏感,但函数本身用精确匹配;
// 锁定当前语义:大写变体不命中(避免后续误改 to_lowercase 引入行为变化)
assert!(!is_noise_dir("Target"));
assert!(!is_noise_dir("DIST"));
assert!(!is_noise_dir("NODE_MODULES"));
}
/// is_noise_file 纯函数:write_file 副产物 / 编辑器临时文件命中
#[test]
fn test_is_noise_file_matches_temp_artifacts() {
// write_file 原子写副产物
assert!(is_noise_file("PROGRESS.md.bak"));
assert!(is_noise_file("PROGRESS.md.tmp-write"));
// 编辑器临时文件
assert!(is_noise_file(".main.rs.swp"));
assert!(is_noise_file("main.rs~"));
// 路径含但非后缀的应不命中(避免误杀)
assert!(!is_noise_file("backup.bak.md"));
assert!(!is_noise_file("main.rs"));
assert!(!is_noise_file(".gitignore"));
assert!(!is_noise_file(""));
}
/// 造一个临时目录树用于 list_dir_recursive 测试
fn build_noise_tree(root: &Path) {
// .git/HEAD (噪音目录,内部文件不应出现)
fs::create_dir_all(root.join(".git")).unwrap();
fs::write(root.join(".git").join("HEAD"), "ref: refs/heads/main").unwrap();
// node_modules/x/index.js (噪音目录,内部文件不应出现)
fs::create_dir_all(root.join("node_modules").join("x")).unwrap();
fs::write(root.join("node_modules").join("x").join("index.js"), "module.exports=1;").unwrap();
// target/debug/bin (噪音目录,内部文件不应出现)
fs::create_dir_all(root.join("target").join("debug")).unwrap();
fs::write(root.join("target").join("debug").join("app"), "binary").unwrap();
// 普通 src/main.rs (应出现)
fs::create_dir_all(root.join("src")).unwrap();
fs::write(root.join("src").join("main.rs"), "fn main(){}").unwrap();
// .gitignore 文件 (文件名不是噪音目录,应出现)
fs::write(root.join(".gitignore"), "/target\n").unwrap();
}
/// list_dir_recursive:skip_noise=true 时噪音目录不出现在 entries(且不深入其内部)
#[tokio::test]
async fn test_list_dir_recursive_filters_noise_dirs() {
let tmp = std::env::temp_dir().join(format!("df_tool_noise_{}", std::process::id()));
let _ = fs::remove_dir_all(&tmp);
fs::create_dir_all(&tmp).unwrap();
build_noise_tree(&tmp);
let root = tmp.to_string_lossy().to_string();
let mut entries = Vec::new();
let truncated = list_dir_recursive(&root, true, 0, 3, 1000, true, &mut entries).await.unwrap();
assert!(!truncated);
let names: Vec<String> = entries.iter()
.filter_map(|e| e["name"].as_str().map(|s| s.to_string()))
.collect();
// 噪音目录本身及其内部文件均不应出现
assert!(!names.contains(&".git".to_string()), ".git 不应作为 entry 返回");
assert!(!names.contains(&"node_modules".to_string()), "node_modules 不应作为 entry 返回");
assert!(!names.contains(&"target".to_string()), "target 不应作为 entry 返回");
// 噪音目录内部文件也不应被递归带入
assert!(!names.iter().any(|n| n == "HEAD"), ".git/HEAD 不应出现");
assert!(!names.iter().any(|n| n == "index.js"), "node_modules/x/index.js 不应出现");
assert!(!names.iter().any(|n| n == "app"), "target/debug/app 不应出现");
// 普通目录与 .gitignore 文件应正常返回
assert!(names.contains(&"src".to_string()), "src 应作为 entry 返回");
assert!(names.contains(&"main.rs".to_string()), "src/main.rs 应被递归返回");
assert!(names.contains(&".gitignore".to_string()), ".gitignore 文件名不是噪音目录,应返回");
fs::remove_dir_all(&tmp).ok();
}
/// list_dir_recursive:skip_noise=false 时噪音目录正常返回
#[tokio::test]
async fn test_list_dir_recursive_keeps_noise_when_disabled() {
let tmp = std::env::temp_dir().join(format!("df_tool_noisefalse_{}", std::process::id()));
let _ = fs::remove_dir_all(&tmp);
fs::create_dir_all(&tmp).unwrap();
build_noise_tree(&tmp);
let root = tmp.to_string_lossy().to_string();
let mut entries = Vec::new();
list_dir_recursive(&root, false, 0, 1, 1000, false, &mut entries).await.unwrap();
let names: Vec<String> = entries.iter()
.filter_map(|e| e["name"].as_str().map(|s| s.to_string()))
.collect();
// 关闭过滤后噪音目录应作为 entry 出现
assert!(names.contains(&".git".to_string()));
assert!(names.contains(&"node_modules".to_string()));
assert!(names.contains(&"target".to_string()));
fs::remove_dir_all(&tmp).ok();
}
/// list_dir_recursive:max_depth 控制递归深度
#[tokio::test]
async fn test_list_dir_recursive_max_depth() {
// 造 a/b/c/main.rs (4 层: a 在 depth0, b depth1, c depth2, main.rs depth3)
let tmp = std::env::temp_dir().join(format!("df_tool_depth_{}", std::process::id()));
let _ = fs::remove_dir_all(&tmp);
fs::create_dir_all(tmp.join("a").join("b").join("c")).unwrap();
fs::write(tmp.join("a").join("b").join("c").join("main.rs"), "fn main(){}").unwrap();
let root = tmp.to_string_lossy().to_string();
// max_depth=1: 只到 depth1 (列出 a, 深入 a 列出 b, 但不再深入 b 因为 depth1 不 < 1)
let mut entries = Vec::new();
list_dir_recursive(&root, true, 0, 1, 1000, true, &mut entries).await.unwrap();
let names: Vec<String> = entries.iter()
.filter_map(|e| e["name"].as_str().map(|s| s.to_string()))
.collect();
assert!(names.contains(&"a".to_string()), "depth1 应含 a");
assert!(names.contains(&"b".to_string()), "max_depth=1 应深入 a 列出 b (depth1<1 为 false, 列 b 自身但不再递归)");
assert!(!names.contains(&"c".to_string()), "max_depth=1 不应到 c (depth1 已达上限)");
assert!(!names.contains(&"main.rs".to_string()), "max_depth=1 不应到 main.rs");
// max_depth=3: 应能到 depth3 的 main.rs
let mut entries = Vec::new();
list_dir_recursive(&root, true, 0, 3, 1000, true, &mut entries).await.unwrap();
let names: Vec<String> = entries.iter()
.filter_map(|e| e["name"].as_str().map(|s| s.to_string()))
.collect();
assert!(names.contains(&"main.rs".to_string()), "max_depth=3 应能递归到 main.rs");
fs::remove_dir_all(&tmp).ok();
}
}