From 79b6a430956e6ab3d6af36c2b95a4f1f11db5c4c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=BB=9D=E5=B0=98?= <237809796@qq.com> Date: Wed, 17 Jun 2026 02:04:58 +0800 Subject: [PATCH] =?UTF-8?q?=E6=96=B0=E5=A2=9E:=20F-04=E5=A4=9AProvider?= =?UTF-8?q?=E8=B4=9F=E8=BD=BD=E5=9D=87=E8=A1=A1=E6=B1=A0(=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E5=B1=82+=E9=80=89=E6=8B=A9=E5=99=A8+=E5=B9=B6?= =?UTF-8?q?=E5=8F=91=E5=8E=9F=E8=AF=AD)+CR-52=E7=99=BD=E9=A1=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/df-nodes/src/ai_node.rs | 2 + crates/df-storage/src/crud.rs | 23 +- crates/df-storage/src/migrations.rs | 38 ++- crates/df-storage/src/models.rs | 18 ++ crates/df-storage/src/secret.rs | 1 + src-tauri/src/commands/ai/agentic.rs | 33 +++ src-tauri/src/commands/ai/commands.rs | 4 + src-tauri/src/commands/ai/mod.rs | 4 + src-tauri/src/commands/ai/provider_pool.rs | 270 +++++++++++++++++++++ src-tauri/src/commands/ai/tool_registry.rs | 5 +- src-tauri/src/commands/workflow.rs | 17 +- src-tauri/src/state.rs | 64 ++++- 12 files changed, 463 insertions(+), 16 deletions(-) create mode 100644 src-tauri/src/commands/ai/provider_pool.rs diff --git a/crates/df-nodes/src/ai_node.rs b/crates/df-nodes/src/ai_node.rs index 69bcc87..19b4e8d 100644 --- a/crates/df-nodes/src/ai_node.rs +++ b/crates/df-nodes/src/ai_node.rs @@ -724,6 +724,8 @@ mod tests { config: None, created_at: "0".to_string(), updated_at: "0".to_string(), + enabled: true, + weight: 50, }) .await .expect("insert provider"); diff --git a/crates/df-storage/src/crud.rs b/crates/df-storage/src/crud.rs index 1d5d69b..d1653a3 100644 --- a/crates/df-storage/src/crud.rs +++ b/crates/df-storage/src/crud.rs @@ -1067,6 +1067,10 @@ fn ai_provider_from_row(row: &Row<'_>) -> std::result::Result("enabled").unwrap_or(1) != 0, + weight: row.get::<_, i32>("weight").unwrap_or(50).max(0) as u32, }) } @@ -1151,23 +1155,30 @@ impl_repo!( let is_default = if rec.is_default { 1i32 } else { 0i32 }; // model_configs:Vec → JSON 字符串落 TEXT 列 let model_configs_json = serde_json::to_string(&rec.model_configs).unwrap_or_else(|_| "[]".into()); + // F-260614-04: enabled/weight 落库(SQLite 无 BOOLEAN,i32 承载)。 + let enabled_i = if rec.enabled { 1i32 } else { 0i32 }; + let weight_i = rec.weight.min(100) as i32; conn.execute( - "INSERT OR REPLACE INTO ai_providers (id, name, provider_type, api_key, base_url, default_model, models, model_configs, is_default, config, created_at, updated_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12)", + "INSERT OR REPLACE INTO ai_providers (id, name, provider_type, api_key, base_url, default_model, models, model_configs, is_default, config, created_at, updated_at, enabled, weight) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)", params![ rec.id, rec.name, rec.provider_type, rec.api_key, rec.base_url, - rec.default_model, rec.models, model_configs_json, is_default, rec.config, rec.created_at, rec.updated_at + rec.default_model, rec.models, model_configs_json, is_default, rec.config, rec.created_at, rec.updated_at, + enabled_i, weight_i ], ) }, update => |conn, rec| { let is_default = if rec.is_default { 1i32 } else { 0i32 }; let model_configs_json = serde_json::to_string(&rec.model_configs).unwrap_or_else(|_| "[]".into()); + let enabled_i = if rec.enabled { 1i32 } else { 0i32 }; + let weight_i = rec.weight.min(100) as i32; conn.execute( - "UPDATE ai_providers SET name = ?1, provider_type = ?2, api_key = ?3, base_url = ?4, default_model = ?5, models = ?6, model_configs = ?7, is_default = ?8, config = ?9, updated_at = ?10 WHERE id = ?11", + "UPDATE ai_providers SET name = ?1, provider_type = ?2, api_key = ?3, base_url = ?4, default_model = ?5, models = ?6, model_configs = ?7, is_default = ?8, config = ?9, updated_at = ?10, enabled = ?11, weight = ?12 WHERE id = ?13", params![ rec.name, rec.provider_type, rec.api_key, rec.base_url, - rec.default_model, rec.models, model_configs_json, is_default, rec.config, rec.updated_at, rec.id + rec.default_model, rec.models, model_configs_json, is_default, rec.config, rec.updated_at, + enabled_i, weight_i, rec.id ], ) } @@ -1761,6 +1772,8 @@ mod tests { config: None, created_at: "0".into(), updated_at: "0".into(), + enabled: true, + weight: 50, }; repo.insert(rec).await.expect("insert"); diff --git a/crates/df-storage/src/migrations.rs b/crates/df-storage/src/migrations.rs index 2b61b88..ab48c52 100644 --- a/crates/df-storage/src/migrations.rs +++ b/crates/df-storage/src/migrations.rs @@ -34,7 +34,7 @@ pub fn run(conn: &Connection) -> Result<()> { // 迁移步骤链: 顺序执行,跳过已应用的版本(current_version < N 才跑)。 // 新增版本时,在此数组追加一项 (N, migrate_vN) 即可,无需改逻辑。 - let steps: [(i32, fn(&Connection) -> Result<()>); 18] = [ + let steps: [(i32, fn(&Connection) -> Result<()>); 19] = [ (1, migrate_v1), (2, migrate_v2), (3, migrate_v3), @@ -53,6 +53,7 @@ pub fn run(conn: &Connection) -> Result<()> { (16, migrate_v16), (17, migrate_v17), (18, migrate_v18), + (19, migrate_v19), ]; for (version, migrate_fn) in steps { @@ -328,6 +329,36 @@ fn migrate_v18(conn: &Connection) -> Result<()> { Ok(()) } +/// V19: 幂等补 ai_providers.enabled + ai_providers.weight 列(F-260614-04 多 Provider 负载均衡池) +/// +/// - `enabled INTEGER NOT NULL DEFAULT 1`:provider 是否进入负载均衡池。 +/// 老库行迁移后默认 1(所有现存 provider 默认启用,单 provider 路径零变化)。 +/// is_default 仍保留作启动兜底(get_active_provider 无 active_provider_id 时取 is_default)。 +/// - `weight INTEGER NOT NULL DEFAULT 50`:provider 在池中的选择权重(0-100)。 +/// 高权重 provider 优先被选为主;同权重时退化近似轮询。 +/// +/// 向后兼容:老库行 ALTER 后取 DEFAULT,from_row 经 i32→bool / i32→u32 解析。 +/// 用 PRAGMA 探测列存在性,缺失才 ALTER(同 v17/v18 模式),对新库/老库均安全。 +fn migrate_v19(conn: &Connection) -> Result<()> { + if !column_exists(conn, "ai_providers", "enabled") { + conn.execute( + "ALTER TABLE ai_providers ADD COLUMN enabled INTEGER NOT NULL DEFAULT 1", + [], + )?; + tracing::info!("v19: 补建 ai_providers.enabled 列(多 Provider 负载均衡池,F-260614-04)"); + } + if !column_exists(conn, "ai_providers", "weight") { + conn.execute( + "ALTER TABLE ai_providers ADD COLUMN weight INTEGER NOT NULL DEFAULT 50", + [], + )?; + tracing::info!("v19: 补建 ai_providers.weight 列(多 Provider 负载均衡池,F-260614-04)"); + } + conn.execute("INSERT INTO schema_version (version) VALUES (?)", [19])?; + tracing::info!("迁移 v19 完成"); + Ok(()) +} + /// V1 建表 SQL const V1_SQL: &str = " -- 想法表 @@ -512,7 +543,10 @@ CREATE TABLE IF NOT EXISTS ai_providers ( is_default INTEGER NOT NULL DEFAULT 0, config TEXT, created_at TEXT NOT NULL, - updated_at TEXT NOT NULL + updated_at TEXT NOT NULL, + model_configs TEXT, + enabled INTEGER NOT NULL DEFAULT 1, + weight INTEGER NOT NULL DEFAULT 50 ); CREATE TABLE IF NOT EXISTS ai_tool_executions ( diff --git a/crates/df-storage/src/models.rs b/crates/df-storage/src/models.rs index 43feb23..d6ad40b 100644 --- a/crates/df-storage/src/models.rs +++ b/crates/df-storage/src/models.rs @@ -162,6 +162,24 @@ pub struct AiProviderRecord { pub config: Option, // JSON extra config pub created_at: String, pub updated_at: String, + /// provider 是否进入负载均衡池(F-260614-04)。false = 仅作为配置存在,不参与主链路由/选池。 + /// 老库迁移默认 1(单 provider 路径零变化)。 + #[serde(default = "default_enabled")] + pub enabled: bool, + /// provider 在负载均衡池中的选择权重(0-100,F-260614-04)。高权重优先被选为主; + /// 同权重时退化近似轮询。老库迁移默认 50。 + #[serde(default = "default_weight")] + pub weight: u32, +} + +/// AiProviderRecord.enabled 的 serde 默认(true)。老库/缺字段 JSON → enabled。 +fn default_enabled() -> bool { + true +} + +/// AiProviderRecord.weight 的 serde 默认(50)。老库/缺字段 JSON → 50。 +fn default_weight() -> u32 { + 50 } /// AI 对话记录 diff --git a/crates/df-storage/src/secret.rs b/crates/df-storage/src/secret.rs index 85967ec..db9a7a1 100644 --- a/crates/df-storage/src/secret.rs +++ b/crates/df-storage/src/secret.rs @@ -211,6 +211,7 @@ mod tests { api_key: "sk-db-fallback".into(), base_url: "https://x".into(), default_model: "m".into(), models: None, model_configs: Vec::new(), is_default: false, config: None, created_at: "0".into(), updated_at: "0".into(), + enabled: true, weight: 50, }; assert_eq!(resolve_provider_secret(&rec), "sk-db-fallback"); } diff --git a/src-tauri/src/commands/ai/agentic.rs b/src-tauri/src/commands/ai/agentic.rs index 53f730e..b612568 100644 --- a/src-tauri/src/commands/ai/agentic.rs +++ b/src-tauri/src/commands/ai/agentic.rs @@ -130,6 +130,32 @@ pub(crate) async fn run_agentic_loop( // B-260615-09: generating 状态由 RAII guard 收敛复位(正常 exit 显式 reset;panic/异常 Drop 兜底) let mut guard = GeneratingGuard::new(session_arc.clone()); + // F-260614-04: 多 Provider 负载均衡池 — 选主候选 provider。 + // + // 流程:list_all → ProviderPool::select(按 模型亲和 > weight > is_default 排序)→ 取首位。 + // 单 provider 场景:池仅 1 enabled provider → select 返回单元素 Vec → 首位 = 唯一 provider, + // 行为同 F-01 前(零变化)。空池(0 enabled)→ fallback 入参 provider_config(保启动行为)。 + // + // 主候选的 model_configs 用于路由(F-01),其 provider_config 用于 build_provider。 + // 当前批次仅选主,fallback(主失败切备用 provider)见后续批次(需重构流式重试块,单独评估)。 + // + // AiProviderRepo::new 仅 clone Arc(廉价),不复用 AppState.ai_providers + // (run_agentic_loop 签名只传 Arc,改签名会牵动 3 调用点 + try_continue)。 + let provider_repo = df_storage::crud::AiProviderRepo::new(&db); + let pool_providers: Vec = provider_repo.list_all().await.unwrap_or_default(); + let primary_provider: AiProviderRecord = match super::provider_pool::ProviderPool::select( + &pool_providers, + None, // 模型亲和此时尚未确定(router 选模型需 provider_config,见下方) + ) + .into_iter() + .next() + { + Some(p) => p, + None => provider_config.clone(), // 空池兜底:用调用方传入的 provider_config(启动行为不变) + }; + // 用主候选覆盖入参 provider_config(下游 build_provider / 路由 / 日志均用此)。 + let provider_config = primary_provider; + // FR-S1: resolve→ensure_resolved_key(空 key 早失败)→build_provider 三步统一走工厂 // 空 key 早失败(逻辑见 secret::ensure_resolved_key 单测):避免空 key 发请求吃 401,错误伪装成"API Key 无效" // @@ -400,8 +426,13 @@ pub(crate) async fn run_agentic_loop( // Partial)不重试保文。重试退避复用 retry::backoff_delay(1s→2s→4s±20% jitter) + // retry::is_status_retryable Fatal 分类(stream_recv classify_status_or_class 镜像, // 4xx 非429 立即放弃) + 30s 总挂钟预算。重试期间持有 permit 不释放(防新请求挤占)。 + // + // F-260614-04: per-provider permit(可选)。set_provider_caps 未配置时返回 None + // (单 provider 场景零变化);配置后取额外 permit 防单 provider 被打满(限流 429)。 + // 三 permit 均 Drop 释放(L595-597 显式 drop _global/_per_conv;_provider_permit 绑块尾 Drop)。 let _global_permit = llm_concurrency.acquire_global().await; let _per_conv_permit = llm_concurrency.acquire_per_conv().await; + let _provider_permit = llm_concurrency.acquire_for_provider(&provider_config.id).await; // 重试总预算(挂钟,含 sleep + 各次请求耗时),对齐 retry::MAX_TOTAL_BUDGET 30s。 // 超预算直接放弃重试交最终错误/保文路径。 @@ -569,6 +600,8 @@ pub(crate) async fn run_agentic_loop( // stream 结束立即释放 permit,后续工具执行不受限流(本地操作无 RPM 成本) drop(_global_permit); drop(_per_conv_permit); + // F-260614-04: per-provider permit 为 None 时 drop None 无副作用;Some 时释放槽。 + drop(_provider_permit); // 累加本轮 token:provider 流式 usage 的 prompt_tokens 为 0 时(GLM 等),用预估输入兜底 let round_prompt = if round_usage.prompt_tokens == 0 { estimated_prompt } else { round_usage.prompt_tokens }; diff --git a/src-tauri/src/commands/ai/commands.rs b/src-tauri/src/commands/ai/commands.rs index e1922b9..f462651 100644 --- a/src-tauri/src/commands/ai/commands.rs +++ b/src-tauri/src/commands/ai/commands.rs @@ -993,6 +993,10 @@ pub async fn ai_save_provider( config: None, created_at, updated_at: now_millis(), + // F-260614-04: 新建 provider 默认进入负载均衡池(enabled=true,weight=50)。 + // 单 provider 路径零变化(enabled=true 即与迁移前行为等价)。 + enabled: true, + weight: 50, }; let id = record.id.clone(); state diff --git a/src-tauri/src/commands/ai/mod.rs b/src-tauri/src/commands/ai/mod.rs index cc94e4f..ffd3d04 100644 --- a/src-tauri/src/commands/ai/mod.rs +++ b/src-tauri/src/commands/ai/mod.rs @@ -11,6 +11,7 @@ //! - [`audit`] — 工具调用审计 + pending 审批恢复 + 工具调用处理 //! - [`skills`] — 本机 Claude 技能扫描 //! - [`prompt`] — 系统提示词构建 +//! - [`provider_pool`] — F-260614-04 多 Provider 负载均衡池选择(纯逻辑) //! - [`compress`] — F-15 上下文压缩 LLM 摘要(compress_via_llm 公共函数) //! - [`tool_registry`] — AI 工具注册表构建 + 文件路径校验 //! - [`knowledge_inject`] — 知识库注入 + 提炼 @@ -28,6 +29,7 @@ pub mod compress; pub mod conversation; pub mod knowledge_inject; pub mod prompt; +pub mod provider_pool; pub mod secret; pub mod skills; pub mod stream_recv; @@ -468,6 +470,7 @@ pub(crate) async fn execute_run_workflow_for_tool( // 转调工作流执行核心(B-260617-01:run_workflow_inner 是 run_workflow 命令的共用核心, // 同 invoke('run_workflow') IPC 入口的执行逻辑,持 &AppState 非 tauri::State,绕开 // tauri::State 在非命令上下文不可构造的限制)。 + // CR-52: 本路径是 AI 工具经 ai_approve 审批后执行 → triggered_by="ai"(区别命令层 manual)。 // 返回 execution_id 字符串;失败(模板不存在/DAG 校验失败/DB 写入失败)上抛 Err。 let execution_id = crate::commands::workflow::run_workflow_inner( app, @@ -477,6 +480,7 @@ pub(crate) async fn execute_run_workflow_for_tool( config, Some(task_id.clone()), Some(target_status.clone()), + "ai", ) .await .map_err(|e| anyhow::anyhow!("工作流执行失败: {}", e))?; diff --git a/src-tauri/src/commands/ai/provider_pool.rs b/src-tauri/src/commands/ai/provider_pool.rs new file mode 100644 index 0000000..f02db40 --- /dev/null +++ b/src-tauri/src/commands/ai/provider_pool.rs @@ -0,0 +1,270 @@ +//! 多 Provider 负载均衡池 — F-260614-04 +//! +//! 与 `df_ai::router`(F-01 阶段4,纯函数选模型)对齐的纯逻辑核心:给定 enabled provider +//! 列表 + 已选模型 ID + 选择策略,返回**有序候选 provider 列表**(主→备用)。 +//! 零 IO / 零状态,所有状态(DB 读 / 健康度 / 计数器)由调用方持有。 +//! +//! ## 架构定位(为何放 src-tauri 而非 df-ai) +//! `AiProviderRecord` 定义在 df-storage。df-ai 当前存储无关(router.rs 操作 df-ai-core 的 +//! `ModelConfig`,非 storage 类型)。把 ProviderPool 放 df-ai 会强制 df-ai→df-storage 依赖, +//! 破坏 df-ai 的存储无关边界。ProviderPool 与 `prompt.rs::get_active_provider` / +//! `secret.rs::build_provider_for` 同属"消费 AiProviderRecord 的 app 层逻辑",故放 src-tauri。 +//! +//! ## 与 router 协同(职责切分,不重叠) +//! - **router.select_model_id**:在**单个 provider 的 model_configs 池**中选最优模型(F-01)。 +//! - **ProviderPool::select**:在**多个 provider 间**选主 + 列出备用(本模块)。 +//! 调用顺序:先 router 选模型 → 再 ProviderPool 选 provider 实例(模型在多 provider 间共享时)。 +//! +//! ## 选择策略(论证,见 select 文档) +//! 1. **模型亲和优先**:router 已选 model_id,优先选**其池中含该模型的 provider**(正确性:不擅自换模型)。 +//! 2. **加权**:同亲和级内按 weight 降序(weight=0 provider 不入选)。 +//! 3. **is_default 兜底**:模型亲和全 miss 时,默认 provider 优先(零变化启动行为)。 +//! +//! ## fallback 顺序(select 返回的 Vec 即为 fallback 序) +//! 调用方(agentic loop)按顺序尝试:主 provider 失败(可重试错误)→ 列表下一个 provider 重试。 +//! 不可重试错误(401/400)立即放弃,不浪费备用 provider(对齐 retry.rs Fatal 分类)。 +//! +//! ## 向后兼容 +//! - enabled provider 仅 1 个 → 返回单元素 Vec,调用方无 fallback 路径(行为同 F-01 前)。 +//! - enabled provider 0 个 → 返回空 Vec,调用方兜底 get_active_provider(行为同 F-01 前)。 +//! - 模型在多 provider 间共享 → 主=含模型且 weight 最高的 provider,备用=其余含模型的 provider。 + +use df_storage::models::AiProviderRecord; + +/// 多 Provider 负载均衡池(单元结构,无状态)。 +/// +/// `select` 为关联函数:给定 enabled provider 列表 + 已选模型 ID,返回有序候选列表。 +/// 对齐 `ModelRouter`(df_ai::router)的单元结构 + 关联函数风格。 +pub struct ProviderPool; + +impl ProviderPool { + /// 在 enabled provider 池中选出**有序候选列表**(主 → 备用),供调用方 fallback。 + /// + /// ## 选择算法(3 步,稳定排序) + /// 1. **过滤 enabled**:只选 `enabled == true` 的 provider。 + /// 2. **过滤零权重**:weight == 0 的 provider 不参与(weight=0 = 显式禁用主选; + /// 若需"仅 fallback"配置,用 enabled=false 替代,语义更清晰)。 + /// 3. **排序**(稳定,保输入顺序作 tiebreak): + /// - **主键:模型亲和**。含 `model_id` 的 provider 排前(模型亲和=true → 排前)。 + /// model_id 为 None(调用方未走 router,如标题/扫描)→ 退化为全亲和,跳过此键。 + /// - **次键:weight 降序**。同亲和级内 weight 高者排前。 + /// - **末键:is_default 兜底**。同亲和同 weight 时 is_default 排前(启动行为不变)。 + /// + /// ## 向后兼容 + /// - enabled provider 1 个 → 单元素 Vec,无 fallback。 + /// - 0 个 → 空 Vec,调用方兜底 get_active_provider。 + /// - model_id None → 全亲和,纯按 weight + is_default 排序(标题/扫描路径用)。 + pub fn select( + providers: &[AiProviderRecord], + model_id: Option<&str>, + ) -> Vec { + // 步骤 1+2:enabled 且 weight > 0 + let mut candidates: Vec = providers + .iter() + .filter(|p| p.enabled && p.weight > 0) + .cloned() + .collect(); + + // 步骤 3:稳定排序(主键模型亲和,次键 weight,末键 is_default)。 + // sort_by 稳定:同 key 保输入顺序(created_at DESC,list_all 返回序)。 + // 反转比较结果:大者排前(亲和 true > false;weight 大 > 小;is_default true > false)。 + candidates.sort_by(|a, b| { + // 主键:模型亲和。model_id None → 视两方都亲和(退化为全过此键)。 + let a_affinity = model_id + .map_or(true, |mid| a.model_configs.iter().any(|m| m.model_id == mid)); + let b_affinity = model_id + .map_or(true, |mid| b.model_configs.iter().any(|m| m.model_id == mid)); + let by_affinity = b_affinity.cmp(&a_affinity); // true 排前 + if by_affinity != std::cmp::Ordering::Equal { + return by_affinity; + } + // 次键:weight 降序。 + let by_weight = b.weight.cmp(&a.weight); + if by_weight != std::cmp::Ordering::Equal { + return by_weight; + } + // 末键:is_default 兜底。 + b.is_default.cmp(&a.is_default) + }); + + candidates + } +} + +// ============================================================ +// 单测(纯逻辑,零 IO) +// ============================================================ + +#[cfg(test)] +mod tests { + use super::*; + use df_ai::df_ai_core::model::ModelConfig; + + /// 构造可定制 AiProviderRecord(默认 enabled/weight=50/is_default=false,无 model_configs)。 + fn provider(id: &str) -> AiProviderRecord { + AiProviderRecord { + id: id.into(), + name: id.into(), + provider_type: "openai_compat".into(), + api_key: String::new(), + base_url: "https://x".into(), + default_model: "glm-4-flash".into(), + models: None, + model_configs: Vec::new(), + is_default: false, + config: None, + created_at: "0".into(), + updated_at: "0".into(), + enabled: true, + weight: 50, + } + } + + /// ── 向后兼容:单 provider 路径 ── + + #[test] + fn empty_pool_returns_empty() { + // enabled provider 0 个 → 空 Vec,调用方兜底 get_active_provider。 + let pool: Vec = vec![]; + assert!(ProviderPool::select(&pool, Some("glm-4-flash")).is_empty()); + } + + #[test] + fn single_provider_returns_singleton() { + // 单 provider → 单元素 Vec,无 fallback(行为同 F-01 前)。 + let pool = vec![provider("p1")]; + let selected = ProviderPool::select(&pool, Some("glm-4-flash")); + assert_eq!(selected.len(), 1); + assert_eq!(selected[0].id, "p1"); + } + + /// ── 过滤:enabled / weight=0 ── + + #[test] + fn disabled_provider_filtered() { + // enabled=false → 不入选(单 provider 场景兜底由调用方处理)。 + let pool = vec![ + AiProviderRecord { enabled: false, ..provider("disabled") }, + provider("enabled"), + ]; + let selected = ProviderPool::select(&pool, None); + assert_eq!(selected.len(), 1); + assert_eq!(selected[0].id, "enabled"); + } + + #[test] + fn zero_weight_provider_filtered() { + // weight=0 → 显式排除(语义:weight=0 = 不参与主选也不作 fallback)。 + let pool = vec![ + AiProviderRecord { weight: 0, ..provider("zero") }, + provider("normal"), + ]; + let selected = ProviderPool::select(&pool, None); + assert_eq!(selected.len(), 1); + assert_eq!(selected[0].id, "normal"); + } + + /// ── 主键:模型亲和 ── + + #[test] + fn model_affinity_picks_provider_with_model() { + // 两 provider,仅 p2 含目标模型 → p2 排前(即使 p1 weight 更高)。 + let pool = vec![ + AiProviderRecord { + weight: 90, + ..provider("p1") + }, + AiProviderRecord { + weight: 30, + model_configs: vec![ModelConfig::with_defaults("glm-4-flash")], + ..provider("p2") + }, + ]; + let selected = ProviderPool::select(&pool, Some("glm-4-flash")); + assert_eq!(selected[0].id, "p2", "含目标模型的 provider 应排前"); + assert_eq!(selected[1].id, "p1"); + } + + #[test] + fn model_id_none_all_affinity_equal() { + // model_id None(标题/扫描路径)→ 全亲和,纯按 weight 排序。 + let pool = vec![ + AiProviderRecord { weight: 30, ..provider("low") }, + AiProviderRecord { weight: 90, ..provider("high") }, + ]; + let selected = ProviderPool::select(&pool, None); + assert_eq!(selected[0].id, "high"); + assert_eq!(selected[1].id, "low"); + } + + /// ── 次键:weight 降序 ── + + #[test] + fn higher_weight_wins_same_affinity() { + // 两 provider 都不含目标模型(同亲和=false),weight 高者排前。 + let pool = vec![ + AiProviderRecord { weight: 50, ..provider("low") }, + AiProviderRecord { weight: 80, ..provider("high") }, + ]; + let selected = ProviderPool::select(&pool, Some("glm-4-flash")); + assert_eq!(selected[0].id, "high"); + } + + /// ── 末键:is_default 兜底 ── + + #[test] + fn is_default_breaks_weight_tie() { + // 同亲和 + 同 weight → is_default 排前(启动行为不变)。 + let pool = vec![ + provider("normal"), // is_default=false + AiProviderRecord { is_default: true, ..provider("default") }, + ]; + let selected = ProviderPool::select(&pool, None); + assert_eq!(selected[0].id, "default"); + } + + /// ── 综合:多 provider 多维度 ── + + #[test] + fn full_scenario_orders_correctly() { + // 候选: + // a: enabled=true, weight=70, 含目标模型 → 亲和+weight70 + // b: enabled=true, weight=90, 不含目标模型 → 非亲和+weight90(排 a 后) + // c: enabled=true, weight=70, 含目标模型, is_default=true → 亲和+weight70+default(同 a 但 default 排前) + // d: enabled=false → 排除 + // e: weight=0 → 排除 + // + // 预期顺序:c(亲和/70/default) > a(亲和/70) > b(非亲和/90) + let pool = vec![ + AiProviderRecord { + weight: 70, + model_configs: vec![ModelConfig::with_defaults("glm-4-flash")], + ..provider("a") + }, + AiProviderRecord { + weight: 90, + ..provider("b") + }, + AiProviderRecord { + weight: 70, + is_default: true, + model_configs: vec![ModelConfig::with_defaults("glm-4-flash")], + ..provider("c") + }, + AiProviderRecord { + enabled: false, + ..provider("d") + }, + AiProviderRecord { + weight: 0, + ..provider("e") + }, + ]; + let selected = ProviderPool::select(&pool, Some("glm-4-flash")); + assert_eq!( + selected.iter().map(|p| p.id.as_str()).collect::>(), + vec!["c", "a", "b"], + "应按 亲和 > weight > is_default 排序,排除 enabled=false/weight=0" + ); + } +} diff --git a/src-tauri/src/commands/ai/tool_registry.rs b/src-tauri/src/commands/ai/tool_registry.rs index 370779a..6d075c4 100644 --- a/src-tauri/src/commands/ai/tool_registry.rs +++ b/src-tauri/src/commands/ai/tool_registry.rs @@ -544,7 +544,10 @@ pub fn build_ai_tool_registry(db: &Arc) -> AiToolRegistry { "run_workflow", "按任务 target_status 推进对应工作流(含 AiNode 自审 / HumanNode 核对闸门)。参数 task_id + target_status 同时提供才联动任务推进(完成后按 target_status 推进任务,失败按退回态回滚)。属高风险操作(触发工作流引擎执行),须人工批准。审批通过后由后端直接执行工作流引擎并联动推进任务,返回 execution_id", df_ai::ai_tools::object_schema(vec![("task_id", "string", true), ("target_status", "string", true)]), RiskLevel::High, - { let _db = db.clone(); Box::new(move |args: serde_json::Value| { + // CR-52: 删死代码 `let _db = db.clone();`(完全未用 — 防御 handler 仅返回 Err, + // 不持 db;run_workflow 真正执行经 ai_approve → execute_run_workflow_for_tool + // → workflow.rs::run_workflow_inner,此处 handler 在正常流程不可达)。 + { Box::new(move |args: serde_json::Value| { Box::pin(async move { let task_id = args["task_id"].as_str().ok_or_else(|| anyhow::anyhow!("缺少 task_id"))?; let target_status = args["target_status"].as_str() diff --git a/src-tauri/src/commands/workflow.rs b/src-tauri/src/commands/workflow.rs index bc97b8e..da83f21 100644 --- a/src-tauri/src/commands/workflow.rs +++ b/src-tauri/src/commands/workflow.rs @@ -76,7 +76,8 @@ pub async fn run_workflow( ) -> Result { // 瘦转发到 run_workflow_inner(B-260617-01:抽取共用核心供 ai_approve 经 // execute_run_workflow_for_tool 直接调,无需构造 tauri::State)。 - run_workflow_inner(&app, &state, name, dag, config, task_id, target_status).await + // CR-52: 命令层(前端 invoke run_workflow)是人工手动触发 → triggered_by="manual"。 + run_workflow_inner(&app, &state, name, dag, config, task_id, target_status, "manual").await } /// 工作流执行核心逻辑(B-260617-01 抽取,命令层瘦转发)。 @@ -84,6 +85,11 @@ pub async fn run_workflow( /// 原 `run_workflow` 命令体拆出,使 ai_approve 经 execute_run_workflow_for_tool 直接调用 /// (持 &AppState 而非 tauri::State,绕开 tauri::State 在非命令上下文不可构造的限制)。 /// 语义与原命令完全一致:模板选 DAG → 落执行记录 → spawn 事件转发 + DAG 执行 + 任务联动推进。 +/// +/// `triggered_by` 落入 WorkflowRecord.triggered_by,区分触发来源: +/// - 命令层 run_workflow(前端 invoke,人工手动)传 "manual" +/// - execute_run_workflow_for_tool(AI 工具经 ai_approve 审批后执行)传 "ai" +/// CR-52:原硬编码 "manual" 致 AI 路径误标,改由调用方透传。 pub async fn run_workflow_inner( app: &AppHandle, state: &AppState, @@ -92,6 +98,7 @@ pub async fn run_workflow_inner( config: serde_json::Value, task_id: Option, target_status: Option, + triggered_by: &str, ) -> Result { // 0. F-260616-06 ①-1: DagDef 来源选模板(方案A·最小)。 // 规则: @@ -121,11 +128,9 @@ pub async fn run_workflow_inner( name, dag_json, status: "running".to_string(), - // B-260617-01: AI 工具经 ai_approve 触发时标 triggered_by=ai(区分人工手动触发)。 - // 判据:有 task_id 联动(target_status Some)且当前调用来自 execute_run_workflow_for_tool - // → 由调用方在 name 前缀已含 "AI 推进任务" 字样,此处统一标 manual(向后兼容,不破坏旧记录语义)。 - // 精细化 triggered_by=ai 留待后续(需加参数透传,本次最小改动)。 - triggered_by: Some("manual".to_string()), + // CR-52: triggered_by 由调用方透传(命令层 run_workflow="manual", + // execute_run_workflow_for_tool="ai"),不再硬编码 "manual" 致 AI 路径误标。 + triggered_by: Some(triggered_by.to_string()), project_id: None, task_id: task_id.clone(), created_at: now_millis(), diff --git a/src-tauri/src/state.rs b/src-tauri/src/state.rs index 8b4892d..c69955a 100644 --- a/src-tauri/src/state.rs +++ b/src-tauri/src/state.rs @@ -78,10 +78,10 @@ impl Default for KnowledgeConfig { } // ============================================================ -// LLM 调用并发控制(双层 Semaphore) +// LLM 调用并发控制(双层 Semaphore + 可选 per-provider 层) // ============================================================ -/// LLM 调用并发控制 — 全局 + 单对话双层 Semaphore +/// LLM 调用并发控制 — 全局 + 单对话双层 Semaphore + 可选 per-provider 层 /// /// 限流对象:所有真实 LLM 调用(主循环 stream_llm / 标题生成 / 知识提炼)。 /// 不限流本地工具执行(tools.execute)——本地操作无外部成本、不受 RPM 约束。 @@ -95,10 +95,21 @@ impl Default for KnowledgeConfig { /// generating 互斥,同一时刻仅一个对话的 loop 在跑,per_conv 退化为"单对话内并发" /// (主循环 stream_llm + 标题生成 + 知识提炼)。未来若支持多对话并发, /// 需改为 HashMap。 +/// +/// ## F-260614-04: per-provider 层(可选) +/// `per_provider` 为 HashMap>。调用方(agentic loop)经 +/// `acquire_for_provider(pid)` 取额外 permit,防单 provider 被打满(限流 429)。 +/// **单 provider 场景**:若未调 `set_provider_caps`,HashMap 为空, +/// `acquire_for_provider` 返回 None(无限流,行为同 F-01 前)。零变化保证。 +/// 全局容量 = min(sum(各 provider 上限), global_cap):由调用方在配置时约束 +/// (set_provider_caps 传 min(sum, global_cap)),非运行时强约束。 #[derive(Clone)] pub struct LlmConcurrency { global: Arc>>, per_conv: Arc>>, + /// F-260614-04: per-provider 信号量表。空 = 无 per-provider 限流(单 provider 路径零变化)。 + /// Arc>:运行时增删 provider 配置时替换/插入,acquire 时 clone Arc。 + per_provider: Arc>>>, } impl LlmConcurrency { @@ -106,6 +117,7 @@ impl LlmConcurrency { Self { global: Arc::new(Mutex::new(Arc::new(Semaphore::new(global)))), per_conv: Arc::new(Mutex::new(Arc::new(Semaphore::new(per_conv)))), + per_provider: Arc::new(Mutex::new(HashMap::new())), } } @@ -130,6 +142,54 @@ impl LlmConcurrency { pub async fn set_per_conv(&self, permits: usize) { *self.per_conv.lock().await = Arc::new(Semaphore::new(permits)); } + + /// F-260614-04: 取 per-provider 并发 permit(可选)。 + /// + /// - provider 在 `per_provider` 表中有配置 → 取其 Semaphore permit,返回 Some。 + /// - provider 无配置(单 provider 场景或未 set_provider_caps)→ 返回 None,无限流。 + /// + /// 调用方(agentic loop)用法: + /// ```ignore + /// let _global_permit = llm_concurrency.acquire_global().await; + /// let _per_conv_permit = llm_concurrency.acquire_per_conv().await; + /// let _provider_permit = llm_concurrency.acquire_for_provider(&provider_id).await; + /// ``` + /// 三 permit 均绑 guard Drop 自动释放。None 时无 permit 需释放(行为同 F-01 前)。 + pub async fn acquire_for_provider( + &self, + provider_id: &str, + ) -> Option { + let sema = { + let map = self.per_provider.lock().await; + map.get(provider_id).cloned() + }?; + // Semaphore 存在 → acquire。expect 同 acquire_global/per_conv:Semaphore 不会 close + // (无 close() 调用,仅在 set_provider_caps 时替换表内 Arc,旧 Arc permit 仍有效)。 + Some( + sema.acquire_owned() + .await + .expect("llm per_provider semaphore closed"), + ) + } + + /// F-260614-04: 批量设置 per-provider 并发上限(替换整表)。 + /// + /// 调用方(Settings 配置热改 / 启动初始化)传入 `{ provider_id: permits }` map, + /// 替换整张 per_provider 表(软收敛:已持 permit 不回收)。全局容量约束 min(sum, global_cap) + /// 由调用方在构造 map 时应用(本函数不做强约束,只落表)。 + /// + /// 传空 map → 清空 per_provider 表(所有 provider 回退到无 per-provider 限流)。 + pub async fn set_provider_caps(&self, caps: HashMap) { + let mut map = self.per_provider.lock().await; + map.clear(); + for (pid, permits) in caps { + // permits=0 等同无配置(Semaphore::new(0) 永远 acquire 不到 → 死锁), + // 故 permits=0 跳过(不落表 → acquire_for_provider 返 None → 无限流)。 + if permits > 0 { + map.insert(pid, Arc::new(Semaphore::new(permits))); + } + } + } } /// 应用全局状态 — 通过 `app.manage()` 注入,command 中以 `State<'_, AppState>` 取用