- [根因] 连接池补 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 重算
1357 lines
49 KiB
Rust
1357 lines
49 KiB
Rust
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 index:topic_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 index(WHERE 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);
|
||
}
|