feat(model): 话题级模型选择持久化与 API(topics 表加 provider/model 列 + 双写内存缓存)

- topics 表迁移新增 provider/model 列,记录话题级显式模型选择
- storage 层新增 update_topic_model / list_topic_model_selections
- 新增 POST /api/topic/select-model 与 GET /api/topic/selected-model
- 双写 key 取 topic 行自带 session_id(不信任请求体,防污染其他会话)
- 清除时同步清 session 级选择,避免'重置为默认'被物化逻辑重新覆盖
This commit is contained in:
oudecheng 2026-08-15 15:33:58 +08:00
parent 9e503f2672
commit 414105d419
7 changed files with 245 additions and 5 deletions

View File

@ -1217,3 +1217,154 @@ pub async fn session_selected_model(
.unwrap_or((None, None)); .unwrap_or((None, None));
Json(SessionSelectedModelResponse { provider, model }) Json(SessionSelectedModelResponse { provider, model })
} }
#[derive(Deserialize)]
pub struct SelectTopicModelRequest {
pub session_id: String,
pub topic_id: String,
pub provider: Option<String>,
pub model: Option<String>,
}
/// 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<Arc<GatewayState>>,
Json(req): Json<SelectTopicModelRequest>,
) -> (StatusCode, Json<SelectModelResponse>) {
// 规范化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 storekey 一律取 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<Arc<GatewayState>>,
Query(q): Query<TopicSelectedModelQuery>,
) -> Json<SessionSelectedModelResponse> {
// 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 })
}

View File

@ -75,6 +75,8 @@ pub struct GatewayState {
pub subagent_executor: Option<Arc<dyn crate::tools::SubAgentRuntime>>, pub subagent_executor: Option<Arc<dyn crate::tools::SubAgentRuntime>>,
/// per-session 的用户模型选择(覆盖专家配置) /// per-session 的用户模型选择(覆盖专家配置)
pub model_selections: Arc<model_selection::ModelSelectionStore>, pub model_selections: Arc<model_selection::ModelSelectionStore>,
/// per-topic 的用户模型选择(最高优先级;物化保证话题模型不随 session 选择漂移)
pub topic_model_selections: Arc<model_selection::ModelSelectionStore>,
/// Prometheus metrics handle/metrics 端点渲染用)。 /// Prometheus metrics handle/metrics 端点渲染用)。
/// None 表示 recorder 安装失败;热重启时从 OnceLock 缓存复用,不会因重复安装而变为 None。 /// None 表示 recorder 安装失败;热重启时从 OnceLock 缓存复用,不会因重复安装而变为 None。
pub prometheus_handle: Option<metrics_exporter_prometheus::PrometheusHandle>, pub prometheus_handle: Option<metrics_exporter_prometheus::PrometheusHandle>,
@ -108,7 +110,7 @@ impl GatewayState {
mcp_servers: config.mcp_servers.clone(), 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( build_session_manager_with_sender(
agent_prompt_reinject_every, agent_prompt_reinject_every,
show_tool_results, show_tool_results,
@ -155,6 +157,7 @@ impl GatewayState {
subagent_runtime, subagent_runtime,
subagent_executor, subagent_executor,
model_selections, model_selections,
topic_model_selections,
prometheus_handle, prometheus_handle,
}) })
} }
@ -369,6 +372,14 @@ pub async fn run(
"/api/session/selected-model", "/api/session/selected-model",
routing::get(http::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("/ws", routing::get(ws::ws_handler))
.route("/metrics", routing::get(http::metrics_handler)); .route("/metrics", routing::get(http::metrics_handler));

View File

@ -67,6 +67,7 @@ pub(crate) fn build_session_manager(
Option<Arc<McpClientManager>>, Option<Arc<McpClientManager>>,
Arc<SubagentRuntime>, Arc<SubagentRuntime>,
Arc<ModelSelectionStore>, Arc<ModelSelectionStore>,
Arc<ModelSelectionStore>,
Option<Arc<dyn SubAgentRuntime>>, Option<Arc<dyn SubAgentRuntime>>,
), ),
AgentError, AgentError,
@ -120,6 +121,7 @@ pub(crate) fn build_session_manager_with_sender(
Option<Arc<McpClientManager>>, Option<Arc<McpClientManager>>,
Arc<SubagentRuntime>, Arc<SubagentRuntime>,
Arc<ModelSelectionStore>, Arc<ModelSelectionStore>,
Arc<ModelSelectionStore>,
Option<Arc<dyn SubAgentRuntime>>, Option<Arc<dyn SubAgentRuntime>>,
), ),
AgentError, AgentError,
@ -128,6 +130,22 @@ pub(crate) fn build_session_manager_with_sender(
SessionStore::new() SessionStore::new()
.map_err(|err| AgentError::Other(format!("session store init error: {}", err)))?, .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::<HashSet<_>>(); let known_agents = provider_configs.keys().cloned().collect::<HashSet<_>>();
let provider_configs = ProviderConfigService::new( let provider_configs = ProviderConfigService::new(
provider_config.clone(), provider_config.clone(),
@ -254,6 +272,8 @@ pub(crate) fn build_session_manager_with_sender(
bus.clone(), bus.clone(),
store.clone(), store.clone(),
skills.clone(), skills.clone(),
Some(model_selections.clone()),
Some(topic_model_selections.clone()),
)); ));
// 注册 task 工具到子代理工具集(需在 runtime 创建之后,打破循环依赖) // 注册 task 工具到子代理工具集(需在 runtime 创建之后,打破循环依赖)
@ -320,7 +340,6 @@ pub(crate) fn build_session_manager_with_sender(
); );
let prompt_repository: Arc<dyn PromptInjectionRepository> = store.clone(); let prompt_repository: Arc<dyn PromptInjectionRepository> = store.clone();
let model_selections = Arc::new(ModelSelectionStore::new());
let observer: Arc<dyn crate::observability::Observer> = let observer: Arc<dyn crate::observability::Observer> =
crate::observability::metrics::default_observer(); crate::observability::metrics::default_observer();
let agent_factory = AgentFactory::new( let agent_factory = AgentFactory::new(
@ -332,6 +351,8 @@ pub(crate) fn build_session_manager_with_sender(
prompt_repository.clone(), prompt_repository.clone(),
model_resolver.clone(), model_resolver.clone(),
model_selections.clone(), model_selections.clone(),
topic_model_selections.clone(),
store.clone(),
compaction_config, compaction_config,
Some(observer), Some(observer),
); );
@ -376,6 +397,7 @@ pub(crate) fn build_session_manager_with_sender(
mcp_manager, mcp_manager,
subagent_runtime, subagent_runtime,
model_selections, model_selections,
topic_model_selections,
subagent_executor, subagent_executor,
)) ))
} }

