//! 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 在 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, } impl AiNode { /// 构造节点(注册表工厂调用,注入数据库句柄) pub fn new(db: Arc) -> 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 = 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 { 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 = 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(与被测代码同位,纯搬运)。 }