1900 lines
65 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.

pub mod background_task;
pub mod error;
pub mod memory;
pub mod message;
pub mod scheduler;
pub mod session;
pub use background_task::BackgroundTask;
pub use error::StorageError;
pub use scheduler::{DeliveryPolicy, JobKind, JobRun, ScheduledJob};
use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions, SqliteSynchronous};
use sqlx::{Pool, Row, Sqlite};
use std::path::Path;
use tokio::time::{Duration, sleep};
const SCHEMA_VERSION: i64 = 3;
const INSERT_MESSAGE_SQL: &str = r#"
INSERT INTO messages (id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#;
fn insert_message_query<'a>(
session_id: &'a str,
msg: &'a crate::storage::message::MessageMeta,
) -> sqlx::query::Query<'a, Sqlite, sqlx::sqlite::SqliteArguments<'a>> {
sqlx::query(INSERT_MESSAGE_SQL)
.bind(&msg.id)
.bind(session_id)
.bind(msg.seq)
.bind(&msg.role)
.bind(&msg.content)
.bind(&msg.reasoning_content)
.bind(&msg.media_refs)
.bind(&msg.tool_call_id)
.bind(&msg.tool_name)
.bind(&msg.tool_calls)
.bind(&msg.source)
.bind(msg.created_at)
}
pub struct Storage {
pub(crate) pool: Pool<Sqlite>,
}
impl Storage {
/// 打开或创建数据库
pub async fn new(db_path: &Path) -> Result<Self, StorageError> {
let options = SqliteConnectOptions::new()
.filename(db_path)
.create_if_missing(true)
.journal_mode(SqliteJournalMode::Wal)
.synchronous(SqliteSynchronous::Normal)
.busy_timeout(Duration::from_secs(5))
.foreign_keys(true);
let pool = SqlitePoolOptions::new()
.max_connections(8)
.connect_with(options)
.await?;
let storage = Self { pool };
storage.init_schema().await?;
Ok(storage)
}
/// 初始化数据库 schema
async fn init_schema(&self) -> Result<(), StorageError> {
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS sessions (
id TEXT PRIMARY KEY,
channel TEXT NOT NULL,
chat_id TEXT NOT NULL,
dialog_id TEXT NOT NULL,
title TEXT NOT NULL DEFAULT '新对话',
created_at INTEGER NOT NULL,
last_active_at INTEGER NOT NULL,
message_count INTEGER DEFAULT 0,
routing_info TEXT,
archived_at INTEGER,
deleted_at INTEGER,
last_consolidated_at INTEGER,
last_compressed_message_at INTEGER,
UNIQUE(channel, chat_id, dialog_id)
)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_sessions_chat
ON sessions(channel, chat_id, deleted_at)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS messages (
id TEXT PRIMARY KEY,
session_id TEXT NOT NULL,
seq INTEGER NOT NULL,
role TEXT NOT NULL,
content TEXT NOT NULL,
media_refs TEXT,
tool_call_id TEXT,
tool_name TEXT,
tool_calls TEXT,
source TEXT,
reasoning_content TEXT,
created_at INTEGER NOT NULL,
FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE CASCADE
)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_messages_session_seq
ON messages(session_id, seq)
"#,
)
.execute(&self.pool)
.await?;
// Background tasks table — for async sub-agent tasks.
// Note: No FOREIGN KEY on session_id because sessions use soft delete (deleted_at IS NULL).
// Session and task association is maintained at the application level.
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS background_tasks (
id TEXT PRIMARY KEY,
session_id TEXT NOT NULL,
channel TEXT NOT NULL,
chat_id TEXT NOT NULL,
prompt TEXT NOT NULL,
allowed_tools TEXT,
status TEXT NOT NULL DEFAULT 'pending',
result TEXT,
error TEXT,
tool_calls_count INTEGER DEFAULT 0,
iterations INTEGER DEFAULT 0,
started_at INTEGER,
finished_at INTEGER,
created_at INTEGER NOT NULL
)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_bg_tasks_session ON background_tasks(session_id)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE INDEX IF NOT EXISTS idx_bg_tasks_status ON background_tasks(status)
"#,
)
.execute(&self.pool)
.await?;
// Session-scoped task plans. A session may have at most one active plan,
// while independent items can be executed concurrently.
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS task_plans (
id TEXT PRIMARY KEY,
session_id TEXT NOT NULL,
objective TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'active',
version INTEGER NOT NULL DEFAULT 1,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
closed_at INTEGER
)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE UNIQUE INDEX IF NOT EXISTS idx_task_plans_active_session
ON task_plans(session_id) WHERE status = 'active'
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS task_items (
id TEXT NOT NULL,
plan_id TEXT NOT NULL,
ordinal INTEGER NOT NULL,
title TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'pending',
executor_kind TEXT,
execution_id TEXT,
result_summary TEXT,
error TEXT,
version INTEGER NOT NULL DEFAULT 1,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
PRIMARY KEY(plan_id, id),
FOREIGN KEY(plan_id) REFERENCES task_plans(id) ON DELETE CASCADE,
UNIQUE(plan_id, ordinal)
)
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
"CREATE INDEX IF NOT EXISTS idx_task_items_plan ON task_items(plan_id, ordinal)",
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS memories (
id TEXT PRIMARY KEY,
key TEXT NOT NULL UNIQUE,
content TEXT NOT NULL,
category TEXT NOT NULL DEFAULT 'knowledge',
importance REAL NOT NULL DEFAULT 0.5,
session_id TEXT,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
)
"#,
)
.execute(&self.pool)
.await?;
let memory_fts_exists: bool = sqlx::query_scalar(
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'memory_fts')",
)
.fetch_one(&self.pool)
.await?;
// FTS5 virtual table for full-text search on memories
sqlx::query(
r#"
CREATE VIRTUAL TABLE IF NOT EXISTS memory_fts USING fts5(
key,
content,
content=memories,
content_rowid=rowid
)
"#,
)
.execute(&self.pool)
.await?;
// Triggers to keep FTS5 index in sync with memories table
sqlx::query(
r#"
CREATE TRIGGER IF NOT EXISTS memories_ai AFTER INSERT ON memories BEGIN
INSERT INTO memory_fts(rowid, key, content) VALUES (new.rowid, new.key, new.content);
END
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE TRIGGER IF NOT EXISTS memories_ad AFTER DELETE ON memories BEGIN
INSERT INTO memory_fts(memory_fts, rowid, key, content)
VALUES ('delete', old.rowid, old.key, old.content);
END
"#,
)
.execute(&self.pool)
.await?;
sqlx::query(
r#"
CREATE TRIGGER IF NOT EXISTS memories_au AFTER UPDATE ON memories BEGIN
INSERT INTO memory_fts(memory_fts, rowid, key, content)
VALUES ('delete', old.rowid, old.key, old.content);
INSERT INTO memory_fts(rowid, key, content)
VALUES (new.rowid, new.key, new.content);
END
"#,
)
.execute(&self.pool)
.await?;
// Only a newly-created index needs a backfill. Triggers keep an
// existing index current, so rebuilding it on every startup is wasted
// work proportional to the total memory corpus.
if !memory_fts_exists {
sqlx::query("INSERT INTO memory_fts(memory_fts) VALUES ('rebuild')")
.execute(&self.pool)
.await?;
}
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS llm_calls (
id INTEGER PRIMARY KEY AUTOINCREMENT,
created_at INTEGER NOT NULL,
provider TEXT NOT NULL,
model TEXT NOT NULL,
request_body TEXT NOT NULL,
response_body TEXT,
error TEXT,
duration_ms INTEGER
)
"#,
)
.execute(&self.pool)
.await?;
Self::init_scheduler_schema(&self.pool).await?;
self.migrate_schema().await?;
Ok(())
}
/// Apply ordered, atomic migrations. Existing installations predate
/// `user_version`, so each step also checks the actual table shape.
async fn migrate_schema(&self) -> Result<(), StorageError> {
let current: i64 = sqlx::query_scalar("PRAGMA user_version")
.fetch_one(&self.pool)
.await?;
if current > SCHEMA_VERSION {
return Err(StorageError::Migration(format!(
"database schema version {current} is newer than supported version {SCHEMA_VERSION}"
)));
}
if current == SCHEMA_VERSION {
return Ok(());
}
let mut tx = self.pool.begin().await?;
for (table, column, definition) in [
("messages", "source", "source TEXT"),
("messages", "reasoning_content", "reasoning_content TEXT"),
("sessions", "archived_at", "archived_at INTEGER"),
(
"sessions",
"last_consolidated_at",
"last_consolidated_at INTEGER",
),
(
"sessions",
"last_compressed_message_at",
"last_compressed_message_at INTEGER",
),
("scheduled_jobs", "locked_at", "locked_at INTEGER"),
("scheduled_jobs", "lock_owner", "lock_owner TEXT"),
("scheduled_jobs", "lease_until", "lease_until INTEGER"),
(
"scheduled_jobs",
"job_kind",
"job_kind TEXT NOT NULL DEFAULT 'task'",
),
(
"scheduled_jobs",
"delivery_policy",
"delivery_policy TEXT NOT NULL DEFAULT 'direct'",
),
("job_runs", "result_kind", "result_kind TEXT"),
("job_runs", "delivery_status", "delivery_status TEXT"),
("job_runs", "delivery_error", "delivery_error TEXT"),
] {
let pragma = format!("PRAGMA table_info({table})");
let columns = sqlx::query(&pragma).fetch_all(&mut *tx).await?;
if !columns
.iter()
.any(|row| row.get::<String, _>("name") == column)
{
let alter = format!("ALTER TABLE {table} ADD COLUMN {definition}");
sqlx::query(&alter).execute(&mut *tx).await?;
}
}
let duplicate: Option<(String, i64, i64)> = sqlx::query_as(
r#"
SELECT session_id, seq, COUNT(*)
FROM messages
GROUP BY session_id, seq
HAVING COUNT(*) > 1
LIMIT 1
"#,
)
.fetch_optional(&mut *tx)
.await?;
if let Some((session_id, seq, count)) = duplicate {
return Err(StorageError::Migration(format!(
"cannot enforce unique message sequence: session {session_id} has {count} rows at seq {seq}"
)));
}
sqlx::query(
"CREATE UNIQUE INDEX IF NOT EXISTS idx_messages_session_seq_unique ON messages(session_id, seq)",
)
.execute(&mut *tx)
.await?;
sqlx::query(
"CREATE INDEX IF NOT EXISTS idx_jobs_claimable ON scheduled_jobs(enabled, next_run_at, lease_until)",
)
.execute(&mut *tx)
.await?;
sqlx::query(&format!("PRAGMA user_version = {SCHEMA_VERSION}"))
.execute(&mut *tx)
.await?;
tx.commit().await?;
Ok(())
}
/// Initialize the scheduler tables (idempotent).
pub(crate) async fn init_scheduler_schema(pool: &Pool<Sqlite>) -> Result<(), StorageError> {
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS scheduled_jobs (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
schedule TEXT NOT NULL,
prompt TEXT NOT NULL,
channel TEXT NOT NULL,
chat_id TEXT NOT NULL,
model TEXT,
job_kind TEXT NOT NULL DEFAULT 'task',
delivery_policy TEXT NOT NULL DEFAULT 'direct',
enabled INTEGER NOT NULL DEFAULT 1,
delete_after_run INTEGER NOT NULL DEFAULT 0,
next_run_at INTEGER NOT NULL,
last_run_at INTEGER,
last_status TEXT,
last_error TEXT,
locked_at INTEGER,
lock_owner TEXT,
lease_until INTEGER,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
)
"#,
)
.execute(pool)
.await?;
sqlx::query(
r#"
CREATE TABLE IF NOT EXISTS job_runs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
job_id TEXT NOT NULL REFERENCES scheduled_jobs(id) ON DELETE CASCADE,
started_at INTEGER NOT NULL,
finished_at INTEGER NOT NULL,
status TEXT NOT NULL,
output TEXT,
error TEXT,
duration_ms INTEGER NOT NULL,
result_kind TEXT,
delivery_status TEXT,
delivery_error TEXT
)
"#,
)
.execute(pool)
.await?;
sqlx::query(
"CREATE INDEX IF NOT EXISTS idx_jobs_next_run ON scheduled_jobs(enabled, next_run_at)",
)
.execute(pool)
.await?;
sqlx::query("CREATE INDEX IF NOT EXISTS idx_runs_job_id ON job_runs(job_id)")
.execute(pool)
.await?;
Ok(())
}
pub async fn append_llm_call(
&self,
provider: &str,
model: &str,
request_body: &str,
response_body: Option<&str>,
error: Option<&str>,
duration_ms: u64,
) -> Result<(), StorageError> {
let now = chrono::Utc::now().timestamp_millis();
sqlx::query(
r#"
INSERT INTO llm_calls (created_at, provider, model, request_body, response_body, error, duration_ms)
VALUES (?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(now)
.bind(provider)
.bind(model)
.bind(request_body)
.bind(response_body)
.bind(error)
.bind(duration_ms as i64)
.execute(self.pool())
.await?;
// Prune to keep last 1000 records
self.prune_llm_calls(1000).await?;
Ok(())
}
async fn prune_llm_calls(&self, max_records: i64) -> Result<(), StorageError> {
sqlx::query(
r#"
DELETE FROM llm_calls WHERE id <= (
SELECT COALESCE(MAX(id), 0) - ? FROM llm_calls
)
"#,
)
.bind(max_records)
.execute(self.pool())
.await?;
Ok(())
}
/// 获取连接池引用(供内部 CRUD 使用)
pub(crate) fn pool(&self) -> &Pool<Sqlite> {
&self.pool
}
pub async fn upsert_session(
&self,
meta: &crate::storage::session::SessionMeta,
) -> Result<(), StorageError> {
sqlx::query(
r#"
INSERT INTO sessions (id, channel, chat_id, dialog_id, title, created_at, last_active_at, message_count, routing_info, archived_at, deleted_at, last_consolidated_at, last_compressed_message_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
title = excluded.title,
last_active_at = excluded.last_active_at,
message_count = excluded.message_count,
routing_info = excluded.routing_info,
archived_at = excluded.archived_at,
deleted_at = excluded.deleted_at,
last_consolidated_at = excluded.last_consolidated_at,
last_compressed_message_at = excluded.last_compressed_message_at
"#,
)
.bind(&meta.id)
.bind(&meta.channel)
.bind(&meta.chat_id)
.bind(&meta.dialog_id)
.bind(&meta.title)
.bind(meta.created_at)
.bind(meta.last_active_at)
.bind(meta.message_count)
.bind(&meta.routing_info)
.bind(meta.archived_at)
.bind(meta.deleted_at)
.bind(meta.last_consolidated_at)
.bind(meta.last_compressed_message_at)
.execute(self.pool())
.await?;
Ok(())
}
pub async fn get_session(
&self,
id: &str,
) -> Result<crate::storage::session::SessionMeta, StorageError> {
let row = sqlx::query(
r#"
SELECT id, channel, chat_id, dialog_id, title, created_at, last_active_at, message_count, routing_info, archived_at, deleted_at, last_consolidated_at, last_compressed_message_at
FROM sessions WHERE id = ? AND deleted_at IS NULL
"#,
)
.bind(id)
.fetch_optional(self.pool())
.await?
.ok_or_else(|| StorageError::NotFound(id.to_string()))?;
Ok(crate::storage::session::SessionMeta {
id: row.get("id"),
channel: row.get("channel"),
chat_id: row.get("chat_id"),
dialog_id: row.get("dialog_id"),
title: row.get("title"),
created_at: row.get("created_at"),
last_active_at: row.get("last_active_at"),
message_count: row.get("message_count"),
routing_info: row.get("routing_info"),
archived_at: row.get("archived_at"),
deleted_at: row.get("deleted_at"),
last_consolidated_at: row.get("last_consolidated_at"),
last_compressed_message_at: row.get("last_compressed_message_at"),
})
}
pub async fn list_sessions(
&self,
channel: &str,
chat_id: &str,
limit: i64,
include_archived: bool,
) -> Result<Vec<crate::storage::session::SessionMeta>, StorageError> {
let rows = sqlx::query(
r#"
SELECT id, channel, chat_id, dialog_id, title, created_at, last_active_at, message_count, routing_info, archived_at, deleted_at, last_consolidated_at, last_compressed_message_at
FROM sessions
WHERE channel = ? AND chat_id = ? AND deleted_at IS NULL
AND (? OR archived_at IS NULL)
ORDER BY last_active_at DESC
LIMIT ?
"#,
)
.bind(channel)
.bind(chat_id)
.bind(include_archived)
.bind(limit)
.fetch_all(self.pool())
.await?;
Ok(rows
.into_iter()
.map(|row| crate::storage::session::SessionMeta {
id: row.get("id"),
channel: row.get("channel"),
chat_id: row.get("chat_id"),
dialog_id: row.get("dialog_id"),
title: row.get("title"),
created_at: row.get("created_at"),
last_active_at: row.get("last_active_at"),
message_count: row.get("message_count"),
routing_info: row.get("routing_info"),
archived_at: row.get("archived_at"),
deleted_at: row.get("deleted_at"),
last_consolidated_at: row.get("last_consolidated_at"),
last_compressed_message_at: row.get("last_compressed_message_at"),
})
.collect())
}
pub async fn touch_session(
&self,
id: &str,
message_count: i64,
last_active_at: i64,
) -> Result<(), StorageError> {
sqlx::query(
r#"
UPDATE sessions SET message_count = ?, last_active_at = ?
WHERE id = ?
"#,
)
.bind(message_count)
.bind(last_active_at)
.bind(id)
.execute(self.pool())
.await?;
Ok(())
}
pub async fn soft_delete_session(&self, id: &str) -> Result<(), StorageError> {
let now = chrono::Utc::now().timestamp_millis();
sqlx::query(r#"UPDATE sessions SET deleted_at = ? WHERE id = ?"#)
.bind(now)
.bind(id)
.execute(self.pool())
.await?;
Ok(())
}
pub async fn archive_session(&self, id: &str) -> Result<(), StorageError> {
let now = chrono::Utc::now().timestamp_millis();
sqlx::query(r#"UPDATE sessions SET archived_at = ? WHERE id = ? AND deleted_at IS NULL"#)
.bind(now)
.bind(id)
.execute(self.pool())
.await?;
Ok(())
}
pub async fn find_most_recent_session(
&self,
channel: &str,
chat_id: &str,
) -> Result<Option<crate::storage::session::SessionMeta>, StorageError> {
let row = sqlx::query(
r#"
SELECT id, channel, chat_id, dialog_id, title, created_at, last_active_at, message_count, routing_info, archived_at, deleted_at, last_consolidated_at, last_compressed_message_at
FROM sessions
WHERE channel = ? AND chat_id = ? AND deleted_at IS NULL AND archived_at IS NULL
ORDER BY last_active_at DESC
LIMIT 1
"#,
)
.bind(channel)
.bind(chat_id)
.fetch_optional(self.pool())
.await?;
match row {
Some(row) => Ok(Some(crate::storage::session::SessionMeta {
id: row.get("id"),
channel: row.get("channel"),
chat_id: row.get("chat_id"),
dialog_id: row.get("dialog_id"),
title: row.get("title"),
created_at: row.get("created_at"),
last_active_at: row.get("last_active_at"),
message_count: row.get("message_count"),
routing_info: row.get("routing_info"),
archived_at: row.get("archived_at"),
deleted_at: row.get("deleted_at"),
last_consolidated_at: row.get("last_consolidated_at"),
last_compressed_message_at: row.get("last_compressed_message_at"),
})),
None => Ok(None),
}
}
pub async fn append_message(
&self,
session_id: &str,
msg: &crate::storage::message::MessageMeta,
) -> Result<i64, StorageError> {
insert_message_query(session_id, msg)
.execute(self.pool())
.await?;
Ok(msg.seq)
}
/// Atomically persist all messages produced by one logical turn together
/// with the resulting session metadata. A turn is either fully visible
/// after restart or not visible at all.
pub async fn persist_message_batch(
&self,
session_id: &str,
msgs: &[crate::storage::message::MessageMeta],
meta: &crate::storage::session::SessionMeta,
) -> Result<(), StorageError> {
let mut tx = self.pool.begin().await?;
for msg in msgs {
insert_message_query(session_id, msg)
.execute(&mut *tx)
.await?;
}
sqlx::query(
r#"
INSERT INTO sessions (id, channel, chat_id, dialog_id, title, created_at, last_active_at, message_count, routing_info, archived_at, deleted_at, last_consolidated_at, last_compressed_message_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
title = excluded.title,
last_active_at = excluded.last_active_at,
message_count = excluded.message_count,
routing_info = excluded.routing_info,
archived_at = excluded.archived_at,
deleted_at = excluded.deleted_at,
last_consolidated_at = excluded.last_consolidated_at,
last_compressed_message_at = excluded.last_compressed_message_at
"#,
)
.bind(&meta.id)
.bind(&meta.channel)
.bind(&meta.chat_id)
.bind(&meta.dialog_id)
.bind(&meta.title)
.bind(meta.created_at)
.bind(meta.last_active_at)
.bind(meta.message_count)
.bind(&meta.routing_info)
.bind(meta.archived_at)
.bind(meta.deleted_at)
.bind(meta.last_consolidated_at)
.bind(meta.last_compressed_message_at)
.execute(&mut *tx)
.await?;
tx.commit().await?;
Ok(())
}
/// Persist a turn with bounded retry. Retrying the whole transaction keeps
/// message rows and metadata consistent on transient SQLite failures.
pub async fn persist_message_batch_with_retry(
&self,
session_id: &str,
msgs: &[crate::storage::message::MessageMeta],
meta: &crate::storage::session::SessionMeta,
) -> Result<(), StorageError> {
let delays = [100, 200, 300];
for (attempt, delay) in delays.iter().enumerate() {
match self.persist_message_batch(session_id, msgs, meta).await {
Ok(()) => return Ok(()),
Err(error) if attempt < delays.len() - 1 && error.is_transient() => {
tracing::warn!(attempt = attempt + 1, error = %error, "Turn persistence failed; retrying");
sleep(Duration::from_millis(*delay)).await;
}
Err(error) => return Err(error),
}
}
unreachable!()
}
pub async fn load_messages(
&self,
session_id: &str,
from_seq: i64,
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
let rows = sqlx::query(
r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at
FROM messages
WHERE session_id = ? AND seq >= ?
ORDER BY seq ASC
"#,
)
.bind(session_id)
.bind(from_seq)
.fetch_all(self.pool())
.await?;
Ok(rows
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect())
}
pub async fn get_message(
&self,
session_id: &str,
message_id: &str,
) -> Result<Option<crate::storage::message::MessageMeta>, StorageError> {
let row = sqlx::query(
r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs,
tool_call_id, tool_name, tool_calls, source, created_at
FROM messages
WHERE session_id = ? AND id = ?
"#,
)
.bind(session_id)
.bind(message_id)
.fetch_optional(self.pool())
.await?;
Ok(row.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
}))
}
pub async fn get_max_message_seq(&self, session_id: &str) -> Result<i64, StorageError> {
let row = sqlx::query(
"SELECT COALESCE(MAX(seq), 0) as max_seq FROM messages WHERE session_id = ?",
)
.bind(session_id)
.fetch_one(self.pool())
.await?;
Ok(row.get::<i64, _>("max_seq"))
}
/// Load a bounded tail of one session while preserving chronological order.
pub async fn load_recent_session_messages(
&self,
session_id: &str,
limit: u32,
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
let limit = limit.clamp(1, 2_000);
let rows = sqlx::query(
r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs,
tool_call_id, tool_name, tool_calls, source, created_at
FROM (
SELECT id, session_id, seq, role, content, reasoning_content, media_refs,
tool_call_id, tool_name, tool_calls, source, created_at
FROM messages
WHERE session_id = ?
ORDER BY seq DESC
LIMIT ?
)
ORDER BY seq ASC
"#,
)
.bind(session_id)
.bind(i64::from(limit))
.fetch_all(self.pool())
.await?;
Ok(rows
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect())
}
pub async fn load_messages_after_timestamp(
&self,
session_id: &str,
after_ts: i64,
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
let rows = sqlx::query(
r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at
FROM messages
WHERE session_id = ? AND created_at > ?
ORDER BY seq ASC
"#,
)
.bind(session_id)
.bind(after_ts)
.fetch_all(self.pool())
.await?;
Ok(rows
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect())
}
pub async fn query_sessions_range(
&self,
offset: i64,
limit: i64,
) -> Result<(Vec<crate::storage::session::SessionMeta>, i64), StorageError> {
let count_row =
sqlx::query("SELECT COUNT(*) as total FROM sessions WHERE deleted_at IS NULL")
.fetch_one(self.pool())
.await?;
let total: i64 = count_row.get("total");
let rows = sqlx::query(
r#"
SELECT id, channel, chat_id, dialog_id, title, created_at, last_active_at, message_count, routing_info, archived_at, deleted_at, last_consolidated_at, last_compressed_message_at
FROM sessions
WHERE deleted_at IS NULL
ORDER BY last_active_at DESC
LIMIT ? OFFSET ?
"#,
)
.bind(limit)
.bind(offset)
.fetch_all(self.pool())
.await?;
let sessions: Vec<_> = rows
.into_iter()
.map(|row| crate::storage::session::SessionMeta {
id: row.get("id"),
channel: row.get("channel"),
chat_id: row.get("chat_id"),
dialog_id: row.get("dialog_id"),
title: row.get("title"),
created_at: row.get("created_at"),
last_active_at: row.get("last_active_at"),
message_count: row.get("message_count"),
routing_info: row.get("routing_info"),
archived_at: row.get("archived_at"),
deleted_at: row.get("deleted_at"),
last_consolidated_at: row.get("last_consolidated_at"),
last_compressed_message_at: row.get("last_compressed_message_at"),
})
.collect();
Ok((sessions, total))
}
pub async fn list_recent_messages(
&self,
session_id: &str,
count: i64,
) -> Result<Vec<crate::storage::message::MessageMeta>, StorageError> {
let rows = sqlx::query(
r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at
FROM messages
WHERE session_id = ?
ORDER BY seq DESC
LIMIT ?
"#,
)
.bind(session_id)
.bind(count)
.fetch_all(self.pool())
.await?;
let mut messages: Vec<_> = rows
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect();
messages.reverse();
Ok(messages)
}
pub async fn query_messages_range(
&self,
session_id: &str,
before_time: Option<i64>,
after_time: Option<i64>,
offset: i64,
limit: i64,
) -> Result<(Vec<crate::storage::message::MessageMeta>, i64), StorageError> {
let mut where_extra = String::new();
if before_time.is_some() {
where_extra.push_str(" AND created_at < ?");
}
if after_time.is_some() {
where_extra.push_str(" AND created_at > ?");
}
let count_sql = format!(
"SELECT COUNT(*) as total FROM messages WHERE session_id = ?{}",
where_extra
);
let select_sql = format!(
r#"
SELECT id, session_id, seq, role, content, reasoning_content, media_refs, tool_call_id, tool_name, tool_calls, source, created_at
FROM messages
WHERE session_id = ?{}
ORDER BY seq ASC
LIMIT ? OFFSET ?
"#,
where_extra
);
let mut count_query = sqlx::query(&count_sql).bind(session_id);
if let Some(bt) = before_time {
count_query = count_query.bind(bt);
}
if let Some(at) = after_time {
count_query = count_query.bind(at);
}
let count_row = count_query.fetch_one(self.pool()).await?;
let total: i64 = count_row.get("total");
let mut select_query = sqlx::query(&select_sql).bind(session_id);
if let Some(bt) = before_time {
select_query = select_query.bind(bt);
}
if let Some(at) = after_time {
select_query = select_query.bind(at);
}
let rows = select_query
.bind(limit)
.bind(offset)
.fetch_all(self.pool())
.await?;
let messages: Vec<_> = rows
.into_iter()
.map(|row| crate::storage::message::MessageMeta {
id: row.get("id"),
session_id: row.get("session_id"),
seq: row.get("seq"),
role: row.get("role"),
content: row.get("content"),
reasoning_content: row.get("reasoning_content"),
media_refs: row.get("media_refs"),
tool_call_id: row.get("tool_call_id"),
tool_name: row.get("tool_name"),
tool_calls: row.get("tool_calls"),
source: row.get("source"),
created_at: row.get("created_at"),
})
.collect();
Ok((messages, total))
}
pub async fn clear_messages(&self, session_id: &str) -> Result<(), StorageError> {
sqlx::query(r#"DELETE FROM messages WHERE session_id = ?"#)
.bind(session_id)
.execute(self.pool())
.await?;
Ok(())
}
/// 追加消息,带重试逻辑
/// 重试 3 次100/200/300ms 退避),仍失败返回错误
pub async fn append_message_with_retry(
&self,
session_id: &str,
msg: &crate::storage::message::MessageMeta,
) -> Result<i64, StorageError> {
let delays = [100, 200, 300];
for (i, delay) in delays.iter().enumerate() {
match self.append_message(session_id, msg).await {
Ok(seq) => return Ok(seq),
Err(e) if i < delays.len() - 1 => {
sleep(Duration::from_millis(*delay)).await;
tracing::warn!("Storage write failed, retrying: {}", e);
}
Err(e) => {
tracing::error!("Storage write failed after retries: {}", e);
return Err(e);
}
}
}
unreachable!()
}
// ── Background Task CRUD ──
pub async fn create_background_task(
&self,
task: &crate::storage::background_task::BackgroundTask,
) -> Result<(), StorageError> {
sqlx::query(
r#"
INSERT INTO background_tasks (id, session_id, channel, chat_id, prompt, allowed_tools, status, result, error, tool_calls_count, iterations, started_at, finished_at, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&task.id)
.bind(&task.session_id)
.bind(&task.channel)
.bind(&task.chat_id)
.bind(&task.prompt)
.bind(&task.allowed_tools)
.bind(&task.status)
.bind(&task.result)
.bind(&task.error)
.bind(task.tool_calls_count)
.bind(task.iterations)
.bind(task.started_at)
.bind(task.finished_at)
.bind(task.created_at)
.execute(self.pool())
.await?;
Ok(())
}
pub(crate) async fn update_background_task_status(
&self,
id: &str,
update: crate::storage::background_task::BackgroundTaskUpdate<'_>,
) -> Result<(), StorageError> {
sqlx::query(
r#"
UPDATE background_tasks
SET status = ?, result = COALESCE(?, result), error = COALESCE(?, error),
started_at = COALESCE(?, started_at), finished_at = COALESCE(?, finished_at),
tool_calls_count = COALESCE(?, tool_calls_count),
iterations = COALESCE(?, iterations)
WHERE id = ?
"#,
)
.bind(update.status)
.bind(update.result)
.bind(update.error)
.bind(update.started_at)
.bind(update.finished_at)
.bind(update.tool_calls_count)
.bind(update.iterations)
.bind(id)
.execute(self.pool())
.await?;
Ok(())
}
pub async fn get_background_task(
&self,
id: &str,
) -> Result<crate::storage::background_task::BackgroundTask, StorageError> {
let row = sqlx::query(
r#"
SELECT id, session_id, channel, chat_id, prompt, allowed_tools, status, result, error,
tool_calls_count, iterations, started_at, finished_at, created_at
FROM background_tasks WHERE id = ?
"#,
)
.bind(id)
.fetch_optional(self.pool())
.await?
.ok_or_else(|| StorageError::NotFound(id.to_string()))?;
Ok(crate::storage::background_task::BackgroundTask {
id: row.get("id"),
session_id: row.get("session_id"),
channel: row.get("channel"),
chat_id: row.get("chat_id"),
prompt: row.get("prompt"),
allowed_tools: row.get("allowed_tools"),
status: row.get("status"),
result: row.get("result"),
error: row.get("error"),
tool_calls_count: row.get("tool_calls_count"),
iterations: row.get("iterations"),
started_at: row.get("started_at"),
finished_at: row.get("finished_at"),
created_at: row.get("created_at"),
})
}
pub async fn list_background_tasks(
&self,
session_id: &str,
) -> Result<Vec<crate::storage::background_task::BackgroundTask>, StorageError> {
let rows = sqlx::query(
r#"
SELECT id, session_id, channel, chat_id, prompt, allowed_tools, status, result, error,
tool_calls_count, iterations, started_at, finished_at, created_at
FROM background_tasks
WHERE session_id = ?
ORDER BY created_at DESC
"#,
)
.bind(session_id)
.fetch_all(self.pool())
.await?;
Ok(rows
.into_iter()
.map(|row| crate::storage::background_task::BackgroundTask {
id: row.get("id"),
session_id: row.get("session_id"),
channel: row.get("channel"),
chat_id: row.get("chat_id"),
prompt: row.get("prompt"),
allowed_tools: row.get("allowed_tools"),
status: row.get("status"),
result: row.get("result"),
error: row.get("error"),
tool_calls_count: row.get("tool_calls_count"),
iterations: row.get("iterations"),
started_at: row.get("started_at"),
finished_at: row.get("finished_at"),
created_at: row.get("created_at"),
})
.collect())
}
/// List recent background tasks across sessions for the management UI.
pub async fn list_recent_background_tasks(
&self,
limit: usize,
) -> Result<Vec<crate::storage::background_task::BackgroundTask>, StorageError> {
let rows = sqlx::query(
r#"
SELECT id, session_id, channel, chat_id, prompt, allowed_tools, status, result, error,
tool_calls_count, iterations, started_at, finished_at, created_at
FROM background_tasks
ORDER BY created_at DESC
LIMIT ?
"#,
)
.bind(limit as i64)
.fetch_all(self.pool())
.await?;
Ok(rows
.into_iter()
.map(|row| crate::storage::background_task::BackgroundTask {
id: row.get("id"),
session_id: row.get("session_id"),
channel: row.get("channel"),
chat_id: row.get("chat_id"),
prompt: row.get("prompt"),
allowed_tools: row.get("allowed_tools"),
status: row.get("status"),
result: row.get("result"),
error: row.get("error"),
tool_calls_count: row.get("tool_calls_count"),
iterations: row.get("iterations"),
started_at: row.get("started_at"),
finished_at: row.get("finished_at"),
created_at: row.get("created_at"),
})
.collect())
}
pub async fn cleanup_old_tasks(&self, ttl_ms: i64) -> Result<usize, StorageError> {
let cutoff = chrono::Utc::now().timestamp_millis() - ttl_ms;
let result = sqlx::query(
"DELETE FROM background_tasks WHERE status IN ('completed', 'failed', 'cancelled') AND finished_at IS NOT NULL AND finished_at < ?",
)
.bind(cutoff)
.execute(self.pool())
.await?;
Ok(result.rows_affected() as usize)
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
async fn create_test_storage() -> (Storage, TempDir) {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("test.db");
let storage = Storage::new(&db_path).await.unwrap();
(storage, dir)
}
#[tokio::test]
async fn sqlite_runtime_guards_are_enabled() {
let (storage, _dir) = create_test_storage().await;
let journal_mode: String = sqlx::query_scalar("PRAGMA journal_mode")
.fetch_one(storage.pool())
.await
.unwrap();
let foreign_keys: i64 = sqlx::query_scalar("PRAGMA foreign_keys")
.fetch_one(storage.pool())
.await
.unwrap();
let busy_timeout: i64 = sqlx::query_scalar("PRAGMA busy_timeout")
.fetch_one(storage.pool())
.await
.unwrap();
let schema_version: i64 = sqlx::query_scalar("PRAGMA user_version")
.fetch_one(storage.pool())
.await
.unwrap();
assert_eq!(journal_mode, "wal");
assert_eq!(foreign_keys, 1);
assert_eq!(busy_timeout, 5000);
assert_eq!(schema_version, SCHEMA_VERSION);
let orphan = sqlx::query(
r#"
INSERT INTO messages (id, session_id, seq, role, content, created_at)
VALUES ('orphan', 'missing', 1, 'user', 'no parent', 1)
"#,
)
.execute(storage.pool())
.await;
assert!(orphan.is_err());
}
#[tokio::test]
async fn reopening_database_does_not_rebuild_existing_fts_index() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("fts.db");
let storage = Storage::new(&db_path).await.unwrap();
sqlx::query(
"INSERT INTO memory_fts(rowid, key, content) VALUES (999999, 'startup_sentinel', 'startup_sentinel')",
)
.execute(storage.pool())
.await
.unwrap();
drop(storage);
let reopened = Storage::new(&db_path).await.unwrap();
let sentinel_count: i64 = sqlx::query_scalar(
"SELECT count(*) FROM memory_fts WHERE memory_fts MATCH 'startup_sentinel'",
)
.fetch_one(reopened.pool())
.await
.unwrap();
assert_eq!(sentinel_count, 1);
}
#[tokio::test]
async fn background_task_completion_persists_execution_metrics() {
let (storage, _dir) = create_test_storage().await;
let task = crate::storage::BackgroundTask {
id: "task-metrics".into(),
session_id: "cli:test:dialog".into(),
channel: "cli".into(),
chat_id: "test".into(),
prompt: "measure".into(),
allowed_tools: None,
status: "pending".into(),
result: None,
error: None,
tool_calls_count: 0,
iterations: 0,
started_at: None,
finished_at: None,
created_at: 1,
};
storage.create_background_task(&task).await.unwrap();
storage
.update_background_task_status(
&task.id,
crate::storage::background_task::BackgroundTaskUpdate {
status: "completed",
result: Some("done"),
error: None,
started_at: Some(2),
finished_at: Some(3),
tool_calls_count: Some(4),
iterations: Some(5),
},
)
.await
.unwrap();
let persisted = storage.get_background_task(&task.id).await.unwrap();
assert_eq!(persisted.tool_calls_count, 4);
assert_eq!(persisted.iterations, 5);
}
#[tokio::test]
async fn webui_lists_recent_tasks_across_sessions() {
let (storage, _dir) = create_test_storage().await;
for (id, session_id, created_at) in [("old", "cli:a:d1", 1), ("new", "cli:b:d2", 2)] {
storage
.create_background_task(&crate::storage::BackgroundTask {
id: id.into(),
session_id: session_id.into(),
channel: "cli".into(),
chat_id: "chat".into(),
prompt: id.into(),
allowed_tools: None,
status: "pending".into(),
result: None,
error: None,
tool_calls_count: 0,
iterations: 0,
started_at: None,
finished_at: None,
created_at,
})
.await
.unwrap();
}
let tasks = storage.list_recent_background_tasks(10).await.unwrap();
assert_eq!(
tasks
.iter()
.map(|task| task.id.as_str())
.collect::<Vec<_>>(),
vec!["new", "old"]
);
assert_eq!(tasks[0].session_id, "cli:b:d2");
}
#[tokio::test]
async fn webui_lists_and_filters_memories_without_search_text() {
let (storage, _dir) = create_test_storage().await;
for (key, category, updated_at) in [
(
"fact",
crate::memory::MemoryCategory::Knowledge,
"2026-01-01T00:00:00Z",
),
(
"summary",
crate::memory::MemoryCategory::Timeline,
"2026-01-02T00:00:00Z",
),
] {
storage
.upsert_memory(&crate::memory::MemoryEntry {
id: key.into(),
key: key.into(),
content: format!("content {key}"),
category,
importance: 0.5,
session_id: Some("cli:test:dialog".into()),
created_at: updated_at.into(),
updated_at: updated_at.into(),
})
.await
.unwrap();
}
let all = storage.list_memories(None, None, 10).await.unwrap();
assert_eq!(
all.iter()
.map(|entry| entry.key.as_str())
.collect::<Vec<_>>(),
vec!["summary", "fact"]
);
let knowledge = storage
.list_memories(Some(&crate::memory::MemoryCategory::Knowledge), None, 10)
.await
.unwrap();
assert_eq!(knowledge.len(), 1);
assert_eq!(knowledge[0].key, "fact");
}
#[tokio::test]
async fn legacy_schema_is_migrated_without_rebuild() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("legacy.db");
let pool = SqlitePoolOptions::new()
.connect_with(
SqliteConnectOptions::new()
.filename(&db_path)
.create_if_missing(true),
)
.await
.unwrap();
sqlx::query(
r#"
CREATE TABLE sessions (
id TEXT PRIMARY KEY, channel TEXT NOT NULL, chat_id TEXT NOT NULL,
dialog_id TEXT NOT NULL, title TEXT NOT NULL DEFAULT 'new',
created_at INTEGER NOT NULL, last_active_at INTEGER NOT NULL,
message_count INTEGER DEFAULT 0, routing_info TEXT, deleted_at INTEGER,
UNIQUE(channel, chat_id, dialog_id)
)
"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"
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 TEXT,
tool_call_id TEXT, tool_name TEXT, tool_calls TEXT,
created_at INTEGER NOT NULL
)
"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"
CREATE TABLE scheduled_jobs (
id TEXT PRIMARY KEY, name TEXT NOT NULL, schedule TEXT NOT NULL,
prompt TEXT NOT NULL, channel TEXT NOT NULL, chat_id TEXT NOT NULL,
model TEXT, enabled INTEGER NOT NULL DEFAULT 1,
delete_after_run INTEGER NOT NULL DEFAULT 0, next_run_at INTEGER NOT NULL,
last_run_at INTEGER, last_status TEXT, last_error TEXT,
created_at INTEGER NOT NULL, updated_at INTEGER NOT NULL
)
"#,
)
.execute(&pool)
.await
.unwrap();
drop(pool);
let storage = Storage::new(&db_path).await.unwrap();
for (table, expected) in [
("messages", vec!["source", "reasoning_content"]),
(
"sessions",
vec![
"archived_at",
"last_consolidated_at",
"last_compressed_message_at",
],
),
(
"scheduled_jobs",
vec!["locked_at", "lock_owner", "lease_until"],
),
] {
let columns = sqlx::query(&format!("PRAGMA table_info({table})"))
.fetch_all(storage.pool())
.await
.unwrap();
for column in expected {
assert!(
columns
.iter()
.any(|row| row.get::<String, _>("name") == column),
"missing migrated column {table}.{column}"
);
}
}
let schema_version: i64 = sqlx::query_scalar("PRAGMA user_version")
.fetch_one(storage.pool())
.await
.unwrap();
assert_eq!(schema_version, SCHEMA_VERSION);
for table in ["task_plans", "task_items"] {
let exists: i64 = sqlx::query_scalar(
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = ?",
)
.bind(table)
.fetch_one(storage.pool())
.await
.unwrap();
assert_eq!(exists, 1, "missing migrated table {table}");
}
}
#[tokio::test]
async fn test_upsert_and_get_session() {
let (storage, _dir) = create_test_storage().await;
let meta = crate::storage::session::SessionMeta {
id: "cli_chat:sid123:dialog1".to_string(),
channel: "cli_chat".to_string(),
chat_id: "sid123".to_string(),
dialog_id: "dialog1".to_string(),
title: "测试会话".to_string(),
created_at: 1000,
last_active_at: 1000,
message_count: 0,
routing_info: Some(r#"{"type":"cli"}"#.to_string()),
archived_at: None,
deleted_at: None,
last_consolidated_at: None,
last_compressed_message_at: None,
};
storage.upsert_session(&meta).await.unwrap();
let loaded = storage.get_session(&meta.id).await.unwrap();
assert_eq!(loaded.title, "测试会话");
assert_eq!(loaded.channel, "cli_chat");
}
#[tokio::test]
async fn test_get_nonexistent_session() {
let (storage, _dir) = create_test_storage().await;
let result = storage.get_session("nonexistent").await;
assert!(result.is_err());
matches!(result.unwrap_err(), StorageError::NotFound(_));
}
#[tokio::test]
async fn test_list_sessions() {
let (storage, _dir) = create_test_storage().await;
for i in 0..5 {
let meta = crate::storage::session::SessionMeta {
id: format!("cli_chat:sid123:dialog{}", i),
channel: "cli_chat".to_string(),
chat_id: "sid123".to_string(),
dialog_id: format!("dialog{}", i),
title: format!("会话{}", i),
created_at: i * 1000,
last_active_at: i * 1000,
message_count: i,
routing_info: None,
archived_at: None,
deleted_at: None,
last_consolidated_at: None,
last_compressed_message_at: None,
};
storage.upsert_session(&meta).await.unwrap();
}
let sessions = storage
.list_sessions("cli_chat", "sid123", 10, false)
.await
.unwrap();
assert_eq!(sessions.len(), 5);
// 按 last_active_at DESC 排序
assert_eq!(sessions[0].dialog_id, "dialog4");
}
#[tokio::test]
async fn test_soft_delete() {
let (storage, _dir) = create_test_storage().await;
let meta = crate::storage::session::SessionMeta {
id: "cli_chat:sid123:dialog1".to_string(),
channel: "cli_chat".to_string(),
chat_id: "sid123".to_string(),
dialog_id: "dialog1".to_string(),
title: "测试".to_string(),
created_at: 1000,
last_active_at: 1000,
message_count: 0,
routing_info: None,
archived_at: None,
deleted_at: None,
last_consolidated_at: None,
last_compressed_message_at: None,
};
storage.upsert_session(&meta).await.unwrap();
storage.soft_delete_session(&meta.id).await.unwrap();
let result = storage.get_session(&meta.id).await;
assert!(result.is_err());
matches!(result.unwrap_err(), StorageError::NotFound(_));
}
#[tokio::test]
async fn test_append_and_load_messages() {
let (storage, _dir) = create_test_storage().await;
let session_meta = crate::storage::session::SessionMeta {
id: "cli_chat:sid123:dialog1".to_string(),
channel: "cli_chat".to_string(),
chat_id: "sid123".to_string(),
dialog_id: "dialog1".to_string(),
title: "测试".to_string(),
created_at: 1000,
last_active_at: 1000,
message_count: 0,
routing_info: None,
archived_at: None,
deleted_at: None,
last_consolidated_at: None,
last_compressed_message_at: None,
};
storage.upsert_session(&session_meta).await.unwrap();
let msg = crate::storage::message::MessageMeta {
id: "msg1".to_string(),
session_id: session_meta.id.clone(),
seq: 1,
role: "user".to_string(),
content: "你好".to_string(),
reasoning_content: None,
media_refs: None,
tool_call_id: None,
tool_name: None,
tool_calls: None,
source: None,
created_at: 1000,
};
let seq = storage
.append_message(&session_meta.id, &msg)
.await
.unwrap();
assert_eq!(seq, 1);
let loaded = storage.load_messages(&session_meta.id, 0).await.unwrap();
assert_eq!(loaded.len(), 1);
assert_eq!(loaded[0].content, "你好");
for seq in 2..=5 {
let mut message = msg.clone();
message.id = format!("msg{seq}");
message.seq = seq;
message.content = format!("message {seq}");
storage
.append_message(&session_meta.id, &message)
.await
.unwrap();
}
let recent = storage
.load_recent_session_messages(&session_meta.id, 2)
.await
.unwrap();
assert_eq!(recent.len(), 2);
assert_eq!(recent[0].seq, 4);
assert_eq!(recent[1].seq, 5);
}
#[tokio::test]
async fn test_persist_message_batch_is_atomic() {
let (storage, _dir) = create_test_storage().await;
let session_meta = crate::storage::session::SessionMeta {
id: "cli_chat:atomic:dialog1".to_string(),
channel: "cli_chat".to_string(),
chat_id: "atomic".to_string(),
dialog_id: "dialog1".to_string(),
title: "Atomic turn".to_string(),
created_at: 1000,
last_active_at: 2000,
message_count: 1,
routing_info: None,
archived_at: None,
deleted_at: None,
last_consolidated_at: None,
last_compressed_message_at: None,
};
storage.upsert_session(&session_meta).await.unwrap();
let message = crate::storage::message::MessageMeta {
id: "duplicate-id".to_string(),
session_id: session_meta.id.clone(),
seq: 1,
role: "assistant".to_string(),
content: "must roll back".to_string(),
reasoning_content: None,
media_refs: None,
tool_call_id: None,
tool_name: None,
tool_calls: None,
source: None,
created_at: 2000,
};
let result = storage
.persist_message_batch(&session_meta.id, &[message.clone(), message], &session_meta)
.await;
assert!(result.is_err());
assert!(
storage
.load_messages(&session_meta.id, 0)
.await
.unwrap()
.is_empty()
);
}
#[tokio::test]
async fn test_touch_session() {
let (storage, _dir) = create_test_storage().await;
let meta = crate::storage::session::SessionMeta {
id: "cli_chat:sid123:dialog1".to_string(),
channel: "cli_chat".to_string(),
chat_id: "sid123".to_string(),
dialog_id: "dialog1".to_string(),
title: "测试".to_string(),
created_at: 1000,
last_active_at: 1000,
message_count: 0,
routing_info: None,
archived_at: None,
deleted_at: None,
last_consolidated_at: None,
last_compressed_message_at: None,
};
storage.upsert_session(&meta).await.unwrap();
storage.touch_session(&meta.id, 5, 2000).await.unwrap();
let loaded = storage.get_session(&meta.id).await.unwrap();
assert_eq!(loaded.message_count, 5);
assert_eq!(loaded.last_active_at, 2000);
}
}