后端: - 工作流推进链(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 + 总线/技术债审查新文档)
424 lines
18 KiB
Rust
424 lines
18 KiB
Rust
//! 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 ID(FR-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_id;FR-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 的 config(provider 经 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())
|
||
}
|
||
|
||
/// 路径 1:provider_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"
|
||
);
|
||
}
|
||
|
||
/// 路径 1:config.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");
|
||
}
|
||
|
||
/// 路径 1:provider_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);
|
||
}
|
||
|
||
/// 路径 1:provider_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(与被测代码同位,纯搬运)。
|
||
}
|