View File

@ -309,6 +309,8 @@ impl Session {
prompt_repository.clone(), prompt_repository.clone(),
model_resolver, model_resolver,
Arc::new(super::model_selection::ModelSelectionStore::new()), Arc::new(super::model_selection::ModelSelectionStore::new()),
Arc::new(super::model_selection::ModelSelectionStore::new()),
store.clone(),
crate::config::CompactionConfig::default(), crate::config::CompactionConfig::default(),
None, None,
); );
@ -994,7 +996,7 @@ impl SessionManager {
model_resolver, model_resolver,
crate::config::CompactionConfig::default(), crate::config::CompactionConfig::default(),
) )
.map(|(session_manager, _, _, _, _, _)| session_manager) .map(|(session_manager, _, _, _, _, _, _)| session_manager)
} }
pub fn tools(&self) -> Arc<ToolRegistry> { pub fn tools(&self) -> Arc<ToolRegistry> {

View File

@ -93,6 +93,18 @@ pub(super) fn ensure_messages_schema(conn: &Connection) -> Result<(), StorageErr
Ok(()) 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> { pub(super) fn ensure_scheduler_schema(conn: &Connection) -> Result<(), StorageError> {
if !has_column(conn, "scheduler_jobs", "schedule_json")? { if !has_column(conn, "scheduler_jobs", "schedule_json")? {
conn.execute( conn.execute(

View File

@ -231,6 +231,7 @@ impl SessionStore {
ensure_sessions_schema(&conn)?; ensure_sessions_schema(&conn)?;
ensure_messages_schema(&conn)?; ensure_messages_schema(&conn)?;
ensure_topics_schema(&conn)?;
ensure_scheduler_schema(&conn)?; ensure_scheduler_schema(&conn)?;
ensure_memory_scope_key_migration(&conn)?; ensure_memory_scope_key_migration(&conn)?;
ensure_todos_schema(&conn)?; ensure_todos_schema(&conn)?;
@ -468,7 +469,7 @@ impl SessionStore {
pub fn get_topic(&self, topic_id: &str) -> Result<Option<TopicRecord>, StorageError> { pub fn get_topic(&self, topic_id: &str) -> Result<Option<TopicRecord>, StorageError> {
let conn = self.pool.get()?; let conn = self.pool.get()?;
let mut stmt = conn.prepare( 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| { stmt.query_row(params![topic_id], |row| {
@ -481,6 +482,8 @@ impl SessionStore {
updated_at: row.get(5)?, updated_at: row.get(5)?,
last_active_at: row.get(6)?, last_active_at: row.get(6)?,
message_count: row.get(7)?, message_count: row.get(7)?,
provider: row.get(8)?,
model: row.get(9)?,
}) })
}) })
.optional() .optional()
@ -490,7 +493,7 @@ impl SessionStore {
pub fn list_topics(&self, session_id: &str) -> Result<Vec<TopicRecord>, StorageError> { pub fn list_topics(&self, session_id: &str) -> Result<Vec<TopicRecord>, StorageError> {
let conn = self.pool.get()?; let conn = self.pool.get()?;
let mut stmt = conn.prepare( 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| { let rows = stmt.query_map(params![session_id], |row| {
@ -503,6 +506,8 @@ impl SessionStore {
updated_at: row.get(5)?, updated_at: row.get(5)?,
last_active_at: row.get(6)?, last_active_at: row.get(6)?,
message_count: row.get(7)?, message_count: row.get(7)?,
provider: row.get(8)?,
model: row.get(9)?,
}) })
})?; })?;
@ -544,6 +549,40 @@ impl SessionStore {
Ok(()) 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<Vec<(String, Option<String>, Option<String>)>, 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> { pub fn touch_topic(&self, topic_id: &str) -> Result<(), StorageError> {
let now = current_timestamp(); let now = current_timestamp();
let conn = self.pool.get()?; let conn = self.pool.get()?;

View File

@ -109,6 +109,9 @@ pub struct TopicRecord {
pub updated_at: i64, pub updated_at: i64,
pub last_active_at: i64, pub last_active_at: i64,
pub message_count: i64, pub message_count: i64,
/// 话题级用户模型选择NULL 表示无显式选择,运行时按继承链解析)
pub provider: Option<String>,
pub model: Option<String>,
} }
/// pending_subagents 表的记录,跟踪异步子代理执行状态。 /// pending_subagents 表的记录,跟踪异步子代理执行状态。