新增: 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 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::models::AiProviderRecord;
|
||||
use df_storage::secret::{ensure_resolved_key, resolve_provider_secret};
|
||||
use df_workflow::node::{Node, NodeContext, NodeOutput, NodeResult, NodeSchema};
|
||||
|
||||
/// 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)]
|
||||
struct AiNodeParams {
|
||||
base_url: String,
|
||||
api_key: String,
|
||||
/// 解析后的 provider 构造要素(含明文 api_key,仅 AiNode 内存可见)
|
||||
provider: ResolvedProvider,
|
||||
prompt: String,
|
||||
system_prompt: Option<String>,
|
||||
model: String,
|
||||
temperature: Option<f32>,
|
||||
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,
|
||||
/// 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,
|
||||
}
|
||||
|
||||
/// 从节点 config + 上游输入解析 AI 节点参数
|
||||
/// 经 ai_providers 表 + df_storage::secret 解析 provider 构造要素(FR-S1 注入链核心)。
|
||||
///
|
||||
/// prompt 取值优先级:上游 `inputs["prompt"]` > `config.prompt`,两者皆无则报错。
|
||||
/// model 为空时 default_model 兜底为 "gpt-4o-mini"。protocol 默认 openai_compat。
|
||||
fn parse_params(
|
||||
/// 三路径(优先级从高到低):
|
||||
/// 1. **provider_id 优先**:config["provider_id"] 存在 → `AiProviderRepo::get_by_id` 查 record →
|
||||
/// `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,
|
||||
inputs: &HashMap<String, NodeOutput>,
|
||||
) -> anyhow::Result<AiNodeParams> {
|
||||
// ── provider 配置(必填)──
|
||||
) -> anyhow::Result<ResolvedProvider> {
|
||||
let repo = AiProviderRepo::new(db);
|
||||
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
|
||||
.get("base_url")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| anyhow::anyhow!("AiNode 缺少必填参数: base_url"))?
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let api_key = config
|
||||
.get("api_key")
|
||||
.and_then(|v| v.as_str())
|
||||
.ok_or_else(|| anyhow::anyhow!("AiNode 缺少必填参数: api_key"))?
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
// 空 key 早失败:避免空 key 发请求吃 401,错误伪装成"API Key 无效"(与 src-tauri secret::ensure_resolved_key 行为对齐)
|
||||
if api_key.trim().is_empty() {
|
||||
anyhow::bail!("AiNode api_key 为空(请检查节点配置或密钥注入链路),已阻止请求避免 401 误导");
|
||||
ensure_resolved_key("(明文注入)", &api_key)
|
||||
.map_err(anyhow::Error::msg)?;
|
||||
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 ──
|
||||
let prompt = inputs
|
||||
.get("prompt")
|
||||
@@ -86,28 +196,14 @@ fn parse_params(
|
||||
.get("system_prompt")
|
||||
.and_then(|v| v.as_str())
|
||||
.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 {
|
||||
base_url,
|
||||
api_key,
|
||||
provider,
|
||||
prompt,
|
||||
system_prompt,
|
||||
model,
|
||||
temperature,
|
||||
max_tokens,
|
||||
protocol,
|
||||
default_model,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -132,10 +228,16 @@ impl Node for AiNode {
|
||||
async fn execute(&self, ctx: NodeContext) -> NodeResult {
|
||||
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> =
|
||||
df_ai::build_provider(&p.protocol, &p.base_url, &p.api_key, &p.default_model);
|
||||
let provider: Box<dyn LlmProvider> = df_ai::build_provider(
|
||||
&p.provider.protocol,
|
||||
&p.provider.base_url,
|
||||
&p.provider.api_key,
|
||||
&p.provider.default_model,
|
||||
);
|
||||
|
||||
// ── 构建消息 ──
|
||||
let mut messages = Vec::with_capacity(2);
|
||||
@@ -156,8 +258,8 @@ impl Node for AiNode {
|
||||
|
||||
tracing::info!(
|
||||
"AiNode 调用 LLM: model={}, base_url={}",
|
||||
p.default_model,
|
||||
p.base_url
|
||||
p.provider.default_model,
|
||||
p.provider.base_url
|
||||
);
|
||||
let response = provider.complete(request).await?;
|
||||
|
||||
@@ -204,14 +306,14 @@ impl Node for AiNode {
|
||||
"properties": {
|
||||
"prompt": { "type": "string", "description": "用户提示词(若无则取上游 prompt 输出)" },
|
||||
"system_prompt": { "type": "string", "description": "系统提示词(可选)" },
|
||||
"protocol": { "type": "string", "description": "协议类型:openai_compat(默认)或 anthropic(GLM订阅/Claude官方)" },
|
||||
"base_url": { "type": "string", "description": "API 地址,OpenAI 兼容如 https://api.deepseek.com;Anthropic 如 https://open.bigmodel.cn/api/anthropic" },
|
||||
"api_key": { "type": "string", "description": "API 密钥" },
|
||||
"model": { "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)" }
|
||||
"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!({
|
||||
"type": "object",
|
||||
@@ -327,7 +429,9 @@ impl Node for AiSelfReviewNode {
|
||||
async fn execute(&self, ctx: NodeContext) -> NodeResult {
|
||||
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 必填(自审对象是具体任务的产出);缺失报错而非静默跳过。
|
||||
@@ -354,8 +458,12 @@ impl Node for AiSelfReviewNode {
|
||||
|
||||
let prompt = Self::build_review_prompt(&task.description, &output_text);
|
||||
|
||||
let provider: Box<dyn LlmProvider> =
|
||||
df_ai::build_provider(&p.protocol, &p.base_url, &p.api_key, &p.default_model);
|
||||
let provider: Box<dyn LlmProvider> = df_ai::build_provider(
|
||||
&p.provider.protocol,
|
||||
&p.provider.base_url,
|
||||
&p.provider.api_key,
|
||||
&p.provider.default_model,
|
||||
);
|
||||
|
||||
let request = CompletionRequest {
|
||||
model: p.model.clone(),
|
||||
@@ -372,8 +480,8 @@ impl Node for AiSelfReviewNode {
|
||||
|
||||
tracing::info!(
|
||||
"AiSelfReviewNode 调用 LLM: model={}, base_url={}",
|
||||
p.default_model,
|
||||
p.base_url
|
||||
p.provider.default_model,
|
||||
p.provider.base_url
|
||||
);
|
||||
let response = provider.complete(request).await?;
|
||||
|
||||
@@ -454,13 +562,11 @@ impl Node for AiSelfReviewNode {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"task_id": { "type": "string", "description": "自审目标任务 ID(必填)" },
|
||||
"protocol": { "type": "string", "description": "协议类型:openai_compat(默认)或 anthropic" },
|
||||
"base_url": { "type": "string" },
|
||||
"api_key": { "type": "string" },
|
||||
"model": { "type": "string" },
|
||||
"provider_id": { "type": "string", "description": "AI Provider ID(FR-S1:密钥经 secret 解析不进 config;留空走默认 provider)" },
|
||||
"model": { "type": "string", "description": "模型名(可选,留空用 record.default_model)" },
|
||||
"max_tokens": { "type": "integer" }
|
||||
},
|
||||
"required": ["task_id", "base_url", "api_key"]
|
||||
"required": ["task_id", "provider_id"]
|
||||
}),
|
||||
output: serde_json::json!({
|
||||
"type": "object",
|
||||
@@ -483,13 +589,24 @@ impl Node for AiSelfReviewNode {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use df_storage::crud::ProjectRepo;
|
||||
use df_storage::models::{ProjectRecord, TaskRecord};
|
||||
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 {
|
||||
let mut base = json!({
|
||||
"base_url": "https://api.example.com",
|
||||
"api_key": "sk-test",
|
||||
"prompt": "config-prompt"
|
||||
});
|
||||
if let (serde_json::Value::Object(b), serde_json::Value::Object(o)) = (&mut base, overrides) {
|
||||
@@ -504,30 +621,19 @@ mod tests {
|
||||
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]
|
||||
fn missing_prompt_errors() {
|
||||
let config = json!({ "base_url": "http://x", "api_key": "k" });
|
||||
let err = parse_params(&config, &empty_inputs()).unwrap_err().to_string();
|
||||
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()).unwrap();
|
||||
let p =
|
||||
parse_params(&config_with(json!({})), &empty_inputs(), provider_stub()).unwrap();
|
||||
assert_eq!(p.prompt, "config-prompt");
|
||||
}
|
||||
|
||||
@@ -538,42 +644,175 @@ mod tests {
|
||||
"prompt".to_string(),
|
||||
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");
|
||||
}
|
||||
|
||||
#[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]
|
||||
fn optional_fields_parsed_when_present() {
|
||||
let p = parse_params(
|
||||
&config_with(json!({
|
||||
"model": "glm-4",
|
||||
"protocol": "anthropic",
|
||||
"temperature": 0.3,
|
||||
"max_tokens": 1024,
|
||||
"system_prompt": "你是助手"
|
||||
})),
|
||||
&empty_inputs(),
|
||||
provider_stub(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(p.model, "glm-4");
|
||||
assert_eq!(p.default_model, "glm-4", "model 非空时 default_model = model");
|
||||
assert_eq!(p.protocol, "anthropic");
|
||||
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,
|
||||
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 端到端)
|
||||
///
|
||||
/// `#[ignore]`:需真实 GLM 配置(env var),默认不跑。
|
||||
@@ -619,8 +858,6 @@ mod tests {
|
||||
|
||||
/// 内存 DB 构造 helper:插父项目 + 指定 task,返回 (db, task_id)。
|
||||
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");
|
||||
ProjectRepo::new(&db)
|
||||
.insert(ProjectRecord {
|
||||
|
||||
Reference in New Issue
Block a user