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:
parent
9e503f2672
commit
414105d419
@ -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<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 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<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 })
|
||||
}
|
||||
|
||||
@ -75,6 +75,8 @@ pub struct GatewayState {
|
||||
pub subagent_executor: Option<Arc<dyn crate::tools::SubAgentRuntime>>,
|
||||
/// per-session 的用户模型选择(覆盖专家配置)
|
||||
pub model_selections: Arc<model_selection::ModelSelectionStore>,
|
||||
/// per-topic 的用户模型选择(最高优先级;物化保证话题模型不随 session 选择漂移)
|
||||
pub topic_model_selections: Arc<model_selection::ModelSelectionStore>,
|
||||
/// Prometheus metrics handle(/metrics 端点渲染用)。
|
||||
/// None 表示 recorder 安装失败;热重启时从 OnceLock 缓存复用,不会因重复安装而变为 None。
|
||||
pub prometheus_handle: Option<metrics_exporter_prometheus::PrometheusHandle>,
|
||||
@ -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));
|
||||
|
||||
|
||||
@ -67,6 +67,7 @@ pub(crate) fn build_session_manager(
|
||||
Option<Arc<McpClientManager>>,
|
||||
Arc<SubagentRuntime>,
|
||||
Arc<ModelSelectionStore>,
|
||||
Arc<ModelSelectionStore>,
|
||||
Option<Arc<dyn SubAgentRuntime>>,
|
||||
),
|
||||
AgentError,
|
||||
@ -120,6 +121,7 @@ pub(crate) fn build_session_manager_with_sender(
|
||||
Option<Arc<McpClientManager>>,
|
||||
Arc<SubagentRuntime>,
|
||||
Arc<ModelSelectionStore>,
|
||||
Arc<ModelSelectionStore>,
|
||||
Option<Arc<dyn SubAgentRuntime>>,
|
||||
),
|
||||
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::<HashSet<_>>();
|
||||
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<dyn PromptInjectionRepository> = store.clone();
|
||||
let model_selections = Arc::new(ModelSelectionStore::new());
|
||||
let observer: Arc<dyn crate::observability::Observer> =
|
||||
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,
|
||||
))
|
||||
}
|
||||
|
||||
@ -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<ToolRegistry> {
|
||||
|
||||
@ -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(
|
||||
|
||||
@ -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<Option<TopicRecord>, 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<Vec<TopicRecord>, 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<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> {
|
||||
let now = current_timestamp();
|
||||
let conn = self.pool.get()?;
|
||||
|
||||
@ -109,6 +109,9 @@ pub struct TopicRecord {
|
||||
pub updated_at: i64,
|
||||
pub last_active_at: i64,
|
||||
pub message_count: i64,
|
||||
/// 话题级用户模型选择(NULL 表示无显式选择,运行时按继承链解析)
|
||||
pub provider: Option<String>,
|
||||
pub model: Option<String>,
|
||||
}
|
||||
|
||||
/// pending_subagents 表的记录,跟踪异步子代理执行状态。
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user