新增: AiNode provider_id注入链(resolve_provider三路径+FR-S1 api_key不进config)
This commit is contained in:
@@ -11,50 +11,160 @@ use std::sync::Arc;
|
|||||||
|
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
use df_ai::provider::{ChatMessage, CompletionRequest, LlmProvider};
|
use df_ai::provider::{ChatMessage, CompletionRequest, LlmProvider};
|
||||||
use df_storage::crud::TaskRepo;
|
use df_storage::crud::{AiProviderRepo, TaskRepo};
|
||||||
use df_storage::db::Database;
|
use df_storage::db::Database;
|
||||||
|
use df_storage::models::AiProviderRecord;
|
||||||
|
use df_storage::secret::{ensure_resolved_key, resolve_provider_secret};
|
||||||
use df_workflow::node::{Node, NodeContext, NodeOutput, NodeResult, NodeSchema};
|
use df_workflow::node::{Node, NodeContext, NodeOutput, NodeResult, NodeSchema};
|
||||||
|
|
||||||
/// AI 节点解析后的参数(execute 与参数解析解耦,便于单测覆盖取值/默认/校验逻辑)
|
/// AI 节点解析后的参数(execute 与参数解析解耦,便于单测覆盖取值/默认/校验逻辑)
|
||||||
|
///
|
||||||
|
/// provider 配置(base_url/api_key/protocol/default_model)经 `resolve_provider` 从 DB
|
||||||
|
/// ai_providers 表查 record + 经 df_storage::secret 解析密钥得到,**不进 config(FR-S1)**。
|
||||||
|
/// api_key 仅存于本结构体内存(AiNode 进程内存),不落 NodeContext.config / NodeOutput.data。
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
struct AiNodeParams {
|
struct AiNodeParams {
|
||||||
base_url: String,
|
/// 解析后的 provider 构造要素(含明文 api_key,仅 AiNode 内存可见)
|
||||||
api_key: String,
|
provider: ResolvedProvider,
|
||||||
prompt: String,
|
prompt: String,
|
||||||
system_prompt: Option<String>,
|
system_prompt: Option<String>,
|
||||||
model: String,
|
model: String,
|
||||||
temperature: Option<f32>,
|
temperature: Option<f32>,
|
||||||
max_tokens: Option<u32>,
|
max_tokens: Option<u32>,
|
||||||
/// 协议类型:openai_compat(默认)/ anthropic(GLM 订阅 / Claude 官方)
|
}
|
||||||
|
|
||||||
|
/// 经 `resolve_provider` 从 ai_providers 表 + df_storage::secret 解析后的 provider 构造要素。
|
||||||
|
///
|
||||||
|
/// api_key 字段:明文,FR-S1 下全程不出 AiNode 进程内存(不进 config/output/schema)。
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
struct ResolvedProvider {
|
||||||
|
/// 协议类型:openai_compat(默认)/ anthropic(GLM 订阅 / Claude 官方)— 从 record.provider_type 映射
|
||||||
protocol: String,
|
protocol: String,
|
||||||
/// model 为空时的占位,避免 provider 构造 panic
|
base_url: String,
|
||||||
|
/// 明文 api_key,经 `resolve_provider_secret`(DB 优先→keyring) 解析;仅 AiNode 内存可见
|
||||||
|
api_key: String,
|
||||||
|
/// model 为空时的占位(record.default_model 或 "gpt-4o-mini"),避免 provider 构造 panic
|
||||||
default_model: String,
|
default_model: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 从节点 config + 上游输入解析 AI 节点参数
|
/// 经 ai_providers 表 + df_storage::secret 解析 provider 构造要素(FR-S1 注入链核心)。
|
||||||
///
|
///
|
||||||
/// prompt 取值优先级:上游 `inputs["prompt"]` > `config.prompt`,两者皆无则报错。
|
/// 三路径(优先级从高到低):
|
||||||
/// model 为空时 default_model 兜底为 "gpt-4o-mini"。protocol 默认 openai_compat。
|
/// 1. **provider_id 优先**:config["provider_id"] 存在 → `AiProviderRepo::get_by_id` 查 record →
|
||||||
fn parse_params(
|
/// `resolve_provider_secret`(DB 优先→keyring) → `ensure_resolved_key` 空键早失败。
|
||||||
|
/// 2. **老路径兼容(过渡)**:config 显式含 `base_url` + `api_key` 明文(老调用方)→ 走原路径
|
||||||
|
/// + `tracing::warn!`(明文 api_key 经 config 注入已废弃)。兼容期保留,后续移除。
|
||||||
|
/// 3. **空兜底**:无 provider_id 也无明文 → 取 `is_default=true` 首条 provider(对齐
|
||||||
|
/// src-tauri commands/idea.rs:265-279 build_default_provider 模式),走同 resolve+ensure 链。
|
||||||
|
/// 无任何 provider → 友好错误「未配置 AI Provider」。
|
||||||
|
///
|
||||||
|
/// model:config["model"] 非空用之,否则 record.default_model,再否则 "gpt-4o-mini" 占位。
|
||||||
|
async fn resolve_provider(
|
||||||
|
db: &Arc<Database>,
|
||||||
config: &serde_json::Value,
|
config: &serde_json::Value,
|
||||||
inputs: &HashMap<String, NodeOutput>,
|
) -> anyhow::Result<ResolvedProvider> {
|
||||||
) -> anyhow::Result<AiNodeParams> {
|
let repo = AiProviderRepo::new(db);
|
||||||
// ── provider 配置(必填)──
|
let config_model = config
|
||||||
|
.get("model")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
|
.unwrap_or("")
|
||||||
|
.to_string();
|
||||||
|
|
||||||
|
// ── 路径 1:provider_id 优先 ──
|
||||||
|
if let Some(pid) = config.get("provider_id").and_then(|v| v.as_str()) {
|
||||||
|
if !pid.is_empty() {
|
||||||
|
let record = repo
|
||||||
|
.get_by_id(pid)
|
||||||
|
.await
|
||||||
|
.map_err(|e| anyhow::anyhow!("AiNode 查 provider 失败: {}", e))?
|
||||||
|
.ok_or_else(|| anyhow::anyhow!("AiNode provider_id={} 不存在", pid))?;
|
||||||
|
return resolve_from_record(&record, &config_model);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── 路径 2:老明文路径兼容(过渡,warn) ──
|
||||||
|
let has_plain_base = config.get("base_url").and_then(|v| v.as_str()).is_some();
|
||||||
|
let has_plain_key = config.get("api_key").and_then(|v| v.as_str()).is_some();
|
||||||
|
if has_plain_base && has_plain_key {
|
||||||
|
tracing::warn!(
|
||||||
|
"AiNode 明文 api_key/base_url 经 config 注入已废弃, 改用 provider_id (FR-S1). \
|
||||||
|
老路径将在后续版本移除"
|
||||||
|
);
|
||||||
let base_url = config
|
let base_url = config
|
||||||
.get("base_url")
|
.get("base_url")
|
||||||
.and_then(|v| v.as_str())
|
.and_then(|v| v.as_str())
|
||||||
.ok_or_else(|| anyhow::anyhow!("AiNode 缺少必填参数: base_url"))?
|
.unwrap_or("")
|
||||||
.to_string();
|
.to_string();
|
||||||
let api_key = config
|
let api_key = config
|
||||||
.get("api_key")
|
.get("api_key")
|
||||||
.and_then(|v| v.as_str())
|
.and_then(|v| v.as_str())
|
||||||
.ok_or_else(|| anyhow::anyhow!("AiNode 缺少必填参数: api_key"))?
|
.unwrap_or("")
|
||||||
.to_string();
|
.to_string();
|
||||||
// 空 key 早失败:避免空 key 发请求吃 401,错误伪装成"API Key 无效"(与 src-tauri secret::ensure_resolved_key 行为对齐)
|
ensure_resolved_key("(明文注入)", &api_key)
|
||||||
if api_key.trim().is_empty() {
|
.map_err(anyhow::Error::msg)?;
|
||||||
anyhow::bail!("AiNode api_key 为空(请检查节点配置或密钥注入链路),已阻止请求避免 401 误导");
|
let protocol = config
|
||||||
|
.get("protocol")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
|
.unwrap_or("openai_compat")
|
||||||
|
.to_string();
|
||||||
|
let default_model = if config_model.is_empty() {
|
||||||
|
"gpt-4o-mini".to_string()
|
||||||
|
} else {
|
||||||
|
config_model.clone()
|
||||||
|
};
|
||||||
|
return Ok(ResolvedProvider {
|
||||||
|
protocol,
|
||||||
|
base_url,
|
||||||
|
api_key,
|
||||||
|
default_model,
|
||||||
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ── 路径 3:兜底 is_default=true 首条 provider ──
|
||||||
|
let providers = repo
|
||||||
|
.list_all()
|
||||||
|
.await
|
||||||
|
.map_err(|e| anyhow::anyhow!("AiNode 列 provider 失败: {}", e))?;
|
||||||
|
let picked = providers
|
||||||
|
.iter()
|
||||||
|
.find(|p| p.is_default)
|
||||||
|
.cloned()
|
||||||
|
.or_else(|| providers.into_iter().next())
|
||||||
|
.ok_or_else(|| anyhow::anyhow!("未配置 AI Provider,请在设置中添加并保存密钥"))?;
|
||||||
|
resolve_from_record(&picked, &config_model)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 从 record 解析 provider 构造要素:resolve_provider_secret(DB 优先→keyring) + ensure 空键早失败。
|
||||||
|
/// protocol 从 record.provider_type 映射;default_model 取 config_model > record.default_model > 占位。
|
||||||
|
fn resolve_from_record(
|
||||||
|
record: &AiProviderRecord,
|
||||||
|
config_model: &str,
|
||||||
|
) -> anyhow::Result<ResolvedProvider> {
|
||||||
|
let api_key = resolve_provider_secret(record);
|
||||||
|
ensure_resolved_key(&record.name, &api_key).map_err(anyhow::Error::msg)?;
|
||||||
|
let default_model = if !config_model.is_empty() {
|
||||||
|
config_model.to_string()
|
||||||
|
} else if !record.default_model.is_empty() {
|
||||||
|
record.default_model.clone()
|
||||||
|
} else {
|
||||||
|
"gpt-4o-mini".to_string()
|
||||||
|
};
|
||||||
|
Ok(ResolvedProvider {
|
||||||
|
protocol: record.provider_type.clone(),
|
||||||
|
base_url: record.base_url.clone(),
|
||||||
|
api_key,
|
||||||
|
default_model,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 从节点 config + 上游输入解析 AI 节点 prompt 与可选参数(provider 经 `resolve_provider` 异步解析)。
|
||||||
|
///
|
||||||
|
/// prompt 取值优先级:上游 `inputs["prompt"]` > `config.prompt`,两者皆无则报错。
|
||||||
|
fn parse_params(
|
||||||
|
config: &serde_json::Value,
|
||||||
|
inputs: &HashMap<String, NodeOutput>,
|
||||||
|
provider: ResolvedProvider,
|
||||||
|
) -> anyhow::Result<AiNodeParams> {
|
||||||
// ── prompt(必填):优先取上游节点 "prompt" 输出,回退 config.prompt ──
|
// ── prompt(必填):优先取上游节点 "prompt" 输出,回退 config.prompt ──
|
||||||
let prompt = inputs
|
let prompt = inputs
|
||||||
.get("prompt")
|
.get("prompt")
|
||||||
@@ -86,28 +196,14 @@ fn parse_params(
|
|||||||
.get("system_prompt")
|
.get("system_prompt")
|
||||||
.and_then(|v| v.as_str())
|
.and_then(|v| v.as_str())
|
||||||
.map(|s| s.to_string());
|
.map(|s| s.to_string());
|
||||||
let protocol = config
|
|
||||||
.get("protocol")
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.unwrap_or("openai_compat")
|
|
||||||
.to_string();
|
|
||||||
// default_model:留空时给一个占位,避免 provider 构造 panic
|
|
||||||
let default_model = if model.is_empty() {
|
|
||||||
"gpt-4o-mini".to_string()
|
|
||||||
} else {
|
|
||||||
model.clone()
|
|
||||||
};
|
|
||||||
|
|
||||||
Ok(AiNodeParams {
|
Ok(AiNodeParams {
|
||||||
base_url,
|
provider,
|
||||||
api_key,
|
|
||||||
prompt,
|
prompt,
|
||||||
system_prompt,
|
system_prompt,
|
||||||
model,
|
model,
|
||||||
temperature,
|
temperature,
|
||||||
max_tokens,
|
max_tokens,
|
||||||
protocol,
|
|
||||||
default_model,
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -132,10 +228,16 @@ impl Node for AiNode {
|
|||||||
async fn execute(&self, ctx: NodeContext) -> NodeResult {
|
async fn execute(&self, ctx: NodeContext) -> NodeResult {
|
||||||
tracing::info!("AiNode 执行: node_id={}", ctx.node_id);
|
tracing::info!("AiNode 执行: node_id={}", ctx.node_id);
|
||||||
|
|
||||||
let p = parse_params(&ctx.config, &ctx.inputs)?;
|
// FR-S1 注入链:provider 经 df_storage::secret 在 AiNode 内存解析,api_key 不进 config。
|
||||||
|
let provider_cfg = resolve_provider(&self.db, &ctx.config).await?;
|
||||||
|
let p = parse_params(&ctx.config, &ctx.inputs, provider_cfg)?;
|
||||||
|
|
||||||
let provider: Box<dyn LlmProvider> =
|
let provider: Box<dyn LlmProvider> = df_ai::build_provider(
|
||||||
df_ai::build_provider(&p.protocol, &p.base_url, &p.api_key, &p.default_model);
|
&p.provider.protocol,
|
||||||
|
&p.provider.base_url,
|
||||||
|
&p.provider.api_key,
|
||||||
|
&p.provider.default_model,
|
||||||
|
);
|
||||||
|
|
||||||
// ── 构建消息 ──
|
// ── 构建消息 ──
|
||||||
let mut messages = Vec::with_capacity(2);
|
let mut messages = Vec::with_capacity(2);
|
||||||
@@ -156,8 +258,8 @@ impl Node for AiNode {
|
|||||||
|
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"AiNode 调用 LLM: model={}, base_url={}",
|
"AiNode 调用 LLM: model={}, base_url={}",
|
||||||
p.default_model,
|
p.provider.default_model,
|
||||||
p.base_url
|
p.provider.base_url
|
||||||
);
|
);
|
||||||
let response = provider.complete(request).await?;
|
let response = provider.complete(request).await?;
|
||||||
|
|
||||||
@@ -204,14 +306,14 @@ impl Node for AiNode {
|
|||||||
"properties": {
|
"properties": {
|
||||||
"prompt": { "type": "string", "description": "用户提示词(若无则取上游 prompt 输出)" },
|
"prompt": { "type": "string", "description": "用户提示词(若无则取上游 prompt 输出)" },
|
||||||
"system_prompt": { "type": "string", "description": "系统提示词(可选)" },
|
"system_prompt": { "type": "string", "description": "系统提示词(可选)" },
|
||||||
"protocol": { "type": "string", "description": "协议类型:openai_compat(默认)或 anthropic(GLM订阅/Claude官方)" },
|
"provider_id": { "type": "string", "description": "AI Provider ID(FR-S1:密钥经 df_storage::secret 解析,不进 config;留空走默认 provider)" },
|
||||||
"base_url": { "type": "string", "description": "API 地址,OpenAI 兼容如 https://api.deepseek.com;Anthropic 如 https://open.bigmodel.cn/api/anthropic" },
|
"model": { "type": "string", "description": "模型名(可选,留空用 record.default_model)" },
|
||||||
"api_key": { "type": "string", "description": "API 密钥" },
|
|
||||||
"model": { "type": "string", "description": "模型名(可选,留空用默认)" },
|
|
||||||
"temperature": { "type": "number", "description": "温度 0.0~2.0(可选)" },
|
"temperature": { "type": "number", "description": "温度 0.0~2.0(可选)" },
|
||||||
"max_tokens": { "type": "integer", "description": "最大生成 token(可选,anthropic 协议无值时默认 4096)" }
|
"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 解析" }
|
||||||
},
|
},
|
||||||
"required": ["base_url", "api_key"]
|
"required": ["provider_id"]
|
||||||
}),
|
}),
|
||||||
output: serde_json::json!({
|
output: serde_json::json!({
|
||||||
"type": "object",
|
"type": "object",
|
||||||
@@ -327,7 +429,9 @@ impl Node for AiSelfReviewNode {
|
|||||||
async fn execute(&self, ctx: NodeContext) -> NodeResult {
|
async fn execute(&self, ctx: NodeContext) -> NodeResult {
|
||||||
tracing::info!("AiSelfReviewNode 执行: node_id={}", ctx.node_id);
|
tracing::info!("AiSelfReviewNode 执行: node_id={}", ctx.node_id);
|
||||||
|
|
||||||
let p = parse_params(&ctx.config, &ctx.inputs)?;
|
// FR-S1 注入链:provider 经 df_storage::secret 在 AiNode 内存解析,api_key 不进 config。
|
||||||
|
let provider_cfg = resolve_provider(&self.db, &ctx.config).await?;
|
||||||
|
let p = parse_params(&ctx.config, &ctx.inputs, provider_cfg)?;
|
||||||
|
|
||||||
// ── 读任务(需求 + 产出) ──
|
// ── 读任务(需求 + 产出) ──
|
||||||
// task_id 必填(自审对象是具体任务的产出);缺失报错而非静默跳过。
|
// task_id 必填(自审对象是具体任务的产出);缺失报错而非静默跳过。
|
||||||
@@ -354,8 +458,12 @@ impl Node for AiSelfReviewNode {
|
|||||||
|
|
||||||
let prompt = Self::build_review_prompt(&task.description, &output_text);
|
let prompt = Self::build_review_prompt(&task.description, &output_text);
|
||||||
|
|
||||||
let provider: Box<dyn LlmProvider> =
|
let provider: Box<dyn LlmProvider> = df_ai::build_provider(
|
||||||
df_ai::build_provider(&p.protocol, &p.base_url, &p.api_key, &p.default_model);
|
&p.provider.protocol,
|
||||||
|
&p.provider.base_url,
|
||||||
|
&p.provider.api_key,
|
||||||
|
&p.provider.default_model,
|
||||||
|
);
|
||||||
|
|
||||||
let request = CompletionRequest {
|
let request = CompletionRequest {
|
||||||
model: p.model.clone(),
|
model: p.model.clone(),
|
||||||
@@ -372,8 +480,8 @@ impl Node for AiSelfReviewNode {
|
|||||||
|
|
||||||
tracing::info!(
|
tracing::info!(
|
||||||
"AiSelfReviewNode 调用 LLM: model={}, base_url={}",
|
"AiSelfReviewNode 调用 LLM: model={}, base_url={}",
|
||||||
p.default_model,
|
p.provider.default_model,
|
||||||
p.base_url
|
p.provider.base_url
|
||||||
);
|
);
|
||||||
let response = provider.complete(request).await?;
|
let response = provider.complete(request).await?;
|
||||||
|
|
||||||
@@ -454,13 +562,11 @@ impl Node for AiSelfReviewNode {
|
|||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": {
|
"properties": {
|
||||||
"task_id": { "type": "string", "description": "自审目标任务 ID(必填)" },
|
"task_id": { "type": "string", "description": "自审目标任务 ID(必填)" },
|
||||||
"protocol": { "type": "string", "description": "协议类型:openai_compat(默认)或 anthropic" },
|
"provider_id": { "type": "string", "description": "AI Provider ID(FR-S1:密钥经 secret 解析不进 config;留空走默认 provider)" },
|
||||||
"base_url": { "type": "string" },
|
"model": { "type": "string", "description": "模型名(可选,留空用 record.default_model)" },
|
||||||
"api_key": { "type": "string" },
|
|
||||||
"model": { "type": "string" },
|
|
||||||
"max_tokens": { "type": "integer" }
|
"max_tokens": { "type": "integer" }
|
||||||
},
|
},
|
||||||
"required": ["task_id", "base_url", "api_key"]
|
"required": ["task_id", "provider_id"]
|
||||||
}),
|
}),
|
||||||
output: serde_json::json!({
|
output: serde_json::json!({
|
||||||
"type": "object",
|
"type": "object",
|
||||||
@@ -483,13 +589,24 @@ impl Node for AiSelfReviewNode {
|
|||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
use df_storage::crud::ProjectRepo;
|
||||||
|
use df_storage::models::{ProjectRecord, TaskRecord};
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
|
|
||||||
/// 造带基础三字段(base_url/api_key/prompt)的 config,overrides 覆盖或追加
|
/// 测试用 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(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 造带基础 prompt 的 config(provider 经 resolve_provider 解析,parse_params 只看 prompt),
|
||||||
|
/// overrides 覆盖或追加。
|
||||||
fn config_with(overrides: serde_json::Value) -> serde_json::Value {
|
fn config_with(overrides: serde_json::Value) -> serde_json::Value {
|
||||||
let mut base = json!({
|
let mut base = json!({
|
||||||
"base_url": "https://api.example.com",
|
|
||||||
"api_key": "sk-test",
|
|
||||||
"prompt": "config-prompt"
|
"prompt": "config-prompt"
|
||||||
});
|
});
|
||||||
if let (serde_json::Value::Object(b), serde_json::Value::Object(o)) = (&mut base, overrides) {
|
if let (serde_json::Value::Object(b), serde_json::Value::Object(o)) = (&mut base, overrides) {
|
||||||
@@ -504,30 +621,19 @@ mod tests {
|
|||||||
HashMap::new()
|
HashMap::new()
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn missing_base_url_errors() {
|
|
||||||
let config = json!({ "api_key": "k", "prompt": "p" });
|
|
||||||
let err = parse_params(&config, &empty_inputs()).unwrap_err().to_string();
|
|
||||||
assert!(err.contains("base_url"), "缺 base_url 应报错, 实际: {}", err);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn missing_api_key_errors() {
|
|
||||||
let config = json!({ "base_url": "http://x", "prompt": "p" });
|
|
||||||
let err = parse_params(&config, &empty_inputs()).unwrap_err().to_string();
|
|
||||||
assert!(err.contains("api_key"), "缺 api_key 应报错, 实际: {}", err);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn missing_prompt_errors() {
|
fn missing_prompt_errors() {
|
||||||
let config = json!({ "base_url": "http://x", "api_key": "k" });
|
let config = json!({});
|
||||||
let err = parse_params(&config, &empty_inputs()).unwrap_err().to_string();
|
let err = parse_params(&config, &empty_inputs(), provider_stub())
|
||||||
|
.unwrap_err()
|
||||||
|
.to_string();
|
||||||
assert!(err.contains("prompt"), "缺 prompt 应报错, 实际: {}", err);
|
assert!(err.contains("prompt"), "缺 prompt 应报错, 实际: {}", err);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn prompt_from_config_when_no_upstream() {
|
fn prompt_from_config_when_no_upstream() {
|
||||||
let p = parse_params(&config_with(json!({})), &empty_inputs()).unwrap();
|
let p =
|
||||||
|
parse_params(&config_with(json!({})), &empty_inputs(), provider_stub()).unwrap();
|
||||||
assert_eq!(p.prompt, "config-prompt");
|
assert_eq!(p.prompt, "config-prompt");
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -538,42 +644,175 @@ mod tests {
|
|||||||
"prompt".to_string(),
|
"prompt".to_string(),
|
||||||
NodeOutput::from_value(json!("upstream-prompt")),
|
NodeOutput::from_value(json!("upstream-prompt")),
|
||||||
);
|
);
|
||||||
let p = parse_params(&config_with(json!({})), &inputs).unwrap();
|
let p =
|
||||||
|
parse_params(&config_with(json!({})), &inputs, provider_stub()).unwrap();
|
||||||
assert_eq!(p.prompt, "upstream-prompt", "上游输入应优先于 config.prompt");
|
assert_eq!(p.prompt, "upstream-prompt", "上游输入应优先于 config.prompt");
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn defaults_when_optional_fields_missing() {
|
|
||||||
let p = parse_params(&config_with(json!({})), &empty_inputs()).unwrap();
|
|
||||||
assert_eq!(p.model, "");
|
|
||||||
assert_eq!(p.default_model, "gpt-4o-mini", "model 空时 default_model 兜底");
|
|
||||||
assert_eq!(p.protocol, "openai_compat", "protocol 默认 openai_compat");
|
|
||||||
assert_eq!(p.temperature, None);
|
|
||||||
assert_eq!(p.max_tokens, None);
|
|
||||||
assert_eq!(p.system_prompt, None);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn optional_fields_parsed_when_present() {
|
fn optional_fields_parsed_when_present() {
|
||||||
let p = parse_params(
|
let p = parse_params(
|
||||||
&config_with(json!({
|
&config_with(json!({
|
||||||
"model": "glm-4",
|
"model": "glm-4",
|
||||||
"protocol": "anthropic",
|
|
||||||
"temperature": 0.3,
|
"temperature": 0.3,
|
||||||
"max_tokens": 1024,
|
"max_tokens": 1024,
|
||||||
"system_prompt": "你是助手"
|
"system_prompt": "你是助手"
|
||||||
})),
|
})),
|
||||||
&empty_inputs(),
|
&empty_inputs(),
|
||||||
|
provider_stub(),
|
||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
assert_eq!(p.model, "glm-4");
|
assert_eq!(p.model, "glm-4");
|
||||||
assert_eq!(p.default_model, "glm-4", "model 非空时 default_model = model");
|
assert_eq!(p.provider.protocol, "openai_compat");
|
||||||
assert_eq!(p.protocol, "anthropic");
|
|
||||||
assert_eq!(p.temperature, Some(0.3));
|
assert_eq!(p.temperature, Some(0.3));
|
||||||
assert_eq!(p.max_tokens, Some(1024));
|
assert_eq!(p.max_tokens, Some(1024));
|
||||||
assert_eq!(p.system_prompt.as_deref(), Some("你是助手"));
|
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,
|
||||||
|
is_default,
|
||||||
|
config: None,
|
||||||
|
created_at: "0".to_string(),
|
||||||
|
updated_at: "0".to_string(),
|
||||||
|
})
|
||||||
|
.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 端到端)
|
/// 真调 GLM 验证 provider 调用层(parse_params 已由上方单测覆盖,此处补 complete 端到端)
|
||||||
///
|
///
|
||||||
/// `#[ignore]`:需真实 GLM 配置(env var),默认不跑。
|
/// `#[ignore]`:需真实 GLM 配置(env var),默认不跑。
|
||||||
@@ -619,8 +858,6 @@ mod tests {
|
|||||||
|
|
||||||
/// 内存 DB 构造 helper:插父项目 + 指定 task,返回 (db, task_id)。
|
/// 内存 DB 构造 helper:插父项目 + 指定 task,返回 (db, task_id)。
|
||||||
async fn setup_task_db(task_output_json: Option<&str>) -> (Database, String) {
|
async fn setup_task_db(task_output_json: Option<&str>) -> (Database, String) {
|
||||||
use df_storage::crud::ProjectRepo;
|
|
||||||
use df_storage::models::{ProjectRecord, TaskRecord};
|
|
||||||
let db = Database::open_in_memory().await.expect("open_in_memory");
|
let db = Database::open_in_memory().await.expect("open_in_memory");
|
||||||
ProjectRepo::new(&db)
|
ProjectRepo::new(&db)
|
||||||
.insert(ProjectRecord {
|
.insert(ProjectRecord {
|
||||||
|
|||||||
Reference in New Issue
Block a user