Files
DevFlow/crates/df-nodes/src/ai_node.rs
绝尘 bd6a41fe6e 新增: 批次工作落地(推进链/评估闭环/事件总线/并发/加固) + 技术债清理 + 文档整理
后端:
- 工作流推进链(D-03):advance_task/状态机/闸门走 df-nodes Node trait,conditions 条件引擎扩展
- 想法评估闭环:启发式评分+对抗评估,df-ideas/scoring + df-storage/idea_eval_repo + idea 前端打通
- 全局事件数据总线:df-ai/context+context_helpers+augmentation 跨模块解耦
- AI planner/plan_hint/intent:aichat B 路线并行多轮基础
- patch_file 加固(TD-03/04):读改写整体锁防 lost update,expected_hash 合约闭环
- 压缩超时兜底(F-15 卡死根治)
- F-09 多会话并发:LlmConcurrency per-conv + streamingGuard 前端守护 + verify 脚本
- 知识注入 DRY/skills/audit 扩展

清理:
- aichat 技术债(误报 allow/死导入/过时注释 30 项)
- URGENT.md 删除(11 项加急全解决/迁 todo)
- 文档整理(todo/待决策/待审查/ARCHITECTURE/INDEX + 总线/技术债审查新文档)
2026-06-21 20:51:26 +08:00

