1900 lines
65 KiB
Rust
1900 lines
65 KiB
Rust
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);
|
||
}
|
||
}
|