PicoBot/src/storage/tests.rs
oudecheng 1cde659bb6 fix(storage): 全项目扫描修复——外键约束失效根因与计数一致性
- [根因] 连接池补 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 重算
2026-08-18 23:31:11 +08:00

1357 lines
49 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

use super::migrations::has_column;
use super::*;
use crate::bus::SYSTEM_CONTEXT_AGENT_PROMPT;
use crate::domain::messages::ToolCall;
const TEST_CHANNEL: &str = "test-channel";
#[test]
fn test_persistent_session_id_for_cli_and_channel() {
assert_eq!(persistent_session_id("cli", "abc"), "abc");
// 幂等:已带前缀的 chat_id 会被清理,不会累积前缀
assert_eq!(persistent_session_id("websocket", "websocket:abc"), "abc");
assert_eq!(
persistent_session_id("websocket", "websocket:websocket:abc"),
"abc"
);
assert_eq!(
persistent_session_id(TEST_CHANNEL, "abc"),
"test-channel:abc"
);
// 其他通道也幂等
assert_eq!(
persistent_session_id(TEST_CHANNEL, "test-channel:abc"),
"test-channel:abc"
);
}
#[test]
fn test_session_store_roundtrip_and_lifecycle() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("demo")).unwrap();
assert_eq!(session.title, "demo");
assert_eq!(session.channel_name, "cli");
assert_eq!(session.chat_id, session.id);
assert_eq!(session.message_count, 0);
assert_eq!(session.user_turn_count, 0);
assert_eq!(session.agent_prompt_reinjection_count, 0);
let first = ChatMessage::user("hello");
let second = ChatMessage::assistant("world");
store.append_message(&session.id, &first).unwrap();
store.append_message(&session.id, &second).unwrap();
let stored = store.get_session(&session.id).unwrap().unwrap();
assert_eq!(stored.message_count, 2);
assert!(stored.archived_at.is_none());
assert_eq!(stored.user_turn_count, 1);
assert_eq!(stored.agent_prompt_reinjection_count, 0);
let messages = store.load_messages(&session.id).unwrap();
assert_eq!(messages.len(), 2);
assert_eq!(messages[0].role, "user");
assert_eq!(messages[0].content, "hello");
assert_eq!(messages[1].role, "assistant");
assert_eq!(messages[1].content, "world");
store.rename_session(&session.id, "renamed").unwrap();
let renamed = store.get_session(&session.id).unwrap().unwrap();
assert_eq!(renamed.title, "renamed");
store.archive_session(&session.id).unwrap();
let archived = store.get_session(&session.id).unwrap().unwrap();
assert!(archived.archived_at.is_some());
let active_only = store.list_sessions("cli", false).unwrap();
assert!(active_only.is_empty());
let including_archived = store.list_sessions("cli", true).unwrap();
assert_eq!(including_archived.len(), 1);
store.clear_messages(&session.id).unwrap();
let cleared = store.load_messages(&session.id).unwrap();
assert!(cleared.is_empty());
let cleared_session = store.get_session(&session.id).unwrap().unwrap();
assert_eq!(cleared_session.message_count, 0);
assert_eq!(cleared_session.user_turn_count, 0);
assert_eq!(cleared_session.agent_prompt_reinjection_count, 0);
store.delete_session(&session.id).unwrap();
assert!(store.get_session(&session.id).unwrap().is_none());
}
#[test]
fn test_ensure_channel_session_is_stable() {
let store = SessionStore::in_memory().unwrap();
let first = store
.ensure_channel_session(TEST_CHANNEL, "chat-1")
.unwrap();
let second = store
.ensure_channel_session(TEST_CHANNEL, "chat-1")
.unwrap();
assert_eq!(first.id, second.id);
assert_eq!(first.chat_id, "chat-1");
assert_eq!(second.channel_name, TEST_CHANNEL);
}
#[test]
fn test_assistant_tool_calls_roundtrip() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("tools")).unwrap();
let assistant = ChatMessage::assistant_with_tool_calls(
"calling tool",
vec![ToolCall {
id: "call_1".to_string(),
name: "calculator".to_string(),
arguments: serde_json::json!({ "expression": "3*7" }),
}],
);
store.append_message(&session.id, &assistant).unwrap();
let messages = store.load_messages(&session.id).unwrap();
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].role, "assistant");
assert_eq!(messages[0].tool_calls.as_ref().unwrap().len(), 1);
assert_eq!(messages[0].tool_calls.as_ref().unwrap()[0].id, "call_1");
assert_eq!(
messages[0].tool_calls.as_ref().unwrap()[0].name,
"calculator"
);
}
#[test]
fn test_assistant_reasoning_content_roundtrip() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("reasoning")).unwrap();
let assistant = ChatMessage::assistant_with_reasoning("final answer", "hidden reasoning");
store.append_message(&session.id, &assistant).unwrap();
let messages = store.load_messages(&session.id).unwrap();
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].content, "final answer");
assert_eq!(
messages[0].reasoning_content.as_deref(),
Some("hidden reasoning")
);
}
#[test]
fn test_schema_migration_adds_user_turn_and_reinjection_columns() {
let tmp = std::env::temp_dir().join(format!("picobot_test_mig2_{}.db", uuid::Uuid::new_v4()));
let conn = Connection::open(&tmp).unwrap();
conn.execute_batch(
"
CREATE TABLE sessions (
id TEXT PRIMARY KEY,
title TEXT NOT NULL,
channel_name TEXT NOT NULL,
chat_id TEXT NOT NULL,
summary TEXT,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
last_active_at INTEGER NOT NULL,
archived_at INTEGER,
deleted_at INTEGER,
message_count INTEGER NOT NULL DEFAULT 0
);
CREATE TABLE messages (
id TEXT PRIMARY KEY,
session_id TEXT NOT NULL,
seq INTEGER NOT NULL,
role TEXT NOT NULL,
content TEXT NOT NULL,
media_refs_json TEXT NOT NULL,
tool_call_id TEXT,
tool_name TEXT,
tool_calls_json TEXT,
created_at INTEGER NOT NULL,
FOREIGN KEY(session_id) REFERENCES sessions(id) ON DELETE CASCADE,
UNIQUE(session_id, seq)
);
",
)
.unwrap();
let path_str = tmp.to_string_lossy().to_string();
let store = SessionStore::from_connection(conn, &path_str).unwrap();
let session = store.create_cli_session(Some("migrated")).unwrap();
assert_eq!(session.user_turn_count, 0);
assert_eq!(session.agent_prompt_reinjection_count, 0);
}
#[test]
fn test_schema_migration_adds_reasoning_content_column_to_messages() {
let tmp = std::env::temp_dir().join(format!("picobot_test_mig_{}.db", uuid::Uuid::new_v4()));
let conn = Connection::open(&tmp).unwrap();
conn.execute_batch(
"
CREATE TABLE sessions (
id TEXT PRIMARY KEY,
title TEXT NOT NULL,
channel_name TEXT NOT NULL,
chat_id TEXT NOT NULL,
summary TEXT,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
last_active_at INTEGER NOT NULL,
archived_at INTEGER,
deleted_at INTEGER,
message_count INTEGER NOT NULL DEFAULT 0
);
CREATE TABLE messages (
id TEXT PRIMARY KEY,
session_id TEXT NOT NULL,
seq INTEGER NOT NULL,
role TEXT NOT NULL,
content TEXT NOT NULL,
media_refs_json TEXT NOT NULL,
tool_call_id TEXT,
tool_name TEXT,
tool_calls_json TEXT,
created_at INTEGER NOT NULL,
FOREIGN KEY(session_id) REFERENCES sessions(id) ON DELETE CASCADE,
UNIQUE(session_id, seq)
);
",
)
.unwrap();
let path_str = tmp.to_string_lossy().to_string();
let _store = SessionStore::from_connection(conn, &path_str).unwrap();
let conn = _store.pool.get().unwrap();
assert!(has_column(&conn, "messages", "reasoning_content").unwrap());
}
#[test]
fn test_compact_active_history_rebuilds_active_segment_with_delta_messages() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("compact-history")).unwrap();
let agent_prompt =
ChatMessage::system_with_context("agent", Some(SYSTEM_CONTEXT_AGENT_PROMPT.to_string()));
let seed_messages = vec![
agent_prompt.clone(),
ChatMessage::user("u1"),
ChatMessage::assistant("a1"),
ChatMessage::user("u2"),
ChatMessage::assistant("a2"),
ChatMessage::user("u3"),
ChatMessage::assistant("a3"),
ChatMessage::user("u4"),
ChatMessage::assistant("a4"),
];
for message in &seed_messages {
store.append_message(&session.id, message).unwrap();
}
let snapshot_end_seq = store
.get_session(&session.id)
.unwrap()
.unwrap()
.message_count;
let preserved_messages = store.load_messages(&session.id).unwrap()[3..].to_vec();
let preserved_system_messages = vec![agent_prompt];
store
.append_message(&session.id, &ChatMessage::user("u5"))
.unwrap();
store
.append_message(&session.id, &ChatMessage::assistant("a5"))
.unwrap();
let summary_message = ChatMessage::system("[Compressed History]\n\nsummary");
let compacted = store
.compact_active_history(
&session.id,
snapshot_end_seq,
&preserved_system_messages,
&summary_message,
&preserved_messages,
)
.unwrap();
assert!(compacted);
let active_messages = store.load_messages(&session.id).unwrap();
assert_eq!(active_messages.len(), 10);
assert_eq!(active_messages[0].role, "system");
assert_eq!(active_messages[0].content, "agent");
assert_eq!(
active_messages[0].system_context.as_deref(),
Some(SYSTEM_CONTEXT_AGENT_PROMPT)
);
assert_eq!(active_messages[1].role, "system");
assert_eq!(
active_messages[1].content,
"[Compressed History]\n\nsummary"
);
assert_eq!(active_messages[2].content, "u2");
assert_eq!(active_messages[3].content, "a2");
assert_eq!(active_messages[8].content, "u5");
assert_eq!(active_messages[9].content, "a5");
let stored = store.get_session(&session.id).unwrap().unwrap();
assert_eq!(stored.user_turn_count, 4);
let all_messages = store.load_all_messages(&session.id).unwrap();
assert_eq!(all_messages.len(), 10);
}
#[test]
fn test_mark_agent_prompt_reinjected_increments_counter() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("prompt")).unwrap();
store.mark_agent_prompt_reinjected(&session.id).unwrap();
store.mark_agent_prompt_reinjected(&session.id).unwrap();
let stored = store.get_session(&session.id).unwrap().unwrap();
assert_eq!(stored.agent_prompt_reinjection_count, 2);
}
#[test]
fn test_tool_result_roundtrip() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("tool-result")).unwrap();
let tool_message = ChatMessage::tool("call_9", "write", "saved to /tmp/output.txt");
store.append_message(&session.id, &tool_message).unwrap();
let messages = store.load_messages(&session.id).unwrap();
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].role, "tool");
assert_eq!(messages[0].content, "saved to /tmp/output.txt");
assert_eq!(messages[0].tool_call_id.as_deref(), Some("call_9"));
assert_eq!(messages[0].tool_name.as_deref(), Some("write"));
assert!(messages[0].tool_calls.is_none());
}
#[test]
fn test_skill_events_roundtrip() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("skill-events")).unwrap();
store
.append_skill_event(None, "discovered", None, &serde_json::json!({"count": 2}))
.unwrap();
store
.append_skill_event(
Some(&session.id),
"activated",
Some("code-review"),
&serde_json::json!({"source": "project"}),
)
.unwrap();
let global_events = store.list_skill_events(None).unwrap();
assert_eq!(global_events.len(), 1);
assert_eq!(global_events[0].event_type, "discovered");
assert_eq!(global_events[0].payload["count"], 2);
let session_events = store.list_skill_events(Some(&session.id)).unwrap();
assert_eq!(session_events.len(), 1);
assert_eq!(session_events[0].event_type, "activated");
assert_eq!(session_events[0].skill_name.as_deref(), Some("code-review"));
assert_eq!(session_events[0].payload["source"], "project");
}
#[test]
fn test_memory_roundtrip_with_source_fields() {
let store = SessionStore::in_memory().unwrap();
let saved = store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-1", TEST_CHANNEL),
namespace: "user".to_string(),
memory_key: "language".to_string(),
content: "Rust".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-1", TEST_CHANNEL)),
source_message_id: Some("msg-1".to_string()),
source_message_seq: Some(7),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-1".to_string()),
})
.unwrap();
assert_eq!(saved.content, "Rust");
assert_eq!(saved.source_type, "message");
assert_eq!(
saved.source_session_id.as_deref(),
Some("test-channel:chat-1")
);
assert_eq!(saved.source_message_id.as_deref(), Some("msg-1"));
assert_eq!(saved.source_message_seq, Some(7));
let fetched = store
.get_memory("user", "test-channel:user-1", "user", "language")
.unwrap()
.unwrap();
assert_eq!(fetched.id, saved.id);
assert_eq!(fetched.source_chat_id.as_deref(), Some("chat-1"));
}
#[test]
fn test_memory_fts_tracks_upsert_and_delete() {
let store = SessionStore::in_memory().unwrap();
store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-1", TEST_CHANNEL),
namespace: "user".to_string(),
memory_key: "editor".to_string(),
content: "Prefers rust-analyzer and cargo test output".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-2", TEST_CHANNEL)),
source_message_id: Some("msg-2".to_string()),
source_message_seq: Some(3),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-2".to_string()),
})
.unwrap();
let hits = store
.search_memories("user", "test-channel:user-1", "rust-analyzer", None, 10)
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].memory_key, "editor");
store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-1", TEST_CHANNEL),
namespace: "user".to_string(),
memory_key: "editor".to_string(),
content: "Prefers clippy diagnostics".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-3", TEST_CHANNEL)),
source_message_id: Some("msg-3".to_string()),
source_message_seq: Some(4),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-3".to_string()),
})
.unwrap();
let old_hits = store
.search_memories("user", "test-channel:user-1", "rust-analyzer", None, 10)
.unwrap();
assert!(old_hits.is_empty());
let new_hits = store
.search_memories("user", "test-channel:user-1", "clippy", None, 10)
.unwrap();
assert_eq!(new_hits.len(), 1);
let deleted = store
.delete_memory("user", "test-channel:user-1", "user", "editor")
.unwrap();
assert!(deleted);
let hits_after_delete = store
.search_memories("user", "test-channel:user-1", "clippy", None, 10)
.unwrap();
assert!(hits_after_delete.is_empty());
}
#[test]
fn test_memory_search_matches_memory_key_field() {
let store = SessionStore::in_memory().unwrap();
store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-1", TEST_CHANNEL),
namespace: "user".to_string(),
memory_key: "email_folder_preference".to_string(),
content: "用户提到邮件时默认查看代收邮箱。".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-8", TEST_CHANNEL)),
source_message_id: Some("msg-8".to_string()),
source_message_seq: Some(8),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-8".to_string()),
})
.unwrap();
let hits = store
.search_memories(
"user",
"test-channel:user-1",
"email_folder_preference",
None,
10,
)
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].memory_key, "email_folder_preference");
}
#[test]
fn test_search_memories_any_matches_multiple_keywords_once() {
let store = SessionStore::in_memory().unwrap();
store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-1", TEST_CHANNEL),
namespace: "user".to_string(),
memory_key: "editor".to_string(),
content: "Prefers rust-analyzer and cargo test output".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-2", TEST_CHANNEL)),
source_message_id: Some("msg-2".to_string()),
source_message_seq: Some(3),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-2".to_string()),
})
.unwrap();
store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-1", TEST_CHANNEL),
namespace: "episodic".to_string(),
memory_key: "quality".to_string(),
content: "Tracks clippy warnings before release".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-3", TEST_CHANNEL)),
source_message_id: Some("msg-3".to_string()),
source_message_seq: Some(4),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-3".to_string()),
})
.unwrap();
let hits = store
.search_memories_any(
"user",
"test-channel:user-1",
&["rust-analyzer".to_string(), "clippy".to_string()],
None,
10,
)
.unwrap();
assert_eq!(hits.len(), 2);
assert!(hits.iter().any(|memory| memory.memory_key == "editor"));
assert!(hits.iter().any(|memory| memory.memory_key == "quality"));
}
#[test]
fn test_memory_scope_listing_and_full_scope_read() {
let store = SessionStore::in_memory().unwrap();
store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-2", TEST_CHANNEL),
namespace: "user".to_string(),
memory_key: "style".to_string(),
content: "偏好简洁表达".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-2", TEST_CHANNEL)),
source_message_id: Some("msg-2".to_string()),
source_message_seq: Some(2),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-2".to_string()),
})
.unwrap();
store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-1", TEST_CHANNEL),
namespace: "user".to_string(),
memory_key: "work".to_string(),
content: "用户在做AI产品".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-1", TEST_CHANNEL)),
source_message_id: Some("msg-1".to_string()),
source_message_seq: Some(1),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-1".to_string()),
})
.unwrap();
store
.put_memory(&MemoryUpsert {
scope_kind: "user".to_string(),
scope_key: format!("{}:user-1", TEST_CHANNEL),
namespace: "patterns".to_string(),
memory_key: "workflow".to_string(),
content: "习惯先问方案再要代码".to_string(),
source_type: "message".to_string(),
source_session_id: Some(format!("{}:chat-1", TEST_CHANNEL)),
source_message_id: Some("msg-3".to_string()),
source_message_seq: Some(3),
source_channel_name: Some(TEST_CHANNEL.to_string()),
source_chat_id: Some("chat-1".to_string()),
})
.unwrap();
let scope_keys = store.list_memory_scope_keys("user").unwrap();
assert_eq!(
scope_keys,
vec![
"test-channel:user-1".to_string(),
"test-channel:user-2".to_string()
]
);
let full_scope = store
.list_memories_for_scope("user", "test-channel:user-1")
.unwrap();
assert_eq!(full_scope.len(), 2);
assert!(
full_scope
.iter()
.all(|memory| memory.scope_key == "test-channel:user-1")
);
assert!(full_scope.iter().any(|memory| memory.memory_key == "work"));
assert!(
full_scope
.iter()
.any(|memory| memory.memory_key == "workflow")
);
}
#[test]
fn test_scheduler_job_roundtrip_and_runtime_update() {
let store = SessionStore::in_memory().unwrap();
let saved = store
.upsert_scheduler_job(&SchedulerJobUpsert {
id: "heartbeat".to_string(),
kind: "outbound_message".to_string(),
schedule: serde_json::json!({
"type": "interval",
"seconds": 300,
"startup_delay_secs": 10,
}),
interval_secs: 300,
startup_delay_secs: 10,
target: serde_json::json!({
"channel": "test-channel",
"chat_id": "oc_demo",
}),
payload: serde_json::json!({
"content": "heartbeat",
}),
enabled: true,
state: SchedulerJobState::Scheduled,
last_status: None,
last_error: None,
run_count: 0,
max_runs: Some(3),
last_fired_at: None,
next_fire_at: Some(1_700_000_000_000),
paused_at: None,
completed_at: None,
})
.unwrap();
assert_eq!(saved.id, "heartbeat");
assert_eq!(saved.kind, "outbound_message");
assert_eq!(saved.state, SchedulerJobState::Scheduled);
assert_eq!(saved.max_runs, Some(3));
store
.update_scheduler_job_runtime(
"heartbeat",
SchedulerJobState::Completed,
Some(SchedulerJobStatus::Ok),
None,
1,
Some(1_700_000_000_000),
None,
None,
Some(1_700_000_000_100),
)
.unwrap();
let fetched = store.get_scheduler_job("heartbeat").unwrap().unwrap();
assert_eq!(fetched.state, SchedulerJobState::Completed);
assert_eq!(fetched.last_status, Some(SchedulerJobStatus::Ok));
assert_eq!(fetched.run_count, 1);
assert_eq!(fetched.completed_at, Some(1_700_000_000_100));
}
#[test]
fn test_get_topic_message_count_uses_count_query() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("topic-count")).unwrap();
let topic = store.create_topic(&session.id, "topic-1", None).unwrap();
// 初始计数为 0
assert_eq!(store.get_topic_message_count(&topic.id).unwrap(), 0);
// 追加 3 条带 topic_id 的消息
for content in ["m1", "m2", "m3"] {
store
.append_message_with_topic(&session.id, Some(&topic.id), &ChatMessage::user(content))
.unwrap();
}
// 计数应为 3且不需要加载消息内容
assert_eq!(store.get_topic_message_count(&topic.id).unwrap(), 3);
// 另一个 topic 的计数应为 0隔离验证
let other_topic = store.create_topic(&session.id, "topic-2", None).unwrap();
assert_eq!(store.get_topic_message_count(&other_topic.id).unwrap(), 0);
// 不存在的 topic_id 返回 0
assert_eq!(
store.get_topic_message_count("topic:nonexistent").unwrap(),
0
);
}
/// 回归防护:所有按 topic_id 过滤的查询必须命中索引idx_messages_topic_seq
/// 或 idx_messages_session_*),不允许退化为全表扫描。
/// 消息表是最大的表,长会话场景下全表扫描是数据层主要退化点。
#[test]
fn test_topic_id_queries_use_index() {
let store = SessionStore::in_memory().unwrap();
// 索引必须存在partial indextopic_id IS NOT NULL
let conn = store.pool.get().unwrap();
let index_sql: String = conn
.query_row(
"SELECT sql FROM sqlite_master WHERE type='index' AND name='idx_messages_topic_seq'",
[],
|row| row.get(0),
)
.unwrap();
assert!(index_sql.contains("topic_id"));
// 覆盖所有 topic_id 查询形态(与 mod.rs 中实际 SQL 一致)
let queries: &[&str] = &[
// load_messages_for_topic带/不带 session_id 分支)
"SELECT id FROM messages WHERE topic_id = ?1 AND session_id = ?2 AND is_compacted = 0 ORDER BY seq ASC",
"SELECT id FROM messages WHERE topic_id = ?1 AND is_compacted = 0 ORDER BY seq ASC",
// load_messages_for_topic_full
"SELECT id FROM messages WHERE topic_id = ?1 AND session_id = ?2 ORDER BY seq ASC",
"SELECT id FROM messages WHERE topic_id = ?1 ORDER BY seq ASC",
// count_user_messages_for_topic / get_topic_message_count
"SELECT COUNT(*) FROM messages WHERE topic_id = ?1",
"SELECT COUNT(*) FROM messages WHERE topic_id = ?1 AND role = 'user' AND is_compacted = 0",
"SELECT COUNT(*) FROM messages WHERE topic_id = ?1 AND role = 'user'",
// 按会话+话题删除
"DELETE FROM messages WHERE session_id = ?1 AND topic_id = ?2",
// 话题聚合IN 列表)
"SELECT topic_id, COUNT(*) FROM messages WHERE topic_id IN (?1, ?2) AND role = 'assistant' GROUP BY topic_id",
];
for sql in queries {
let mut stmt = conn.prepare(&format!("EXPLAIN QUERY PLAN {sql}")).unwrap();
// EXPLAIN 不执行语句,但 rusqlite 要求参数计数匹配;按占位符数量传哑参数
let param_count = sql.matches('?').count();
let mut rows = stmt
.query(rusqlite::params_from_iter(vec!["x"; param_count]))
.unwrap();
let mut plan = String::new();
while let Some(row) = rows.next().unwrap() {
let detail: String = row.get(3).unwrap();
plan.push_str(&detail);
plan.push_str(" | ");
}
// partial indexWHERE topic_id IS NOT NULL可被 topic_id = ? / IN
// 等值查询使用;复合条件查询允许规划器选择 session 前缀索引。
assert!(
plan.contains("idx_messages_topic_seq") || plan.contains("idx_messages_session"),
"query degrades to full scan: {sql}\nplan: {plan}"
);
}
}
#[test]
fn test_repair_session_id_prefix_pollution() {
let store = SessionStore::in_memory().unwrap();
let clean = store.ensure_channel_session("websocket", "abc").unwrap();
assert_eq!(clean.id, "abc");
store
.append_message(&clean.id, &ChatMessage::user("clean msg"))
.unwrap();
let mut conn = store.pool.get().unwrap();
conn.execute("PRAGMA user_version = 1", []).unwrap();
conn.execute(
"INSERT INTO sessions (id, title, channel_name, chat_id, summary, created_at, updated_at, last_active_at, archived_at, deleted_at, message_count, user_turn_count, agent_prompt_reinjection_count) VALUES ('websocket:websocket:abc', 'Polluted', 'websocket', 'websocket:websocket:abc', NULL, 100, 100, 100, NULL, NULL, 1, 1, 0)",
[],
)
.unwrap();
conn.execute(
"INSERT INTO messages (id, session_id, seq, role, content, media_refs_json, created_at) VALUES ('m-polluted', 'websocket:websocket:abc', 1, 'user', 'old polluted msg', '[]', 100)",
[],
)
.unwrap();
conn.execute(
"INSERT INTO sessions (id, title, channel_name, chat_id, summary, created_at, updated_at, last_active_at, archived_at, deleted_at, message_count, user_turn_count, agent_prompt_reinjection_count) VALUES ('websocket:xyz', 'PollutedXyz', 'websocket', 'websocket:xyz', NULL, 200, 200, 200, NULL, NULL, 1, 1, 0)",
[],
)
.unwrap();
conn.execute(
"INSERT INTO messages (id, session_id, seq, role, content, media_refs_json, created_at) VALUES ('m-xyz', 'websocket:xyz', 1, 'user', 'xyz msg', '[]', 200)",
[],
)
.unwrap();
super::migrations::repair_session_id_prefix_pollution(&mut conn).unwrap();
assert!(
store
.get_session("websocket:websocket:abc")
.unwrap()
.is_none()
);
assert!(store.get_session("websocket:xyz").unwrap().is_none());
let merged = store.get_session("abc").unwrap().unwrap();
assert_eq!(merged.message_count, 2);
assert_eq!(merged.user_turn_count, 2);
let msgs = store.load_messages("abc").unwrap();
assert_eq!(msgs.len(), 2);
assert_eq!(msgs[0].content, "old polluted msg");
{
let mut stmt = conn
.prepare("SELECT seq FROM messages WHERE session_id = 'abc' ORDER BY seq")
.unwrap();
let seqs: Vec<i64> = stmt
.query_map([], |row| row.get(0))
.unwrap()
.collect::<Result<Vec<_>, _>>()
.unwrap();
assert_eq!(seqs, vec![1, 2]);
}
let xyz = store.get_session("xyz").unwrap().unwrap();
assert_eq!(xyz.chat_id, "xyz");
assert_eq!(store.load_messages("xyz").unwrap().len(), 1);
assert_eq!(store.list_sessions("websocket", false).unwrap().len(), 2);
super::migrations::repair_session_id_prefix_pollution(&mut conn).unwrap();
assert_eq!(store.get_session("abc").unwrap().unwrap().message_count, 2);
assert_eq!(store.list_sessions("websocket", false).unwrap().len(), 2);
}
#[test]
fn test_cleanup_legacy_empty_cli_sessions() {
let store = SessionStore::in_memory().unwrap();
// 空 cli 会话(旧版本遗留的空壳):应被删除
let empty = store.create_cli_session(None).unwrap();
// 有消息的 cli 会话:必须保留
let with_data = store.create_cli_session(None).unwrap();
store
.append_message(&with_data.id, &ChatMessage::user("hello"))
.unwrap();
// 空 websocket 会话:不在 cli 清理范围内
let ws = store.ensure_channel_session("websocket", "chat-1").unwrap();
let mut conn = store.pool.get().unwrap();
// from_connection 已把 user_version 推到 3回退到 2 模拟"尚未清理"
conn.execute("PRAGMA user_version = 2", []).unwrap();
super::migrations::cleanup_legacy_empty_cli_sessions(&mut conn).unwrap();
assert!(store.get_session(&empty.id).unwrap().is_none());
assert!(store.get_session(&with_data.id).unwrap().is_some());
assert!(store.get_session(&ws.id).unwrap().is_some());
// 幂等:版本守卫使第二次调用直接跳过,已有会话不受影响
super::migrations::cleanup_legacy_empty_cli_sessions(&mut conn).unwrap();
assert!(store.get_session(&with_data.id).unwrap().is_some());
assert!(store.get_session(&ws.id).unwrap().is_some());
}
#[test]
fn test_load_messages_for_topic_page_keyset_pagination() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("paged")).unwrap();
let topic = store.create_topic(&session.id, "topic-page", None).unwrap();
for i in 1..=10 {
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&ChatMessage::user(format!("m{i}")),
)
.unwrap();
}
// 首页:无游标,取最新 4 条(正序),且标记还有更早消息
let (page1, has_more1) = store
.load_messages_for_topic_page(&topic.id, Some(&session.id), None, 4)
.unwrap();
assert!(has_more1);
let contents: Vec<_> = page1.iter().map(|m| m.content.as_str()).collect();
assert_eq!(contents, vec!["m7", "m8", "m9", "m10"]);
// 历史加载路径必须填充 seq 作为下一页游标
let oldest1 = page1.first().and_then(|m| m.seq).expect("seq populated");
// 第二页before_seq 游标,取更早的 4 条
let (page2, has_more2) = store
.load_messages_for_topic_page(&topic.id, Some(&session.id), Some(oldest1), 4)
.unwrap();
assert!(has_more2);
let contents: Vec<_> = page2.iter().map(|m| m.content.as_str()).collect();
assert_eq!(contents, vec!["m3", "m4", "m5", "m6"]);
let oldest2 = page2.first().and_then(|m| m.seq).expect("seq populated");
// 末页:剩余 2 条has_more=false
let (page3, has_more3) = store
.load_messages_for_topic_page(&topic.id, Some(&session.id), Some(oldest2), 4)
.unwrap();
assert!(!has_more3);
let contents: Vec<_> = page3.iter().map(|m| m.content.as_str()).collect();
assert_eq!(contents, vec!["m1", "m2"]);
// 游标越过最老消息:空页且 has_more=false
let oldest3 = page3.first().and_then(|m| m.seq).unwrap();
let (page4, has_more4) = store
.load_messages_for_topic_page(&topic.id, Some(&session.id), Some(oldest3), 4)
.unwrap();
assert!(!has_more4);
assert!(page4.is_empty());
// 拼接所有页 == 全量加载(顺序与内容一致)
let mut combined = page3;
combined.extend(page2);
combined.extend(page1);
let full = store.load_messages_for_topic(&topic.id, None).unwrap();
let combined_contents: Vec<_> = combined.iter().map(|m| m.content.as_str()).collect();
let full_contents: Vec<_> = full.iter().map(|m| m.content.as_str()).collect();
assert_eq!(combined_contents, full_contents);
}
#[test]
fn test_delete_messages_by_ids_removes_only_target_rows() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("del")).unwrap();
let topic = store.create_topic(&session.id, "topic-del", None).unwrap();
for i in 1..=5 {
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&ChatMessage::user(format!("m{i}")),
)
.unwrap();
}
let all = store.load_messages_for_topic(&topic.id, None).unwrap();
assert_eq!(all.len(), 5);
// 删除中间 2 条sanitize 回写场景:只删被清理的消息)
let to_delete: Vec<String> = all[1..3].iter().map(|m| m.id.clone()).collect();
let deleted = store.delete_messages_by_ids(&session.id, &to_delete).unwrap();
assert_eq!(deleted, 2);
let remaining = store.load_messages_for_topic(&topic.id, None).unwrap();
assert_eq!(remaining.len(), 3);
let contents: Vec<_> = remaining.iter().map(|m| m.content.as_str()).collect();
assert_eq!(contents, vec!["m1", "m4", "m5"]);
// 空列表 no-op不存在的 id 返回 0
assert_eq!(store.delete_messages_by_ids(&session.id, &[]).unwrap(), 0);
assert_eq!(
store
.delete_messages_by_ids(&session.id, &["nonexistent".to_string()])
.unwrap(),
0
);
// 消息计数同步修正(活查询 + 计数列均一致)
assert_eq!(store.get_topic_message_count(&topic.id).unwrap(), 3);
let session_after = store.get_session(&session.id).unwrap().unwrap();
assert_eq!(session_after.message_count, 3);
}
fn assistant_with_usage(
content: &str,
prompt: u32,
completion: u32,
cached: u32,
ctx_window: Option<u32>,
) -> ChatMessage {
let mut msg = ChatMessage::assistant(content);
msg.usage = Some(crate::bus::MessageUsage {
prompt_tokens: prompt,
completion_tokens: completion,
total_tokens: prompt + completion,
cached_tokens: cached,
context_window_tokens: ctx_window,
});
msg
}
#[test]
fn test_topic_token_stats_incremental_maintenance() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("stats")).unwrap();
let topic = store.create_topic(&session.id, "topic-stats", None).unwrap();
// 无 usage 数据时不返回条目(前端不显示 token 标签)
store
.append_message_with_topic(&session.id, Some(&topic.id), &ChatMessage::user("q1"))
.unwrap();
let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap();
assert!(stats.is_empty());
// 两条带 usage 的 assistant 消息SUM 累加last_* 取最新一条
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("a1", 100, 50, 20, Some(128_000)),
)
.unwrap();
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("a2", 200, 80, 100, Some(200_000)),
)
.unwrap();
let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap();
let s = stats.get(&topic.id).expect("topic stats entry");
assert_eq!(s.prompt_tokens, 300);
assert_eq!(s.completion_tokens, 130);
assert_eq!(s.total_tokens, 430);
assert_eq!(s.cached_tokens, 120);
assert_eq!(s.last_prompt_tokens, Some(200));
assert_eq!(s.context_window_tokens, Some(200_000));
}
#[test]
fn test_topic_token_stats_excludes_subagent_messages() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("stats-sub")).unwrap();
let topic = store.create_topic(&session.id, "topic-sub", None).unwrap();
// 主代理消息计入
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("main", 100, 50, 0, Some(128_000)),
)
.unwrap();
// 子代理消息 topic_id 指向父 topic 但 session_id 以 'sub:' 开头——不得计入
let sub_session = store
.ensure_session("sub:child-1", "cli", "sub-chat", "sub")
.unwrap();
store
.append_message_with_topic(
&sub_session.id,
Some(&topic.id),
&assistant_with_usage("sub", 999, 999, 999, Some(999_999)),
)
.unwrap();
let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap();
let s = stats.get(&topic.id).expect("topic stats entry");
assert_eq!(s.prompt_tokens, 100);
assert_eq!(s.total_tokens, 150);
assert_eq!(s.last_prompt_tokens, Some(100));
assert_eq!(s.context_window_tokens, Some(128_000));
}
#[test]
fn test_topic_token_stats_recompute_after_delete_and_clear() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("stats-recompute")).unwrap();
let topic = store.create_topic(&session.id, "topic-rc", None).unwrap();
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("a1", 100, 50, 0, Some(128_000)),
)
.unwrap();
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("a2", 200, 80, 0, Some(128_000)),
)
.unwrap();
// 删除第二条 assistant 消息sanitize 回写场景):统计重算为仅第一条
let msgs = store.load_messages_for_topic(&topic.id, None).unwrap();
let assistant_ids: Vec<String> = msgs
.iter()
.filter(|m| m.role == "assistant")
.map(|m| m.id.clone())
.collect();
store
.delete_messages_by_ids(&session.id, &assistant_ids[1..2])
.unwrap();
let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap();
let s = stats.get(&topic.id).expect("topic stats entry");
assert_eq!(s.prompt_tokens, 100);
assert_eq!(s.total_tokens, 150);
assert_eq!(s.last_prompt_tokens, Some(100));
// 删除全部 assistant 消息后无条目has_usage=0
store
.delete_messages_by_ids(&session.id, &assistant_ids[0..1])
.unwrap();
let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap();
assert!(stats.is_empty());
// 再次写入后恢复clear_messages 归零
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("a3", 300, 100, 0, None),
)
.unwrap();
assert!(!store
.batch_topic_token_stats(&[topic.id.as_str()])
.unwrap()
.is_empty());
store.clear_messages(&session.id).unwrap();
let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap();
assert!(stats.is_empty());
}
#[test]
fn test_topic_token_stats_recompute_after_replace_topic_history() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("stats-replace")).unwrap();
let topic = store.create_topic(&session.id, "topic-rp", None).unwrap();
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("old", 500, 200, 0, Some(128_000)),
)
.unwrap();
// 整体替换该 topic 历史(压缩器场景):统计应反映新消息
store
.replace_topic_history(
&session.id,
&topic.id,
&[
ChatMessage::user("new-q"),
assistant_with_usage("new-a", 40, 10, 5, Some(64_000)),
],
)
.unwrap();
let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap();
let s = stats.get(&topic.id).expect("topic stats entry");
assert_eq!(s.prompt_tokens, 40);
assert_eq!(s.completion_tokens, 10);
assert_eq!(s.total_tokens, 50);
assert_eq!(s.cached_tokens, 5);
assert_eq!(s.last_prompt_tokens, Some(40));
assert_eq!(s.context_window_tokens, Some(64_000));
}
#[test]
fn test_backfill_topic_usage_stats_restores_incremental_columns() {
let store = SessionStore::in_memory().unwrap();
let session = store.create_cli_session(Some("backfill")).unwrap();
let topic = store.create_topic(&session.id, "topic-bf", None).unwrap();
// 主代理 usage 消息两条(增量列此时已有值)
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("a1", 100, 50, 20, Some(128_000)),
)
.unwrap();
store
.append_message_with_topic(
&session.id,
Some(&topic.id),
&assistant_with_usage("a2", 200, 80, 0, Some(200_000)),
)
.unwrap();
// 子代理消息topic_id 指向父 topic回填必须排除
let sub_session = store
.ensure_session("sub:bf-child", "cli", "sub-chat", "sub")
.unwrap();
store
.append_message_with_topic(
&sub_session.id,
Some(&topic.id),
&assistant_with_usage("sub", 999, 999, 999, Some(999_999)),
)
.unwrap();
// 模拟老库:增量列全部清零,版本回退到 3尚未回填
let mut conn = store.pool.get().unwrap();
conn.execute(
"UPDATE topics SET stat_prompt_tokens = 0, stat_completion_tokens = 0, \
stat_total_tokens = 0, stat_cached_tokens = 0, \
stat_last_prompt_tokens = NULL, stat_context_window_tokens = NULL, \
stat_has_usage = 0",
[],
)
.unwrap();
conn.execute("PRAGMA user_version = 3", []).unwrap();
super::migrations::backfill_topic_usage_stats(&mut conn).unwrap();
// 版本守卫推进到 4
let version: i64 = conn
.query_row("PRAGMA user_version", [], |row| row.get(0))
.unwrap();
assert_eq!(version, 4);
drop(conn);
// 回填后与增量维护的值一致(子代理被排除)
let stats = store.batch_topic_token_stats(&[topic.id.as_str()]).unwrap();
let s = stats.get(&topic.id).expect("topic stats entry");
assert_eq!(s.prompt_tokens, 300);
assert_eq!(s.completion_tokens, 130);
assert_eq!(s.total_tokens, 430);
assert_eq!(s.cached_tokens, 20);
assert_eq!(s.last_prompt_tokens, Some(200));
assert_eq!(s.context_window_tokens, Some(200_000));
// 幂等:版本守卫使第二次调用直接跳过,结果不变
let mut conn = store.pool.get().unwrap();
super::migrations::backfill_topic_usage_stats(&mut conn).unwrap();
drop(conn);
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<String> = 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);
}