diff --git a/src/gateway/agent_factory.rs b/src/gateway/agent_factory.rs index 1991290..61c5c94 100644 --- a/src/gateway/agent_factory.rs +++ b/src/gateway/agent_factory.rs @@ -15,6 +15,7 @@ use crate::observability::Observer; use crate::skills::{SkillPromptProvider, SkillRuntime}; use crate::storage::PromptInjectionRepository; use crate::storage::persistent_session_id; +use crate::storage::SessionStore; use crate::tools::task::runtime::{SubagentPromptProvider, SubagentRuntime}; use crate::tools::task::SubagentResult; use crate::tools::{ToolContext, ToolRegistry, WaitCoordinator}; @@ -56,8 +57,12 @@ pub(crate) struct AgentFactory { prompt_repository: Arc, /// Provider/Model 解析器:按专家 frontmatter 中的 provider/model 字段覆盖基础配置 model_resolver: Arc, - /// per-session 的用户模型选择(最高优先级,覆盖专家配置) + /// per-session 的用户模型选择(覆盖专家配置) model_selections: Arc, + /// per-topic 的用户模型选择(最高优先级;物化后话题模型不再随 session 选择漂移) + topic_model_selections: Arc, + /// 持久化存储:session 级选择首次被话题命中时物化回写 topics 行 + store: Arc, /// 上下文压缩算法配置(所有 agent 共享) compaction_config: CompactionConfig, /// 可观测性 Observer(依赖注入到 AgentLoop,业务层不感知具体实现) @@ -96,6 +101,8 @@ impl AgentFactory { prompt_repository: Arc, model_resolver: Arc, model_selections: Arc, + topic_model_selections: Arc, + store: Arc, compaction_config: CompactionConfig, observer: Option>, ) -> Self { @@ -115,6 +122,8 @@ impl AgentFactory { prompt_repository, model_resolver, model_selections, + topic_model_selections, + store, compaction_config, observer, instance_id, @@ -164,9 +173,20 @@ impl AgentFactory { _ => request.provider_config.clone(), }; - // 按用户手动选择的 provider/model 覆盖(最高优先级,覆盖专家配置)。 + // 用户手动选择的 provider/model 覆盖(最高优先级,覆盖专家配置)。 + // 优先级:topic 级选择 > session 级选择;均未设置时保持专家/基础配置。 + // 物化规则:session 级选择首次被话题命中时回写 topics 行固化——此后该话题的 + // 模型只能被"在该话题内显式改选"改变,不再随 session 级选择漂移; + // 专家/config 默认不物化(保持继承活性)。 // 引用不存在的 provider/model 名时报错并阻止会话(用户主动选择,配置错误应明确反馈)。 - let effective_provider_config = match self.model_selections.get(&session_id) { + let topic_selection = request + .topic_id + .as_deref() + .and_then(|tid| self.topic_model_selections.get(tid)); + let from_topic = topic_selection.is_some(); + let user_selection = topic_selection.or_else(|| self.model_selections.get(&session_id)); + + let effective_provider_config = match user_selection { Some((user_provider, user_model)) if user_provider.is_some() || user_model.is_some() => { @@ -178,9 +198,32 @@ impl AgentFactory { &expert_provider_config, ) .map_err(|e| AgentError::Other(e.to_string()))?; + + // 物化:命中 session 级选择且话题无固化值时,将解析后的具体 + // (provider, model) 写入 topics 行(持久化 + 内存缓存) + if !from_topic { + if let Some(tid) = request.topic_id.as_deref() { + let provider = resolved.name.clone(); + let model = resolved.model_id.clone(); + self.topic_model_selections + .set(tid, Some(provider.clone()), Some(model.clone())); + if let Err(err) = + self.store.update_topic_model(tid, Some(&provider), Some(&model)) + { + tracing::warn!( + error = %err, + topic_id = %tid, + "AgentFactory: failed to materialize topic model selection" + ); + } + } + } + tracing::info!( instance_id = self.instance_id, session_id = %session_id, + topic_id = request.topic_id.as_deref().unwrap_or(""), + source = if from_topic { "topic" } else { "session" }, provider = %resolved.name, model_id = %resolved.model_id, "AgentFactory: applied user model override" diff --git a/src/tools/task/runtime.rs b/src/tools/task/runtime.rs index 2676c3c..da65a71 100644 --- a/src/tools/task/runtime.rs +++ b/src/tools/task/runtime.rs @@ -462,6 +462,10 @@ pub struct DefaultSubAgentRuntime { /// task_id → CancellationToken 映射,用于取消传播 /// Arc 包装以便 spawned task 完成后清理自身条目 cancel_registry: Arc>>, + /// per-session 的用户模型选择(子代理 def 未显式设定模型时继承) + model_selections: Option>, + /// per-topic 的用户模型选择(优先于 session 级继承) + topic_model_selections: Option>, } impl DefaultSubAgentRuntime { @@ -476,6 +480,8 @@ impl DefaultSubAgentRuntime { bus: Option>, store: Arc, skills: Arc, + model_selections: Option>, + topic_model_selections: Option>, ) -> Self { let max_concurrent = config.max_concurrent.max(1); let semaphore = Arc::new(tokio::sync::Semaphore::new(max_concurrent)); @@ -492,6 +498,8 @@ impl DefaultSubAgentRuntime { skills, semaphore, cancel_registry: Arc::new(parking_lot::Mutex::new(HashMap::new())), + model_selections, + topic_model_selections, } } @@ -502,6 +510,55 @@ impl DefaultSubAgentRuntime { .ok_or_else(|| format!("subagent type '{}' is disabled or not found", type_name)) } + /// 解析子代理最终生效的 provider 配置(create_subagent 与 spawn 共用,保证一致)。 + /// + /// 优先级:def frontmatter > 话题级用户选择 > session 级用户选择 > 全局基础配置。 + /// def 仅设定部分字段时,缺失字段从继承链补齐(而非全局), + /// 与主代理"专家覆盖 → 用户选择再覆盖"的解析语义对称。 + fn resolve_effective_provider_config( + &self, + session: &TaskSession, + def_name: Option<&str>, + def_provider: Option<&str>, + def_model: Option<&str>, + ) -> Result { + let resolve_err = |scope: &str, e: crate::config::ConfigError| { + TaskError::AgentCreationFailed(format!( + "subagent '{}' {} model resolution failed: {}", + def_name.unwrap_or("?"), + scope, + e + )) + }; + + // 继承链:topic 级 > session 级(key 为主代理的 persistent session id) + let inherited = session + .parent_topic_id + .as_deref() + .and_then(|tid| self.topic_model_selections.as_ref().and_then(|s| s.get(tid))) + .or_else(|| { + self.model_selections + .as_ref() + .and_then(|s| s.get(&session.parent_session_id)) + }); + + let base = match inherited { + Some((p, m)) if p.is_some() || m.is_some() => self + .model_resolver + .resolve(p.as_deref(), m.as_deref(), &self.provider_config) + .map_err(|e| resolve_err("inherited", e))?, + _ => self.provider_config.clone(), + }; + + match (def_provider.is_some(), def_model.is_some()) { + (true, _) | (_, true) => self + .model_resolver + .resolve(def_provider, def_model, &base) + .map_err(|e| resolve_err("", e)), + _ => Ok(base), + } + } + /// 获取实际执行时间 fn effective_max_execution_secs(&self, def: &SubagentDef) -> u64 { def.max_execution_secs @@ -582,25 +639,14 @@ impl DefaultSubAgentRuntime { let child_depth = parent_nesting_depth + 1; let tools = self.build_subagent_tools_registry(def, child_depth); - // 按 def 中的 provider/model 字段解析覆盖基础 provider_config。 + // 解析优先级:def frontmatter > topic 级用户选择 > session 级用户选择 > 全局基础配置。 // 引用不存在的 provider/model 名时返回错误(反馈给 LLM 重试,与 def 缺失即拒绝的安全范式一致)。 - let effective_provider_config = match def { - Some(d) if d.provider.is_some() || d.model.is_some() => self - .model_resolver - .resolve( - d.provider.as_deref(), - d.model.as_deref(), - &self.provider_config, - ) - .map_err(|e| { - TaskError::AgentCreationFailed(format!( - "subagent '{}' model resolution failed: {}", - def.map(|d| d.name.as_str()).unwrap_or("?"), - e - )) - })?, - _ => self.provider_config.clone(), - }; + let effective_provider_config = self.resolve_effective_provider_config( + session, + def.map(|d| d.name.as_str()), + def.and_then(|d| d.provider.as_deref()), + def.and_then(|d| d.model.as_deref()), + )?; AgentLoop::with_tools_and_system_prompt_provider( AgentRuntimeConfig::from(effective_provider_config), @@ -714,6 +760,8 @@ impl DefaultSubAgentRuntime { summary: extract_summary(&final_message.content), output: final_message.content, task_id: session.id.clone(), + provider: None, + model: None, }) } Ok(Err(e)) => Err(TaskError::ExecutionFailed(e.to_string())), @@ -759,6 +807,8 @@ impl DefaultSubAgentRuntime { summary: extract_summary(&final_message.content), output: final_message.content, task_id: session.id.clone(), + provider: None, + model: None, }) } Ok(Err(e)) => Err(TaskError::ExecutionFailed(e.to_string())), @@ -797,6 +847,8 @@ impl DefaultSubAgentRuntime { summary: error.to_string(), output: String::new(), task_id: session.id.clone(), + provider: None, + model: None, }) } } @@ -926,23 +978,14 @@ impl SubAgentRuntime for DefaultSubAgentRuntime { } else { self.skills.system_index_prompt() }; - // 同步解析 def 中的 provider/model 覆盖,保证环境提示中的模型名与实际使用的模型一致 - let effective_provider_config = match (def.provider.is_some(), def.model.is_some()) { - (true, _) | (_, true) => self - .model_resolver - .resolve( - def.provider.as_deref(), - def.model.as_deref(), - &self.provider_config, - ) - .map_err(|e| { - TaskError::AgentCreationFailed(format!( - "subagent '{}' model resolution failed: {}", - def.name, e - )) - })?, - _ => self.provider_config.clone(), - }; + // 同步解析生效配置(def > topic 级 > session 级 > 全局), + // 与 create_subagent 共用 helper,保证环境提示中的模型名与实际使用的模型一致 + let effective_provider_config = self.resolve_effective_provider_config( + &session, + Some(def.name.as_str()), + def.provider.as_deref(), + def.model.as_deref(), + )?; let system_prompt = SubagentPromptBuilder::build( &def, &task.description, @@ -1177,6 +1220,7 @@ impl SubAgentRuntime for DefaultSubAgentRuntime { // 8c. 立即返回 running 占位结果 // 注意: summary 留空,output 只含引导信息。LLM 看到 running 后应调 wait_for_subagents。 + // provider/model 携带实际生效配置,供前端在偏离主代理模型时差异显示。 return Ok(TaskToolResult { status: "running".to_string(), summary: format!("Task {} spawned asynchronously", task_id), @@ -1185,6 +1229,8 @@ impl SubAgentRuntime for DefaultSubAgentRuntime { task_id ), task_id, + provider: Some(effective_provider_config.name.clone()), + model: Some(effective_provider_config.model_id.clone()), }); } diff --git a/src/tools/task/types.rs b/src/tools/task/types.rs index 74c5643..f3fbcaf 100644 --- a/src/tools/task/types.rs +++ b/src/tools/task/types.rs @@ -301,6 +301,12 @@ pub struct TaskToolResult { pub output: String, /// 会话 ID(用于恢复) pub task_id: String, + /// 子代理实际使用的 provider 名(running 占位时携带,供前端差异显示;其余场景可省略) + #[serde(skip_serializing_if = "Option::is_none")] + pub provider: Option, + /// 子代理实际使用的 model id(同上) + #[serde(skip_serializing_if = "Option::is_none")] + pub model: Option, } /// 异步子代理完成状态