From 1cde659bb675340d5ad591c5a30a4362070ba4d6 Mon Sep 17 00:00:00 2001 From: oudecheng <13802883547@139.com> Date: Tue, 18 Aug 2026 23:31:11 +0800 Subject: [PATCH] =?UTF-8?q?fix(storage):=20=E5=85=A8=E9=A1=B9=E7=9B=AE?= =?UTF-8?q?=E6=89=AB=E6=8F=8F=E4=BF=AE=E5=A4=8D=E2=80=94=E2=80=94=E5=A4=96?= =?UTF-8?q?=E9=94=AE=E7=BA=A6=E6=9D=9F=E5=A4=B1=E6=95=88=E6=A0=B9=E5=9B=A0?= =?UTF-8?q?=E4=B8=8E=E8=AE=A1=E6=95=B0=E4=B8=80=E8=87=B4=E6=80=A7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - [根因] 连接池补 PRAGMA foreign_keys=ON:原先仅引导连接开启,运行时池连接 FK 强制从未生效,delete_topic 的 SET NULL / delete_session 的 CASCADE 形同虚设(悬空 topic_id、topics 孤儿残留) - clear_messages 归零 topics.message_count(原先只清 token 列,/clear 后话题列表显示虚假消息数) - delete_messages_by_ids 补重算 sessions.user_turn_count(sanitize 删除 user 消息后计数永久偏高) - compact/replace_active_history 新增 recompute_session_topic_message_counts(重插消息无 topic_id,topics.message_count 漂移) - switch_topic 三连查与 http 模型选择路由包 spawn_blocking(侧边栏主交互高频路径) - 新增 3 个测试:FK 生效、clear 归零、user_turn_count 重算 --- src/command/handlers/switch_topic.rs | 68 ++++++++++------- src/gateway/http.rs | 70 ++++++++++++------ src/storage/mod.rs | 39 +++++++++- src/storage/tests.rs | 106 +++++++++++++++++++++++++++ 4 files changed, 235 insertions(+), 48 deletions(-) diff --git a/src/command/handlers/switch_topic.rs b/src/command/handlers/switch_topic.rs index 37c6603..bf49410 100644 --- a/src/command/handlers/switch_topic.rs +++ b/src/command/handlers/switch_topic.rs @@ -68,39 +68,52 @@ async fn handle_switch_topic( .ok_or_else(|| CommandError::new("NO_CHAT_ID", "No chat_id in context"))?; // 尝试解析为序号 + // 同步 rusqlite 查询统一移入 blocking 线程池(切换话题是侧边栏主交互) let target_topic_id = if let Ok(index) = topic_id.parse::() { - let topics = handler - .store - .list_topics(session_id) - .map_err(|e| CommandError::new("LIST_TOPICS_ERROR", e.to_string()))?; + let store = handler.store.clone(); + let session_id_bg = session_id.to_string(); + tokio::task::spawn_blocking(move || -> Result { + let topics = store + .list_topics(&session_id_bg) + .map_err(|e| CommandError::new("LIST_TOPICS_ERROR", e.to_string()))?; - let index = index.saturating_sub(1); - if index >= topics.len() { - return Err(CommandError::new( - "INVALID_TOPIC_INDEX", - format!( - "Topic index {} is out of range (1-{})", - index + 1, - topics.len() - ), - )); - } - topics[index].id.clone() + let index = index.saturating_sub(1); + if index >= topics.len() { + return Err(CommandError::new( + "INVALID_TOPIC_INDEX", + format!( + "Topic index {} is out of range (1-{})", + index + 1, + topics.len() + ), + )); + } + Ok(topics[index].id.clone()) + }) + .await + .map_err(|e| CommandError::new("LIST_TOPICS_ERROR", e.to_string()))?? } else { topic_id }; // 验证目标话题存在 - let topic = handler - .store - .get_topic(&target_topic_id) - .map_err(|e| CommandError::new("SWITCH_TOPIC_ERROR", e.to_string()))? + let topic = { + let store = handler.store.clone(); + let topic_id_bg = target_topic_id.clone(); + tokio::task::spawn_blocking(move || { + store + .get_topic(&topic_id_bg) + .map_err(|e| CommandError::new("SWITCH_TOPIC_ERROR", e.to_string())) + }) + .await + .map_err(|e| CommandError::new("SWITCH_TOPIC_ERROR", e.to_string()))?? .ok_or_else(|| { CommandError::new( "TOPIC_NOT_FOUND", format!("Topic not found: {}", target_topic_id), ) - })?; + })? + }; // 如果有 SessionManager,实际切换话题历史 if let Some(ref session_manager) = handler.session_manager @@ -113,10 +126,15 @@ async fn handle_switch_topic( } // 使用辅助方法获取消息数量 - let msg_count = handler - .store - .get_topic_message_count(&target_topic_id) - .unwrap_or(0); + let msg_count = { + let store = handler.store.clone(); + let topic_id_bg = target_topic_id.clone(); + tokio::task::spawn_blocking(move || store.get_topic_message_count(&topic_id_bg)) + .await + .ok() + .and_then(|r| r.ok()) + .unwrap_or(0) + }; let message = format!( "✓ Switched to topic: {} ({} messages)", diff --git a/src/gateway/http.rs b/src/gateway/http.rs index 9dad97c..41dadc3 100644 --- a/src/gateway/http.rs +++ b/src/gateway/http.rs @@ -1246,9 +1246,14 @@ pub async fn topic_select_model( // 同时取 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) => { + let topic_lookup = { + let store_bg = store.clone(); + let topic_id_bg = req.topic_id.clone(); + tokio::task::spawn_blocking(move || store_bg.get_topic(&topic_id_bg)).await + }; + let topic_session_id = match topic_lookup { + Ok(Ok(Some(topic))) => topic.session_id, + Ok(Ok(None)) => { return ( StatusCode::NOT_FOUND, Json(SelectModelResponse { @@ -1257,12 +1262,17 @@ pub async fn topic_select_model( }), ); } - Err(e) => { + other => { + let err = match other { + Ok(Err(e)) => e.to_string(), + Err(e) => e.to_string(), + _ => unreachable!(), + }; return ( StatusCode::INTERNAL_SERVER_ERROR, Json(SelectModelResponse { success: false, - error: Some(format!("failed to load topic: {}", e)), + error: Some(format!("failed to load topic: {}", err)), }), ); } @@ -1295,7 +1305,17 @@ pub async fn topic_select_model( 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()) { + let update_result = { + let store_bg = store.clone(); + let topic_id_bg = req.topic_id.clone(); + let provider_bg = provider.clone(); + let model_bg = model.clone(); + tokio::task::spawn_blocking(move || { + store_bg.update_topic_model(&topic_id_bg, provider_bg.as_deref(), model_bg.as_deref()) + }) + .await + }; + if let Err(e) = update_result.map_err(|e| e.to_string()).and_then(|r| r.map_err(|e| e.to_string())) { return ( StatusCode::INTERNAL_SERVER_ERROR, Json(SelectModelResponse { @@ -1346,23 +1366,31 @@ pub async fn topic_selected_model( // 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(), - ); + None => { + // 同步 rusqlite 查询移入 blocking 线程池:本路由在每次打开/切换 + // 话题且缓存 miss 时被调用,不得阻塞 axum async worker + let store = state.session_manager.store().clone(); + let topic_id_bg = q.topic_id.clone(); + let topic_lookup = + tokio::task::spawn_blocking(move || store.get_topic(&topic_id_bg)).await; + match topic_lookup { + Ok(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)) } - state - .model_selections - .get(&topic.session_id) - .unwrap_or((None, None)) + _ => (None, None), } - _ => (None, None), - }, + } }; Json(SessionSelectedModelResponse { provider, model }) } diff --git a/src/storage/mod.rs b/src/storage/mod.rs index 9a8a35a..d674f6f 100644 --- a/src/storage/mod.rs +++ b/src/storage/mod.rs @@ -250,6 +250,11 @@ impl SessionStore { // 池内每个连接都必须单独设置。WAL + NORMAL 是 SQLite 官方推荐组合: // 消除每次 commit 的 WAL full fsync,仅 checkpoint 时同步。 c.pragma_update(None, "synchronous", "NORMAL")?; + // foreign_keys 同样是 per-connection 且默认关闭。schema 声明的 + // ON DELETE CASCADE / SET NULL(topics/skill_events/messages 引用 + // sessions,messages 引用 topics)只有在该 PRAGMA 开启时才执行—— + // 缺少它会导致删除会话/话题后子表行孤儿残留、消息留下悬空 topic_id。 + c.pragma_update(None, "foreign_keys", "ON")?; Ok(()) }); let pool = Pool::builder().max_size(8).build(manager)?; @@ -684,7 +689,10 @@ impl SessionStore { if total > 0 { conn.execute( "UPDATE sessions SET message_count = - (SELECT COUNT(*) FROM messages WHERE session_id = ?1) + (SELECT COUNT(*) FROM messages WHERE session_id = ?1), + user_turn_count = + (SELECT COUNT(*) FROM messages + WHERE session_id = ?1 AND role = 'user') WHERE id = ?1", params![session_id], )?; @@ -761,9 +769,11 @@ impl SessionStore { ", params![session_id, now], )?; - // token 统计增量列随消息清空归零(消息已全部删除,无需重算) + // topics 增量列随消息清空归零(消息已全部删除,无需重算): + // message_count 必须一并归零,否则前端话题列表显示虚假消息数 conn.execute( "UPDATE topics SET + message_count = 0, stat_prompt_tokens = 0, stat_completion_tokens = 0, stat_total_tokens = 0, @@ -1107,6 +1117,9 @@ impl SessionStore { tx.commit()?; // 会话级压缩重组了消息行(标记/删除/插入摘要),所有 topic 统计可能漂移 self.recompute_session_topic_usage_stats(session_id)?; + // 新插入消息经 INSERT_MESSAGE_SQL 写入(无 topic_id 列),被删的旧行 + // 可能带 topic_id——按活查询重算 topics.message_count 防漂移 + self.recompute_session_topic_message_counts(session_id)?; Ok(true) } @@ -1164,6 +1177,9 @@ impl SessionStore { tx.commit()?; // 整会话消息被替换,所有 topic 统计可能漂移——重算自愈 self.recompute_session_topic_usage_stats(session_id)?; + // 新插入消息经 INSERT_MESSAGE_SQL 写入(无 topic_id 列),topic 消息数 + // 按活查询重算防漂移 + self.recompute_session_topic_message_counts(session_id)?; Ok(()) } @@ -2300,6 +2316,25 @@ impl SessionStore { self.recompute_topic_usage_stats(&ids) } + /// 按活查询重算指定 session 下所有 topic 的 message_count(自愈漂移)。 + /// + /// 用于整会话级变更路径(replace_active_history / compact_active_history): + /// 这些路径重插的消息不带 topic_id,被删的旧行可能带 topic_id, + /// topics.message_count 增量列无法跟随,必须重算。 + pub fn recompute_session_topic_message_counts( + &self, + session_id: &str, + ) -> Result<(), StorageError> { + let conn = self.pool.get()?; + conn.execute( + "UPDATE topics SET message_count = + (SELECT COUNT(*) FROM messages WHERE topic_id = topics.id) + WHERE session_id = ?1", + params![session_id], + )?; + Ok(()) + } + /// 查询单个 session 的 token 消耗统计(cost 累计 + context 瞬时)。 /// /// 按 `session_id` 精确匹配查询,**不过滤** `sub:%`——专门用于子代理 session diff --git a/src/storage/tests.rs b/src/storage/tests.rs index fefb876..089bd03 100644 --- a/src/storage/tests.rs +++ b/src/storage/tests.rs @@ -1248,3 +1248,109 @@ fn test_backfill_topic_usage_stats_restores_incremental_columns() { let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap(); assert_eq!(stats.get(&topic.id).unwrap().total_tokens, 430); } + +#[test] +fn test_foreign_keys_enforced_on_pool_connections() { + let store = SessionStore::in_memory().unwrap(); + let session = store.create_cli_session(Some("fk")).unwrap(); + let topic = store.create_topic(&session.id, "topic-fk", None).unwrap(); + store + .append_message_with_topic(&session.id, Some(&topic.id), &ChatMessage::user("m1")) + .unwrap(); + + // 删除话题:ON DELETE SET NULL 必须生效(消息保留,topic_id 置空) + store.delete_topic(&topic.id).unwrap(); + let conn = store.pool.get().unwrap(); + let dangling: i64 = conn + .query_row( + "SELECT COUNT(*) FROM messages WHERE session_id = ?1 AND topic_id IS NOT NULL", + params![session.id], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(dangling, 0, "SET NULL 未生效:删除话题后消息仍持悬空 topic_id"); + let remaining: i64 = conn + .query_row( + "SELECT COUNT(*) FROM messages WHERE session_id = ?1", + params![session.id], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(remaining, 1, "消息本身应保留,仅 topic_id 置空"); + drop(conn); + + // 删除会话:topics 行必须级联删除,不得孤儿残留 + store.delete_session(&session.id).unwrap(); + let conn = store.pool.get().unwrap(); + let orphan_topics: i64 = conn + .query_row( + "SELECT COUNT(*) FROM topics WHERE session_id = ?1", + params![session.id], + |row| row.get(0), + ) + .unwrap(); + assert_eq!(orphan_topics, 0, "CASCADE 未生效:删除会话后 topics 孤儿残留"); +} + +#[test] +fn test_clear_messages_resets_topic_message_count() { + let store = SessionStore::in_memory().unwrap(); + let session = store.create_cli_session(Some("clear-count")).unwrap(); + let topic = store.create_topic(&session.id, "topic-clear", None).unwrap(); + for i in 1..=3 { + store + .append_message_with_topic( + &session.id, + Some(&topic.id), + &ChatMessage::user(format!("m{i}")), + ) + .unwrap(); + } + assert_eq!(store.get_topic_message_count(&topic.id).unwrap(), 3); + + store.clear_messages(&session.id).unwrap(); + + // 清空后话题消息数必须归零,前端不得显示虚假计数 + assert_eq!(store.get_topic_message_count(&topic.id).unwrap(), 0); + let topic_row = store.get_topic(&topic.id).unwrap().unwrap(); + assert_eq!(topic_row.message_count, 0); +} + +#[test] +fn test_delete_messages_by_ids_recomputes_user_turn_count() { + let store = SessionStore::in_memory().unwrap(); + let session = store.create_cli_session(Some("turn-count")).unwrap(); + let topic = store.create_topic(&session.id, "topic-turn", None).unwrap(); + + store + .append_message_with_topic(&session.id, Some(&topic.id), &ChatMessage::user("u1")) + .unwrap(); + store + .append_message_with_topic( + &session.id, + Some(&topic.id), + &assistant_with_usage("a1", 10, 5, 0, None), + ) + .unwrap(); + store + .append_message_with_topic(&session.id, Some(&topic.id), &ChatMessage::user("u2")) + .unwrap(); + + let before = store.get_session(&session.id).unwrap().unwrap(); + assert_eq!(before.user_turn_count, 2); + + // 删除 user 消息(sanitize 场景):user_turn_count 必须同步重算 + let msgs = store.load_messages_for_topic(&topic.id, None).unwrap(); + let user_ids: Vec = msgs + .iter() + .filter(|m| m.role == "user") + .map(|m| m.id.clone()) + .collect(); + store + .delete_messages_by_ids(&session.id, &user_ids[..1]) + .unwrap(); + + let after = store.get_session(&session.id).unwrap().unwrap(); + assert_eq!(after.user_turn_count, 1); + assert_eq!(after.message_count, 2); +}