From a825b10c48a4704a7f6c88888daf47731e02d7f4 Mon Sep 17 00:00:00 2001 From: oudecheng <13802883547@139.com> Date: Thu, 30 Jul 2026 23:24:51 +0800 Subject: [PATCH] =?UTF-8?q?feat(model):=20=E4=B8=BB=20agent=20=E6=94=AF?= =?UTF-8?q?=E6=8C=81=E4=BC=9A=E8=AF=9D=E7=BA=A7=E6=A8=A1=E5=9E=8B=E8=A6=86?= =?UTF-8?q?=E7=9B=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新增 ModelSelectionStore 存储 session_id -> (provider, model) 映射,职责单一,仅依赖 std,不引入业务模块耦合。 AgentFactory::create 中按链式覆盖应用模型配置: 专家 frontmatter 覆盖 -> 用户手动选择覆盖(最高优先级)。 HTTP API: POST /api/session/select-model 设置/清除用户模型覆盖(校验 provider/model 存在性); GET /api/session/selected-model 读取当前 session 的用户覆盖(与 experts/selected 对称)。 在 build_session_manager 系列函数中创建并注入 ModelSelectionStore,GatewayState 持有 Arc 供 HTTP handler 访问。 --- src/gateway/agent_factory.rs | 34 +++++++++- src/gateway/http.rs | 117 +++++++++++++++++++++++++++++++++ src/gateway/mod.rs | 10 ++- src/gateway/model_selection.rs | 84 +++++++++++++++++++++++ src/gateway/runtime.rs | 9 ++- src/gateway/session.rs | 3 +- 6 files changed, 251 insertions(+), 6 deletions(-) create mode 100644 src/gateway/model_selection.rs diff --git a/src/gateway/agent_factory.rs b/src/gateway/agent_factory.rs index a993315..931adb1 100644 --- a/src/gateway/agent_factory.rs +++ b/src/gateway/agent_factory.rs @@ -6,6 +6,7 @@ use crate::domain::CapabilityPolicy; use crate::experts::ExpertPromptProvider; use crate::experts::ExpertRuntime; use crate::gateway::agent_prompt_provider::AgentPromptProvider; +use crate::gateway::model_selection::ModelSelectionStore; use crate::gateway::tool_prompt_provider::ToolPromptProvider; use crate::skills::{SkillPromptProvider, SkillRuntime}; use crate::storage::persistent_session_id; @@ -50,6 +51,8 @@ pub(crate) struct AgentFactory { prompt_repository: Arc, /// Provider/Model 解析器:按专家 frontmatter 中的 provider/model 字段覆盖基础配置 model_resolver: Arc, + /// per-session 的用户模型选择(最高优先级,覆盖专家配置) + model_selections: Arc, /// 实例创建时间戳(用于区分新旧 AgentFactory 实例) instance_id: u64, } @@ -76,6 +79,7 @@ impl AgentFactory { reinject_every: usize, prompt_repository: Arc, model_resolver: Arc, + model_selections: Arc, ) -> Self { // 使用 Arc 指针地址作为实例标识符,用于区分新旧 AgentFactory 实例 let instance_id = Arc::as_ptr(&tools) as u64; @@ -92,6 +96,7 @@ impl AgentFactory { reinject_every, prompt_repository, model_resolver, + model_selections, instance_id, } } @@ -105,7 +110,7 @@ impl AgentFactory { // 按专家 frontmatter 中的 provider/model 字段解析覆盖基础 provider_config。 // 引用不存在的 provider/model 名时报错并阻止会话(用户主动选择的角色,配置错误应明确反馈)。 - let effective_provider_config = match &expert { + let expert_provider_config = match &expert { Some(e) if e.provider.is_some() || e.model.is_some() => { let resolved = self.model_resolver.resolve( e.provider.as_deref(), @@ -126,6 +131,33 @@ impl AgentFactory { _ => request.provider_config.clone(), }; + // 按用户手动选择的 provider/model 覆盖(最高优先级,覆盖专家配置)。 + // 引用不存在的 provider/model 名时报错并阻止会话(用户主动选择,配置错误应明确反馈)。 + let effective_provider_config = + match self.model_selections.get(&session_id) { + Some((user_provider, user_model)) + if user_provider.is_some() || user_model.is_some() => + { + let resolved = self + .model_resolver + .resolve( + user_provider.as_deref(), + user_model.as_deref(), + &expert_provider_config, + ) + .map_err(|e| AgentError::Other(e.to_string()))?; + tracing::info!( + instance_id = self.instance_id, + session_id = %session_id, + provider = %resolved.name, + model_id = %resolved.model_id, + "AgentFactory: applied user model override" + ); + resolved + } + _ => expert_provider_config, + }; + // 诊断日志:记录 agent 实际使用的配置和实例 ID tracing::info!( instance_id = self.instance_id, diff --git a/src/gateway/http.rs b/src/gateway/http.rs index 1affeae..868b9fc 100644 --- a/src/gateway/http.rs +++ b/src/gateway/http.rs @@ -268,6 +268,15 @@ pub struct ToolsListResponse { pub struct ModelOptionsResponse { pub providers: Vec, pub models: Vec, + /// 当前默认 agent 的 provider/model 名(来自 config.json agents.default)。 + /// 前端用于在"继承默认"选项旁标注当前生效的模型。 + pub current: CurrentModel, +} + +#[derive(Serialize)] +pub struct CurrentModel { + pub provider: String, + pub model: String, } /// GET /api/tools — Return all registered tools (builtin + MCP) with name/description/source. @@ -308,9 +317,22 @@ pub async fn model_options( ) -> Json { let config = state.config.read().await; let resolver = crate::config::ModelResolver::from_config(&config); + // 当前默认 agent 的 provider/model 名(直接引用 providers/models 表的 key) + let current = if let Some(agent) = config.agents.get("default") { + CurrentModel { + provider: agent.provider.clone(), + model: agent.model.clone(), + } + } else { + CurrentModel { + provider: String::new(), + model: String::new(), + } + }; Json(ModelOptionsResponse { providers: resolver.provider_names(), models: resolver.model_names(), + current, }) } @@ -647,6 +669,23 @@ pub struct ExpertSelectResponse { pub error: Option, } +/// POST /api/session/select-model 请求体 +#[derive(Deserialize)] +pub struct SelectModelRequest { + pub session_id: String, + /// None 或空字符串表示清除覆盖(继承默认) + pub provider: Option, + /// None 或空字符串表示清除覆盖(继承默认) + pub model: Option, +} + +#[derive(Serialize)] +pub struct SelectModelResponse { + pub success: bool, + #[serde(skip_serializing_if = "Option::is_none")] + pub error: Option, +} + #[derive(Serialize)] pub struct ExpertResponse { pub name: String, @@ -929,3 +968,81 @@ pub async fn experts_select( ), } } + +/// POST /api/session/select-model — 设置(或清除)session 的用户模型覆盖 +pub async fn session_select_model( + State(state): State>, + Json(req): Json, +) -> (StatusCode, Json) { + // 规范化:trim 后空字符串视为 None(与 frontmatter 解析逻辑一致) + let provider = req + .provider + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()); + let model = req + .model + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()); + + // 校验:provider/model 名必须在 config 的 providers/models 表中存在 + // (与 AgentFactory::create 中的解析失败行为对齐,提前反馈错误) + let config = state.config.read().await; + if let Some(name) = provider.as_ref() { + if !config.providers.contains_key(name) { + return ( + StatusCode::BAD_REQUEST, + Json(SelectModelResponse { + success: false, + error: Some(format!("provider '{}' not found in config", name)), + }), + ); + } + } + if let Some(name) = model.as_ref() { + if !config.models.contains_key(name) { + return ( + StatusCode::BAD_REQUEST, + Json(SelectModelResponse { + success: false, + error: Some(format!("model '{}' not found in config", name)), + }), + ); + } + } + drop(config); + + state + .model_selections + .set(&req.session_id, provider, model); + ( + StatusCode::OK, + Json(SelectModelResponse { + success: true, + error: None, + }), + ) +} + +/// GET /api/session/selected-model?session_id=... — 返回该 session 当前的用户模型覆盖 +#[derive(Deserialize)] +pub struct SessionSelectedModelQuery { + pub session_id: String, +} + +#[derive(Serialize)] +pub struct SessionSelectedModelResponse { + /// None 表示未设置用户覆盖(继承默认 / 专家配置) + pub provider: Option, + pub model: Option, +} + +pub async fn session_selected_model( + State(state): State>, + Query(q): Query, +) -> Json { + let (provider, model) = state + .model_selections + .get(&q.session_id) + .unwrap_or((None, None)); + Json(SessionSelectedModelResponse { provider, model }) +} diff --git a/src/gateway/mod.rs b/src/gateway/mod.rs index 344098f..8d5a178 100644 --- a/src/gateway/mod.rs +++ b/src/gateway/mod.rs @@ -10,6 +10,7 @@ pub mod http; pub mod memory_maintenance; pub mod memory_maintenance_coordinator; pub mod message_prepare; +pub mod model_selection; pub mod outbound_dispatcher; pub mod processor; pub mod prompt; @@ -67,6 +68,8 @@ pub struct GatewayState { pub skills: Arc, pub experts: Arc, pub subagent_runtime: Arc, + /// per-session 的用户模型选择(覆盖专家配置) + pub model_selections: Arc, } impl GatewayState { @@ -92,7 +95,7 @@ impl GatewayState { mcp_servers: config.mcp_servers.clone(), }; - let (session_manager, task_repository, mcp_manager, subagent_runtime) = build_session_manager_with_sender( + let (session_manager, task_repository, mcp_manager, subagent_runtime, model_selections) = build_session_manager_with_sender( agent_prompt_reinject_every, show_tool_results, config.time.timezone.clone(), @@ -131,6 +134,7 @@ impl GatewayState { skills, experts, subagent_runtime, + model_selections, }) } @@ -254,6 +258,8 @@ pub async fn run( .route("/api/experts/delete", routing::delete(http::experts_delete)) .route("/api/experts/selected", routing::get(http::experts_selected)) .route("/api/experts/select", routing::post(http::experts_select)) + .route("/api/session/select-model", routing::post(http::session_select_model)) + .route("/api/session/selected-model", routing::get(http::session_selected_model)) .route("/ws", routing::get(ws::ws_handler)) .fallback(static_handler) .with_state(state.clone()) @@ -278,6 +284,8 @@ pub async fn run( .route("/api/experts/delete", routing::delete(http::experts_delete)) .route("/api/experts/selected", routing::get(http::experts_selected)) .route("/api/experts/select", routing::post(http::experts_select)) + .route("/api/session/select-model", routing::post(http::session_select_model)) + .route("/api/session/selected-model", routing::get(http::session_selected_model)) .route("/ws", routing::get(ws::ws_handler)) .fallback_service(ServeDir::new(&static_dir)) .with_state(state.clone()) diff --git a/src/gateway/model_selection.rs b/src/gateway/model_selection.rs new file mode 100644 index 0000000..d9459ca --- /dev/null +++ b/src/gateway/model_selection.rs @@ -0,0 +1,84 @@ +use std::collections::HashMap; +use std::sync::RwLock; + +/// per-session 的用户模型覆盖选择存储。 +/// +/// 与 ExpertRuntime 的 session_experts 平级独立,职责单一: +/// 只负责存储 session_id -> (provider, model) 的映射,不依赖任何业务模块。 +#[derive(Debug, Default)] +pub struct ModelSelectionStore { + selections: RwLock, Option)>>, +} + +impl ModelSelectionStore { + pub fn new() -> Self { + Self::default() + } + + /// 设置 session 的用户模型覆盖。provider 和 model 均为 None 时清除该 session 的选择。 + pub fn set( + &self, + session_id: &str, + provider: Option, + model: Option, + ) { + let mut selections = self + .selections + .write() + .expect("model selections rwlock poisoned"); + if provider.is_none() && model.is_none() { + selections.remove(session_id); + } else { + selections.insert(session_id.to_string(), (provider, model)); + } + } + + /// 读取 session 的用户模型覆盖。 + pub fn get(&self, session_id: &str) -> Option<(Option, Option)> { + self.selections + .read() + .expect("model selections rwlock poisoned") + .get(session_id) + .cloned() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn set_and_get() { + let store = ModelSelectionStore::new(); + store.set("s1", Some("p1".to_string()), Some("m1".to_string())); + assert_eq!( + store.get("s1"), + Some((Some("p1".to_string()), Some("m1".to_string()))) + ); + } + + #[test] + fn get_missing_returns_none() { + let store = ModelSelectionStore::new(); + assert_eq!(store.get("missing"), None); + } + + #[test] + fn set_none_none_removes_entry() { + let store = ModelSelectionStore::new(); + store.set("s1", Some("p1".to_string()), Some("m1".to_string())); + assert!(store.get("s1").is_some()); + store.set("s1", None, None); + assert!(store.get("s1").is_none()); + } + + #[test] + fn set_only_provider_keeps_entry() { + let store = ModelSelectionStore::new(); + store.set("s1", Some("p1".to_string()), None); + assert_eq!( + store.get("s1"), + Some((Some("p1".to_string()), None)) + ); + } +} diff --git a/src/gateway/runtime.rs b/src/gateway/runtime.rs index fcfebce..d5eaef9 100644 --- a/src/gateway/runtime.rs +++ b/src/gateway/runtime.rs @@ -9,6 +9,7 @@ use tokio::sync::RwLock; use crate::agent::AgentError; use crate::bus::MessageBus; use crate::config::{LLMProviderConfig, MemoryMaintenanceConfig, ModelResolver, SubagentsConfig, TaskConfig}; +use crate::gateway::model_selection::ModelSelectionStore; use crate::gateway::tool_registry_factory::ToolRegistryFactory; use crate::mcp::McpInitializer; use crate::mcp::client::McpClientManager; @@ -54,7 +55,7 @@ pub(crate) fn build_session_manager( mcp_config: crate::mcp::McpConfig, bus: Option>, model_resolver: Arc, -) -> Result<(SessionManager, Arc, Option>, Arc), AgentError> { +) -> Result<(SessionManager, Arc, Option>, Arc, Arc), AgentError> { build_session_manager_with_sender( agent_prompt_reinject_every, show_tool_results, @@ -93,7 +94,7 @@ pub(crate) fn build_session_manager_with_sender( mcp_config: crate::mcp::McpConfig, bus: Option>, model_resolver: Arc, -) -> Result<(SessionManager, Arc, Option>, Arc), AgentError> { +) -> Result<(SessionManager, Arc, Option>, Arc, Arc), AgentError> { let store = Arc::new( SessionStore::new() .map_err(|err| AgentError::Other(format!("session store init error: {}", err)))?, @@ -272,6 +273,7 @@ pub(crate) fn build_session_manager_with_sender( ); let prompt_repository: Arc = store.clone(); + let model_selections = Arc::new(ModelSelectionStore::new()); let agent_factory = AgentFactory::new( tools.clone(), skills.clone(), @@ -280,6 +282,7 @@ pub(crate) fn build_session_manager_with_sender( agent_prompt_reinject_every as usize, prompt_repository.clone(), model_resolver.clone(), + model_selections.clone(), ); let session_factory = SessionFactory::new( provider_config.clone(), @@ -316,5 +319,5 @@ pub(crate) fn build_session_manager_with_sender( scheduled_tasks, memory_maintenance, task_repository: task_repository.clone(), - }), task_repository, mcp_manager, subagent_runtime)) + }), task_repository, mcp_manager, subagent_runtime, model_selections)) } diff --git a/src/gateway/session.rs b/src/gateway/session.rs index 1f9796a..59379eb 100644 --- a/src/gateway/session.rs +++ b/src/gateway/session.rs @@ -271,6 +271,7 @@ impl Session { agent_prompt_reinject_every as usize, prompt_repository.clone(), model_resolver, + Arc::new(super::model_selection::ModelSelectionStore::new()), ); Self::with_factories( channel_name, @@ -716,7 +717,7 @@ impl SessionManager { None, model_resolver, ) - .map(|(session_manager, _, _, _)| session_manager) + .map(|(session_manager, _, _, _, _)| session_manager) } pub fn tools(&self) -> Arc {