diff --git a/src/gateway/http.rs b/src/gateway/http.rs index 6f1b380..8525e6e 100644 --- a/src/gateway/http.rs +++ b/src/gateway/http.rs @@ -1217,3 +1217,154 @@ pub async fn session_selected_model( .unwrap_or((None, None)); Json(SessionSelectedModelResponse { provider, model }) } + +#[derive(Deserialize)] +pub struct SelectTopicModelRequest { + pub session_id: String, + pub topic_id: String, + pub provider: Option, + pub model: Option, +} + +/// POST /api/topic/select-model — 设置(或清除)话题级用户模型覆盖。 +/// +/// 双写语义(session store 的 key 一律取 topic 行自带的 session_id,不信任请求体): +/// - 设置:写 topics 行(话题记忆,物化)+ 内存缓存 + session store(成为新话题默认) +/// - 清除(provider/model 均空):移除话题级选择,并同步清除 session 级选择 +/// (否则 agent 会在下一条消息把 session 级重新物化回 topic 行,重置永远无效) +pub async fn topic_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()); + + // 话题必须存在(防止对已删除话题静默写空)。 + // 同时取 topic 行自带的 session_id 作为双写目标——不信任请求体中的 session_id, + // 防止客户端误传导致污染其他 session 的默认模型。 + let store = state.session_manager.store(); + let topic_session_id = match store.get_topic(&req.topic_id) { + Ok(Some(topic)) => topic.session_id, + Ok(None) => { + return ( + StatusCode::NOT_FOUND, + Json(SelectModelResponse { + success: false, + error: Some(format!("topic '{}' not found", req.topic_id)), + }), + ); + } + Err(e) => { + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(SelectModelResponse { + success: false, + error: Some(format!("failed to load topic: {}", e)), + }), + ); + } + }; + + // 校验:provider/model 名必须在 config 的 providers/models 表中存在 + 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); + + let is_clear = provider.is_none() && model.is_none(); + if let Err(e) = store.update_topic_model(&req.topic_id, provider.as_deref(), model.as_deref()) { + return ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(SelectModelResponse { + success: false, + error: Some(format!("failed to persist topic model: {}", e)), + }), + ); + } + + state + .topic_model_selections + .set(&req.topic_id, provider.clone(), model.clone()); + + // 双写 session store(key 一律取 topic 行的 session_id): + // - 显式设置:成为新话题的初始默认("最近使用") + // - 清除:同步移除 session 级选择。否则"重置为默认"会陷入死循环—— + // topic 行清空后 agent 命中 session 级又重新物化回 topic 行,重置形同虚设 + if is_clear { + state.model_selections.set(&topic_session_id, None, None); + } else { + state.model_selections.set(&topic_session_id, provider, model); + } + + ( + StatusCode::OK, + Json(SelectModelResponse { + success: true, + error: None, + }), + ) +} + +/// GET /api/topic/selected-model?topic_id=... — 返回话题生效的用户模型选择。 +/// +/// 语义与 session 端点一致:只反映用户选择(topic 级优先,miss 回退 session 级), +/// 不解析 expert/config 默认。 +#[derive(Deserialize)] +pub struct TopicSelectedModelQuery { + pub topic_id: String, +} + +pub async fn topic_selected_model( + State(state): State>, + Query(q): Query, +) -> Json { + // topic 级命中直接返回;miss 时按 topic 行的 session_id 回退 session 级 + let (provider, model) = match state.topic_model_selections.get(&q.topic_id) { + Some(selection) => selection, + None => match state.session_manager.store().get_topic(&q.topic_id) { + Ok(Some(topic)) => { + // 双保险:SQLite 有物化值但缓存 miss(理论上不会发生)时回填缓存 + if topic.provider.is_some() || topic.model.is_some() { + state.topic_model_selections.set( + &q.topic_id, + topic.provider.clone(), + topic.model.clone(), + ); + } + state + .model_selections + .get(&topic.session_id) + .unwrap_or((None, None)) + } + _ => (None, None), + }, + }; + Json(SessionSelectedModelResponse { provider, model }) +} diff --git a/src/gateway/mod.rs b/src/gateway/mod.rs index 3d912c8..5773ddf 100644 --- a/src/gateway/mod.rs +++ b/src/gateway/mod.rs @@ -75,6 +75,8 @@ pub struct GatewayState { pub subagent_executor: Option>, /// per-session 的用户模型选择(覆盖专家配置) pub model_selections: Arc, + /// per-topic 的用户模型选择(最高优先级;物化保证话题模型不随 session 选择漂移) + pub topic_model_selections: Arc, /// Prometheus metrics handle(/metrics 端点渲染用)。 /// None 表示 recorder 安装失败;热重启时从 OnceLock 缓存复用,不会因重复安装而变为 None。 pub prometheus_handle: Option, @@ -108,7 +110,7 @@ impl GatewayState { mcp_servers: config.mcp_servers.clone(), }; - let (session_manager, task_repository, mcp_manager, subagent_runtime, model_selections, subagent_executor) = + let (session_manager, task_repository, mcp_manager, subagent_runtime, model_selections, topic_model_selections, subagent_executor) = build_session_manager_with_sender( agent_prompt_reinject_every, show_tool_results, @@ -155,6 +157,7 @@ impl GatewayState { subagent_runtime, subagent_executor, model_selections, + topic_model_selections, prometheus_handle, }) } @@ -369,6 +372,14 @@ pub async fn run( "/api/session/selected-model", routing::get(http::session_selected_model), ) + .route( + "/api/topic/select-model", + routing::post(http::topic_select_model), + ) + .route( + "/api/topic/selected-model", + routing::get(http::topic_selected_model), + ) .route("/ws", routing::get(ws::ws_handler)) .route("/metrics", routing::get(http::metrics_handler)); diff --git a/src/gateway/runtime.rs b/src/gateway/runtime.rs index 0b1ee18..cf5203b 100644 --- a/src/gateway/runtime.rs +++ b/src/gateway/runtime.rs @@ -67,6 +67,7 @@ pub(crate) fn build_session_manager( Option>, Arc, Arc, + Arc, Option>, ), AgentError, @@ -120,6 +121,7 @@ pub(crate) fn build_session_manager_with_sender( Option>, Arc, Arc, + Arc, Option>, ), AgentError, @@ -128,6 +130,22 @@ pub(crate) fn build_session_manager_with_sender( SessionStore::new() .map_err(|err| AgentError::Other(format!("session store init error: {}", err)))?, ); + // 模型选择内存缓存:session 级 + topic 级(topic 级启动时从 topics 表预热) + let model_selections = Arc::new(ModelSelectionStore::new()); + let topic_model_selections = Arc::new(ModelSelectionStore::new()); + match store.list_topic_model_selections() { + Ok(entries) => { + for (topic_id, provider, model) in entries { + topic_model_selections.set(&topic_id, provider, model); + } + } + Err(err) => { + tracing::warn!( + error = %err, + "build_session_manager: failed to preload topic model selections" + ); + } + } let known_agents = provider_configs.keys().cloned().collect::>(); let provider_configs = ProviderConfigService::new( provider_config.clone(), @@ -254,6 +272,8 @@ pub(crate) fn build_session_manager_with_sender( bus.clone(), store.clone(), skills.clone(), + Some(model_selections.clone()), + Some(topic_model_selections.clone()), )); // 注册 task 工具到子代理工具集(需在 runtime 创建之后,打破循环依赖) @@ -320,7 +340,6 @@ pub(crate) fn build_session_manager_with_sender( ); let prompt_repository: Arc = store.clone(); - let model_selections = Arc::new(ModelSelectionStore::new()); let observer: Arc = crate::observability::metrics::default_observer(); let agent_factory = AgentFactory::new( @@ -332,6 +351,8 @@ pub(crate) fn build_session_manager_with_sender( prompt_repository.clone(), model_resolver.clone(), model_selections.clone(), + topic_model_selections.clone(), + store.clone(), compaction_config, Some(observer), ); @@ -376,6 +397,7 @@ pub(crate) fn build_session_manager_with_sender( mcp_manager, subagent_runtime, model_selections, + topic_model_selections, subagent_executor, )) } diff --git a/src/gateway/session.rs b/src/gateway/session.rs index 4e3d859..8392cb5 100644 --- a/src/gateway/session.rs +++ b/src/gateway/session.rs @@ -309,6 +309,8 @@ impl Session { prompt_repository.clone(), model_resolver, Arc::new(super::model_selection::ModelSelectionStore::new()), + Arc::new(super::model_selection::ModelSelectionStore::new()), + store.clone(), crate::config::CompactionConfig::default(), None, ); @@ -994,7 +996,7 @@ impl SessionManager { model_resolver, crate::config::CompactionConfig::default(), ) - .map(|(session_manager, _, _, _, _, _)| session_manager) + .map(|(session_manager, _, _, _, _, _, _)| session_manager) } pub fn tools(&self) -> Arc { diff --git a/src/storage/migrations.rs b/src/storage/migrations.rs index 544baae..52af32c 100644 --- a/src/storage/migrations.rs +++ b/src/storage/migrations.rs @@ -93,6 +93,18 @@ pub(super) fn ensure_messages_schema(conn: &Connection) -> Result<(), StorageErr Ok(()) } +/// topics 表:话题级模型选择列(用户在特定话题内显式选择的 provider/model)。 +/// NULL 表示该话题无显式选择(运行时按 session 级 → expert → config 链解析)。 +pub(super) fn ensure_topics_schema(conn: &Connection) -> Result<(), StorageError> { + if !has_column(conn, "topics", "provider")? { + add_column_if_missing(conn, "ALTER TABLE topics ADD COLUMN provider TEXT")?; + } + if !has_column(conn, "topics", "model")? { + add_column_if_missing(conn, "ALTER TABLE topics ADD COLUMN model TEXT")?; + } + Ok(()) +} + pub(super) fn ensure_scheduler_schema(conn: &Connection) -> Result<(), StorageError> { if !has_column(conn, "scheduler_jobs", "schedule_json")? { conn.execute( diff --git a/src/storage/mod.rs b/src/storage/mod.rs index ecba5e7..1d1c1be 100644 --- a/src/storage/mod.rs +++ b/src/storage/mod.rs @@ -231,6 +231,7 @@ impl SessionStore { ensure_sessions_schema(&conn)?; ensure_messages_schema(&conn)?; + ensure_topics_schema(&conn)?; ensure_scheduler_schema(&conn)?; ensure_memory_scope_key_migration(&conn)?; ensure_todos_schema(&conn)?; @@ -468,7 +469,7 @@ impl SessionStore { pub fn get_topic(&self, topic_id: &str) -> Result, StorageError> { let conn = self.pool.get()?; let mut stmt = conn.prepare( - "SELECT id, session_id, title, description, created_at, updated_at, last_active_at, message_count FROM topics WHERE id = ?1", + "SELECT id, session_id, title, description, created_at, updated_at, last_active_at, message_count, provider, model FROM topics WHERE id = ?1", )?; stmt.query_row(params![topic_id], |row| { @@ -481,6 +482,8 @@ impl SessionStore { updated_at: row.get(5)?, last_active_at: row.get(6)?, message_count: row.get(7)?, + provider: row.get(8)?, + model: row.get(9)?, }) }) .optional() @@ -490,7 +493,7 @@ impl SessionStore { pub fn list_topics(&self, session_id: &str) -> Result, StorageError> { let conn = self.pool.get()?; let mut stmt = conn.prepare( - "SELECT id, session_id, title, description, created_at, updated_at, last_active_at, message_count FROM topics WHERE session_id = ?1 ORDER BY last_active_at DESC" + "SELECT id, session_id, title, description, created_at, updated_at, last_active_at, message_count, provider, model FROM topics WHERE session_id = ?1 ORDER BY last_active_at DESC" )?; let rows = stmt.query_map(params![session_id], |row| { @@ -503,6 +506,8 @@ impl SessionStore { updated_at: row.get(5)?, last_active_at: row.get(6)?, message_count: row.get(7)?, + provider: row.get(8)?, + model: row.get(9)?, }) })?; @@ -544,6 +549,40 @@ impl SessionStore { Ok(()) } + /// 设置/清除话题级模型选择。provider 与 model 均为 None 时清除(恢复继承)。 + pub fn update_topic_model( + &self, + topic_id: &str, + provider: Option<&str>, + model: Option<&str>, + ) -> Result<(), StorageError> { + let now = current_timestamp(); + let conn = self.pool.get()?; + conn.execute( + "UPDATE topics SET provider = ?2, model = ?3, updated_at = ?4 WHERE id = ?1", + params![topic_id, provider, model, now], + )?; + Ok(()) + } + + /// 全量读出话题级模型选择,供启动时预热内存缓存。 + pub fn list_topic_model_selections( + &self, + ) -> Result, Option)>, StorageError> { + let conn = self.pool.get()?; + let mut stmt = conn.prepare( + "SELECT id, provider, model FROM topics WHERE provider IS NOT NULL OR model IS NOT NULL", + )?; + let rows = stmt.query_map([], |row| { + Ok((row.get(0)?, row.get(1)?, row.get(2)?)) + })?; + let mut result = Vec::new(); + for row in rows { + result.push(row?); + } + Ok(result) + } + pub fn touch_topic(&self, topic_id: &str) -> Result<(), StorageError> { let now = current_timestamp(); let conn = self.pool.get()?; diff --git a/src/storage/records.rs b/src/storage/records.rs index 5085f9e..6c52c76 100644 --- a/src/storage/records.rs +++ b/src/storage/records.rs @@ -109,6 +109,9 @@ pub struct TopicRecord { pub updated_at: i64, pub last_active_at: i64, pub message_count: i64, + /// 话题级用户模型选择(NULL 表示无显式选择,运行时按继承链解析) + pub provider: Option, + pub model: Option, } /// pending_subagents 表的记录,跟踪异步子代理执行状态。