424 lines
18 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! AI 节点 — 调用 LLM 完成文本生成/分析任务
//!
//! 工作流中无人值守的 AI 步骤:从节点 config 读取 OpenAI 兼容 provider 配置与
//! prompt自建 client 调用一次 LLM输出文本供下游节点消费。
//!
//! 与 AI Chat侧边栏交互对话的区别AI Node 由 DAG Executor 自动驱动,
//! 适合嵌入自动化链路(如 想法 → AI分析 → 脚本落地 → 人工审批)。
//!
//! 拆分边界(纯结构,逻辑等价):
//! - provider 解析/参数构造/自审 JSON 解析等纯逻辑 → `ai_node_helpers`
//! - AiSelfReviewNode(自审闭环独立节点) → `ai_self_review_node`
//! - 本文件:仅 AiNode 的 struct + impl + tests。
use std::sync::Arc;
use async_trait::async_trait;
use df_ai::provider::{ChatMessage, CompletionRequest, LlmProvider};
use df_storage::crud::TaskRepo;
use df_storage::db::Database;
use df_workflow::node::{Node, NodeContext, NodeOutput, NodeResult, NodeSchema};
// AiNode::execute 用的纯函数(resolve_and_parse / provider_from_params);
// 其余 helpers(parse_review_json/gate_should_block/REVIEW_SYSTEM_PROMPT/ResolvedProvider/
// truncate_for_summary)由 ai_self_review_node 与 tests 引入,不在此处 glob。
use crate::ai_node_helpers::{provider_from_params, resolve_and_parse};
/// AI 节点
///
/// 持有 Arc<Database> 在 NodeRegistry 注册时注入(state.rs build_registry 工厂闭包 move 捕获
/// db.clone(),对齐 TaskAdvanceNode 先例)。execute 从 NodeContext.config 读 provider 配置 +
/// 可选 task_id:有 task_id 时把产出落 task.output_json(AiNode 自审闭环,决策 a)。
pub struct AiNode {
db: Arc<Database>,
}
impl AiNode {
/// 构造节点(注册表工厂调用,注入数据库句柄)
pub fn new(db: Arc<Database>) -> Self {
Self { db }
}
}
#[async_trait]
impl Node for AiNode {
async fn execute(&self, ctx: NodeContext) -> NodeResult {
tracing::info!("AiNode 执行: node_id={}", ctx.node_id);
// FR-S1 注入链provider 经 df_storage::secret 在 AiNode 内存解析api_key 不进 config。
let p = resolve_and_parse(&self.db, &ctx.config, &ctx.inputs).await?;
let provider: Box<dyn LlmProvider> = provider_from_params(&p);
// ── 构建消息 ──
let mut messages = Vec::with_capacity(2);
if let Some(sys) = p.system_prompt {
messages.push(ChatMessage::system(sys));
}
messages.push(ChatMessage::user(p.prompt));
let request = CompletionRequest {
model: p.model.clone(),
messages,
temperature: p.temperature,
max_tokens: p.max_tokens,
stream: false,
tools: None,
tool_choice: None,
reasoning_content: None,
};
tracing::info!(
"AiNode 调用 LLM: model={}, base_url={}",
p.provider.default_model,
p.provider.base_url
);
let response = provider.complete(request).await?;
tracing::info!(
"AiNode 完成: model={}, prompt_tokens={}, completion_tokens={}",
response.model,
response.usage.prompt_tokens,
response.usage.completion_tokens
);
let output = serde_json::json!({
"text": response.text,
"model": response.model,
"usage": {
"prompt_tokens": response.usage.prompt_tokens,
"completion_tokens": response.usage.completion_tokens,
"total_tokens": response.usage.total_tokens,
},
});
// 决策 a(AiNode 自审闭环):有 task_id 时落产出到 task.output_json。
// 走通用 update_field(output_json 已在 tasks 白名单,crud/mod.rs update_field 宏(impl_repo!)),
// 非 status 状态机收口字段,合法可写。落库失败不阻断节点返回(产出已在内存,工作流可继续)。
if let Some(task_id) = ctx.config.get("task_id").and_then(|v| v.as_str()) {
if let Ok(json_str) = serde_json::to_string(&output) {
let repo = TaskRepo::new(&self.db);
if let Err(e) = repo.update_field(task_id, "output_json", &json_str).await {
tracing::warn!(
task_id = task_id,
error = %e,
"AiNode 落 output_json 失败(产出已在内存,节点继续)"
);
}
}
}
Ok(NodeOutput::from_value(output))
}
fn schema(&self) -> NodeSchema {
NodeSchema {
params: serde_json::json!({
"type": "object",
"properties": {
"prompt": { "type": "string", "description": "用户提示词(若无则取上游 prompt 输出)" },
"system_prompt": { "type": "string", "description": "系统提示词(可选)" },
"provider_id": { "type": "string", "description": "AI Provider IDFR-S1密钥经 df_storage::secret 解析,不进 config留空走默认 provider" },
"model": { "type": "string", "description": "模型名(可选,留空用 record.default_model" },
"temperature": { "type": "number", "description": "温度 0.0~2.0(可选)" },
"max_tokens": { "type": "integer", "description": "最大生成 token可选anthropic 协议无值时默认 4096" },
"base_url": { "type": "string", "description": "(已废弃过渡)明文 API 地址,改用 provider_id" },
"api_key": { "type": "string", "description": "(已废弃过渡)明文 API 密钥,改用 provider_idFR-S1 下经 secret 解析" }
},
// SW-260618-15: prompt/provider_id 均"留空走兜底"(prompt 取上游、provider_id 走默认 provider),与 required 矛盾。改 required=[] 对齐 execute 运行时,防前端按 schema 误拒合法配置。
"required": []
}),
output: serde_json::json!({
"type": "object",
"properties": {
"text": { "type": "string" },
"model": { "type": "string" },
"usage": { "type": "object" }
}
}),
}
}
fn node_type(&self) -> &str {
"ai"
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ai_node_helpers::{parse_params, resolve_provider, ResolvedProvider};
use df_storage::crud::AiProviderRepo;
use df_storage::db::Database;
use df_storage::models::AiProviderRecord;
use serde_json::json;
use std::collections::HashMap;
/// 测试用 ResolvedProvider 桩:固定 protocol/base_url/api_key/default_model。
fn provider_stub() -> ResolvedProvider {
ResolvedProvider {
protocol: "openai_compat".to_string(),
base_url: "https://api.example.com".to_string(),
api_key: "sk-test".to_string(),
default_model: "gpt-4o-mini".to_string(),
model_pool: Vec::new(),
}
}
/// 造带基础 prompt 的 configprovider 经 resolve_provider 解析parse_params 只看 prompt
/// overrides 覆盖或追加。
fn config_with(overrides: serde_json::Value) -> serde_json::Value {
let mut base = json!({
"prompt": "config-prompt"
});
if let (serde_json::Value::Object(b), serde_json::Value::Object(o)) = (&mut base, overrides) {
for (k, v) in o {
b.insert(k, v);
}
}
base
}
fn empty_inputs() -> HashMap<String, NodeOutput> {
HashMap::new()
}
#[test]
fn missing_prompt_errors() {
let config = json!({});
let err = parse_params(&config, &empty_inputs(), provider_stub())
.unwrap_err()
.to_string();
assert!(err.contains("prompt"), "缺 prompt 应报错, 实际: {}", err);
}
#[test]
fn prompt_from_config_when_no_upstream() {
let p =
parse_params(&config_with(json!({})), &empty_inputs(), provider_stub()).unwrap();
assert_eq!(p.prompt, "config-prompt");
}
#[test]
fn prompt_prefers_upstream_input_over_config() {
let mut inputs = empty_inputs();
inputs.insert(
"prompt".to_string(),
NodeOutput::from_value(json!("upstream-prompt")),
);
let p =
parse_params(&config_with(json!({})), &inputs, provider_stub()).unwrap();
assert_eq!(p.prompt, "upstream-prompt", "上游输入应优先于 config.prompt");
}
#[test]
fn optional_fields_parsed_when_present() {
let p = parse_params(
&config_with(json!({
"model": "glm-4",
"temperature": 0.3,
"max_tokens": 1024,
"system_prompt": "你是助手"
})),
&empty_inputs(),
provider_stub(),
)
.unwrap();
assert_eq!(p.model, "glm-4");
assert_eq!(p.provider.protocol, "openai_compat");
assert_eq!(p.temperature, Some(0.3));
assert_eq!(p.max_tokens, Some(1024));
assert_eq!(p.system_prompt.as_deref(), Some("你是助手"));
}
// ============================================================
// resolve_provider 双路径测试(FR-S1 注入链核心)
// ============================================================
/// 内存 DB 插 provider(可选 is_default),返回 (db, provider_id)。
async fn setup_provider_db(
provider_id: &str,
is_default: bool,
api_key: &str,
) -> (Database, String) {
let db = Database::open_in_memory().await.expect("open_in_memory");
let repo = AiProviderRepo::new(&db);
repo.insert(AiProviderRecord {
id: provider_id.to_string(),
name: "测试 Provider".to_string(),
provider_type: "openai_compat".to_string(),
api_key: api_key.to_string(),
base_url: "https://api.example.com".to_string(),
default_model: "glm-4-flash".to_string(),
models: None,
model_configs: Vec::new(),
is_default,
config: None,
created_at: "0".to_string(),
updated_at: "0".to_string(),
enabled: true,
weight: 50,
})
.await
.expect("insert provider");
(db, provider_id.to_string())
}
/// 路径 1provider_id 优先 → resolve_provider 走下沉解析(DB api_key 非空→直接返回)。
#[tokio::test]
async fn resolve_provider_by_id_prefers_db_key() {
let (db, pid) = setup_provider_db("prov-1", true, "sk-from-db").await;
let config = json!({ "provider_id": pid });
let resolved = resolve_provider(&Arc::new(db), &config).await.unwrap();
assert_eq!(resolved.api_key, "sk-from-db", "DB api_key 非空应直接返回");
assert_eq!(resolved.base_url, "https://api.example.com");
assert_eq!(resolved.protocol, "openai_compat");
assert_eq!(
resolved.default_model, "glm-4-flash",
"无 config.model 时用 record.default_model"
);
}
/// 路径 1config.model 覆盖 record.default_model。
#[tokio::test]
async fn resolve_provider_config_model_overrides_record() {
let (db, pid) = setup_provider_db("prov-1m", true, "sk-db").await;
let config = json!({ "provider_id": pid, "model": "glm-4" });
let resolved = resolve_provider(&Arc::new(db), &config).await.unwrap();
assert_eq!(resolved.default_model, "glm-4");
}
/// 路径 1provider_id 不存在 → 友好错误。
#[tokio::test]
async fn resolve_provider_unknown_id_errors() {
let (db, _pid) = setup_provider_db("prov-2", true, "sk-db").await;
let config = json!({ "provider_id": "non-existent" });
let err = resolve_provider(&Arc::new(db), &config)
.await
.unwrap_err()
.to_string();
assert!(err.contains("不存在"), "未知 provider_id 应报错, 实际: {}", err);
}
/// 路径 1provider_id 指向但密钥空(DB 空 + keyring 无)→ ensure_resolved_key 报错。
#[tokio::test]
async fn resolve_provider_empty_key_errors() {
// api_key 空DB 无 → keyring 读不到 → 空 → ensure 报错
let (db, pid) = setup_provider_db("prov-empty", true, "").await;
let config = json!({ "provider_id": pid });
let err = resolve_provider(&Arc::new(db), &config)
.await
.unwrap_err()
.to_string();
assert!(
err.contains("密钥") || err.contains("未读取到"),
"空密钥应报错, 实际: {}",
err
);
}
/// 路径 2老明文路径兼容(base_url+api_key 显式注入)→ 返回明文 + 不报错。
#[tokio::test]
async fn resolve_provider_legacy_plain_path_compat() {
// 无 provider 且无 provider_id → 老明文路径兜底
let db = Database::open_in_memory().await.expect("open_in_memory");
let config = json!({
"base_url": "https://legacy.example.com",
"api_key": "sk-legacy"
});
let resolved = resolve_provider(&Arc::new(db), &config).await.unwrap();
assert_eq!(resolved.api_key, "sk-legacy");
assert_eq!(resolved.base_url, "https://legacy.example.com");
assert_eq!(resolved.default_model, "gpt-4o-mini", "无 model 时占位兜底");
}
/// 路径 2老明文路径空 key → ensure_resolved_key 报错。
#[tokio::test]
async fn resolve_provider_legacy_empty_key_errors() {
let db = Database::open_in_memory().await.expect("open_in_memory");
let config = json!({ "base_url": "https://x", "api_key": "" });
let err = resolve_provider(&Arc::new(db), &config)
.await
.unwrap_err()
.to_string();
assert!(err.contains("密钥"), "老路径空 key 应报错, 实际: {}", err);
}
/// 路径 3空兜底 → 无 provider_id 无明文 → 取 is_default=true 首条。
#[tokio::test]
async fn resolve_provider_fallback_default_provider() {
let (db, _pid) = setup_provider_db("prov-default", true, "sk-default").await;
let config = json!({}); // 无 provider_id 无明文
let resolved = resolve_provider(&Arc::new(db), &config).await.unwrap();
assert_eq!(resolved.api_key, "sk-default", "兜底应取 is_default 首条");
}
/// 路径 3无任何 provider → 友好错误「未配置 AI Provider」。
#[tokio::test]
async fn resolve_provider_no_providers_errors() {
let db = Database::open_in_memory().await.expect("open_in_memory");
let config = json!({});
let err = resolve_provider(&Arc::new(db), &config)
.await
.unwrap_err()
.to_string();
assert!(
err.contains("未配置") || err.contains("Provider"),
"无 provider 应友好报错, 实际: {}",
err
);
}
/// 路径 3兜底无 is_default 时取首条(对齐 idea.rs build_default_provider)。
#[tokio::test]
async fn resolve_provider_fallback_first_when_no_default() {
let (db, _pid) = setup_provider_db("prov-first", false, "sk-first").await;
let config = json!({});
let resolved = resolve_provider(&Arc::new(db), &config).await.unwrap();
assert_eq!(resolved.api_key, "sk-first", "无 is_default 时取首条");
}
/// 真调 GLM 验证 provider 调用层parse_params 已由上方单测覆盖,此处补 complete 端到端)
///
/// `#[ignore]`:需真实 GLM 配置env var默认不跑。
/// 跑法:`GLM_BASE_URL=... GLM_API_KEY=... GLM_MODEL=glm-4-flash \
/// cargo test -p df-nodes --lib glm_live_complete -- --ignored --nocapture`
/// env var 缺失 → 跳过(非失败)。
#[ignore = "需真实 GLM 配置(env var)"]
#[tokio::test]
async fn glm_live_complete() {
let base_url = std::env::var("GLM_BASE_URL").ok().filter(|s| !s.is_empty());
let api_key = std::env::var("GLM_API_KEY").ok().filter(|s| !s.is_empty());
let (base_url, api_key) = match (base_url, api_key) {
(Some(b), Some(k)) => (b, k),
_ => {
eprintln!("跳过: 未设 GLM_BASE_URL / GLM_API_KEY env var");
return;
}
};
let model = std::env::var("GLM_MODEL").unwrap_or_else(|_| "glm-4-flash".to_string());
let provider: Box<dyn LlmProvider> =
df_ai::build_provider("openai_compat", &base_url, &api_key, &model);
let request = CompletionRequest {
model: model.clone(),
messages: vec![ChatMessage::user("只回复两个字:通过")],
temperature: Some(0.0),
max_tokens: Some(16),
stream: false,
tools: None,
tool_choice: None,
reasoning_content: None,
};
let response = provider.complete(request).await.expect("GLM 调用失败");
assert!(!response.text.is_empty(), "GLM 返回空文本");
println!(
"GLM 响应: model={}, text={}, usage={:?}",
response.model, response.text, response.usage
);
}
// 自审闭环相关单测(parse_review_json / truncate_for_summary / gate_should_block /
// build_review_prompt / update_field_writes_output_json)已随 AiSelfReviewNode 迁移至
// ai_self_review_node.rs(与被测代码同位,纯搬运)。
}