use std::collections::HashMap;
use std::fs;
use std::io;
use std::path::PathBuf;
use choreo_proto::{ContextConfig, ReasoningProducer, Turn};
use redb::ReadableDatabase;
use redb::ReadableTable;
use redb::TableDefinition;
use serde::{Deserialize, Serialize};
use tracing::{debug, error, info, warn};
const SESSIONS: TableDefinition<u64, &[u8]> = TableDefinition::new("sessions");
const SESSION_TURNS: TableDefinition<(u64, u32), &[u8]> = TableDefinition::new("session_turns");
const CREDENTIALS: TableDefinition<&str, &[u8]> = TableDefinition::new("credentials");
const META: TableDefinition<&str, u64> = TableDefinition::new("meta");
const SESSION_KV: TableDefinition<(u64, String), Vec<u8>> = TableDefinition::new("session_kv");
const DELETED_SESSIONS: TableDefinition<u64, ()> = TableDefinition::new("deleted_sessions");
type KvRangeIter<'a> = Box<
dyn Iterator<
Item = Result<
(
redb::AccessGuard<'a, (u64, String)>,
redb::AccessGuard<'a, Vec<u8>>,
),
redb::StorageError,
>,
> + 'a,
>;
fn db_err(msg: String) -> io::Error {
io::Error::other(msg)
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionRecord {
pub title: Option<String>,
pub selected_model: Option<String>,
pub parent_session_id: Option<u64>,
pub working_dir: Option<String>,
pub turn_count: u32,
pub created_at: i64,
pub last_modified: i64,
pub active_tool_groups: Vec<String>,
#[serde(default)]
pub context_config: ContextConfig,
#[serde(default)]
pub account_name: Option<String>,
#[serde(default)]
pub reasoning_effort: Option<String>,
#[serde(default)]
pub last_response_id: Option<String>,
#[serde(default)]
pub last_response_id_producer: Option<ReasoningProducer>,
}
pub fn db_path() -> io::Result<PathBuf> {
if let Ok(override_path) = std::env::var("CHOREOGRAPHR_DB_PATH") {
return Ok(PathBuf::from(override_path));
}
let data_dir = dirs::data_dir().ok_or_else(|| {
io::Error::new(
io::ErrorKind::NotFound,
"could not determine data directory",
)
})?;
Ok(data_dir.join("choreographr").join("state.redb"))
}
pub const SCHEMA_VERSION: u64 = 1;
pub const INITIAL_SCHEMA_VERSION: u64 = 1;
const SCHEMA_VERSION_KEY: &str = "schema_version";
struct Migration {
from: u64,
run: fn(&redb::Database) -> io::Result<()>,
}
const MIGRATIONS: &[Migration] = &[];
fn current_schema_version(db: &redb::Database) -> io::Result<u64> {
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = match read_txn.open_table(META) {
Ok(table) => table,
Err(redb::TableError::TableDoesNotExist(_)) => return Ok(0),
Err(e) => return Err(db_err(format!("redb open meta: {e}"))),
};
Ok(table
.get(SCHEMA_VERSION_KEY)
.map_err(|e| db_err(format!("redb get meta: {e}")))?
.map(|guard| guard.value())
.unwrap_or(0))
}
fn stamp_schema_version(db: &redb::Database, version: u64) -> io::Result<()> {
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut table = write_txn
.open_table(META)
.map_err(|e| db_err(format!("redb open meta: {e}")))?;
table
.insert(SCHEMA_VERSION_KEY, version)
.map_err(|e| db_err(format!("redb set schema_version: {e}")))?;
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit schema_version: {e}")))?;
info!(version, "stamped database schema version");
Ok(())
}
fn backup_db_file(path: &std::path::Path, from: u64) -> io::Result<()> {
let file_name = path
.file_name()
.map(|name| name.to_string_lossy().into_owned())
.unwrap_or_else(|| "state.redb".to_string());
let backup_path = path.with_file_name(format!("{file_name}.bak-v{from}"));
fs::copy(path, &backup_path)?;
info!(
from = %path.display(),
to = %backup_path.display(),
"backed up database before applying migrations"
);
Ok(())
}
pub fn run_migrations(db: &redb::Database) -> io::Result<()> {
run_migrations_to(db, SCHEMA_VERSION, MIGRATIONS)
}
fn run_migrations_to(db: &redb::Database, target: u64, migrations: &[Migration]) -> io::Result<()> {
let current = current_schema_version(db)?;
if current > target {
error!(
current,
supported = target,
"refusing to open database: schema version newer than this binary supports"
);
return Err(db_err(format!(
"database schema version {current} is newer than this binary supports ({target}); \
upgrade choreographr before continuing"
)));
}
if current == 0 && target > 1 {
let msg =
"database has no schema version (pre-release data); recreate it or restore a backup";
error!("{msg}");
return Err(db_err(msg.to_string()));
}
if current == target {
return Ok(()); }
let expected: Vec<u64> = (1..target).collect();
let provided: Vec<u64> = migrations.iter().map(|m| m.from).collect();
if provided != expected {
let msg =
format!("migration chain is not contiguous: has {provided:?}, needs {expected:?}");
error!("{msg}");
return Err(db_err(msg));
}
if current == 0 {
warn!(
"database was unversioned; stamping schema version {target} \
(pre-release blobs, if any, are not migrated)"
);
}
if !migrations.is_empty() {
let path = db_path()?;
backup_db_file(&path, current)?; }
for migration in migrations.iter().filter(|m| m.from >= current) {
info!(
from = migration.from,
to = migration.from + 1,
"applying database migration"
);
(migration.run)(db)?;
}
stamp_schema_version(db, target)
}
fn initialize_schema_version(db: &redb::Database) -> io::Result<()> {
stamp_schema_version(db, INITIAL_SCHEMA_VERSION)
.map_err(|e| io::Error::other(format!("failed to initialize schema version: {e}")))
}
pub fn open_db() -> io::Result<redb::Database> {
let path = db_path()?;
info!(path = %path.display(), "opening database");
if let Some(parent) = path.parent() {
fs::create_dir_all(parent)?;
}
if let Ok(metadata) = fs::metadata(&path)
&& metadata.len() == 0
{
warn!("database file exists but is empty (interrupted create?); recreating");
let db = redb::Database::create(&path)
.map_err(|e| io::Error::other(format!("failed to create database: {e}")))?;
initialize_schema_version(&db)?;
return Ok(db);
}
match redb::Database::open(&path) {
Ok(db) => Ok(db),
Err(redb::DatabaseError::Storage(redb::StorageError::Io(io_err)))
if io_err.kind() == io::ErrorKind::NotFound =>
{
info!("database file not found, creating new database");
let db = redb::Database::create(&path)
.map_err(|e| io::Error::other(format!("failed to create database: {e}")))?;
initialize_schema_version(&db)?;
Ok(db)
}
Err(redb::DatabaseError::UpgradeRequired(actual)) => Err(io::Error::other(format!(
"database file format version {actual} is not supported by this binary; \
restore a backup (state.redb.bak-v*) or use the documented dump/restore path"
))),
Err(e) => Err(io::Error::other(format!(
"failed to open database (refusing to recreate a potentially corrupt file): {e}"
))),
}
}
pub fn write_session(
db: &redb::Database,
session_id: u64,
record: &SessionRecord,
) -> io::Result<()> {
let payload = rmp_serde::to_vec_named(record)
.map_err(|e| db_err(format!("codec encode session: {e}")))?;
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut table = write_txn
.open_table(SESSIONS)
.map_err(|e| db_err(format!("redb open sessions: {e}")))?;
table
.insert(session_id, payload.as_slice())
.map_err(|e| db_err(format!("redb insert session: {e}")))?;
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit session: {e}")))?;
debug!("write_session: id={} ok", session_id);
Ok(())
}
pub fn read_session(db: &redb::Database, session_id: u64) -> io::Result<Option<SessionRecord>> {
debug!("read_session: id={}", session_id);
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = read_txn
.open_table(SESSIONS)
.map_err(|e| db_err(format!("redb open sessions: {e}")))?;
match table
.get(session_id)
.map_err(|e| db_err(format!("redb get session: {e}")))?
{
Some(guard) => match rmp_serde::from_slice::<SessionRecord>(guard.value()) {
Ok(record) => Ok(Some(record)),
Err(e) => {
warn!(
session_id,
error = %e,
"undecodable session record, treating as absent"
);
Ok(None)
}
},
None => Ok(None),
}
}
pub fn read_all_sessions(db: &redb::Database) -> io::Result<Vec<(u64, SessionRecord)>> {
debug!("read_all_sessions");
let read_txn = db.begin_read().map_err(|e| {
let msg = format!("redb read txn: {e}");
error!("read_all_sessions: {msg}");
db_err(msg)
})?;
let table = match read_txn.open_table(SESSIONS) {
Ok(t) => t,
Err(e) => {
warn!("read_all_sessions: table 'sessions' not found (first run?): {e}");
return Ok(Vec::new());
}
};
let mut sessions: Vec<(u64, SessionRecord)> = Vec::new();
let iter = match table.iter() {
Ok(it) => it,
Err(e) => {
let msg = format!("redb iter sessions: {e}");
error!("read_all_sessions: {msg}");
return Err(db_err(msg));
}
};
for result in iter {
let (key, value) = match result {
Ok(kv) => kv,
Err(e) => {
warn!("read_all_sessions: skipping bad entry: {e}");
continue;
}
};
match rmp_serde::from_slice::<SessionRecord>(value.value()) {
Ok(record) => {
sessions.push((key.value(), record));
}
Err(e) => {
warn!(
"read_all_sessions: skipping session {} (decode failed: {e})",
key.value()
);
continue;
}
}
}
debug!("read_all_sessions: {} records", sessions.len());
sessions.sort_by_key(|(id, _)| *id);
Ok(sessions)
}
fn session_range_end(session_id: u64) -> u64 {
session_id.saturating_add(1)
}
pub fn delete_session(db: &redb::Database, session_id: u64) -> io::Result<()> {
debug!("delete_session: id={}", session_id);
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut sessions = write_txn
.open_table(SESSIONS)
.map_err(|e| db_err(format!("redb open sessions: {e}")))?;
sessions
.remove(session_id)
.map_err(|e| db_err(format!("redb remove session: {e}")))?;
}
{
let mut turns = write_txn
.open_table(SESSION_TURNS)
.map_err(|e| db_err(format!("redb open turns: {e}")))?;
let keys_to_remove: Vec<(u64, u32)> = turns
.range::<(u64, u32)>((session_id, 0u32)..(session_range_end(session_id), 0u32))
.map_err(|e| db_err(format!("redb range turns: {e}")))?
.filter_map(|result| result.ok())
.map(|(key, _)| key.value())
.collect();
for key in keys_to_remove {
turns
.remove(key)
.map_err(|e| db_err(format!("redb remove turn: {e}")))?;
}
}
{
let mut kv_table = write_txn
.open_table(SESSION_KV)
.map_err(|e| db_err(format!("redb open session_kv: {e}")))?;
let kv_keys: Vec<(u64, String)> = kv_table
.range::<(u64, String)>(
(session_id, String::new())..(session_range_end(session_id), String::new()),
)
.map_err(|e| db_err(format!("redb range session_kv: {e}")))?
.filter_map(|result| result.ok())
.map(|(k, _)| k.value())
.collect();
for key in kv_keys {
kv_table
.remove(key)
.map_err(|e| db_err(format!("redb remove session_kv: {e}")))?;
}
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit delete: {e}")))?;
Ok(())
}
pub fn mark_session_deleted(db: &redb::Database, session_id: u64) -> io::Result<()> {
debug!("mark_session_deleted: id={}", session_id);
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut table = write_txn
.open_table(DELETED_SESSIONS)
.map_err(|e| db_err(format!("redb open deleted_sessions: {e}")))?;
table
.insert(session_id, ())
.map_err(|e| db_err(format!("redb insert tombstone: {e}")))?;
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit tombstone: {e}")))?;
Ok(())
}
pub fn clear_session_tombstone(db: &redb::Database, session_id: u64) -> io::Result<()> {
debug!("clear_session_tombstone: id={}", session_id);
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut table = write_txn
.open_table(DELETED_SESSIONS)
.map_err(|e| db_err(format!("redb open deleted_sessions: {e}")))?;
table
.remove(session_id)
.map_err(|e| db_err(format!("redb remove tombstone: {e}")))?;
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit tombstone: {e}")))?;
Ok(())
}
pub fn purge_tombstoned_sessions(db: &redb::Database) -> io::Result<usize> {
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = match read_txn.open_table(DELETED_SESSIONS) {
Ok(table) => table,
Err(redb::TableError::TableDoesNotExist(_)) => return Ok(0),
Err(e) => return Err(db_err(format!("redb open deleted_sessions: {e}"))),
};
let ids: Vec<u64> = table
.iter()
.map_err(|e| db_err(format!("redb iter deleted_sessions: {e}")))?
.filter_map(|result| result.ok())
.map(|(key, _)| key.value())
.collect();
drop(read_txn);
let mut purged = 0usize;
for id in ids {
if let Err(e) = delete_session(db, id) {
warn!(session_id = id, error = %e, "purge: failed to delete tombstoned session");
continue;
}
if let Err(e) = clear_session_tombstone(db, id) {
warn!(session_id = id, error = %e, "purge: failed to clear tombstone");
}
purged += 1;
info!(
session_id = id,
"purged session record left behind by a deleted-session shutdown"
);
}
Ok(purged)
}
pub fn write_turn(
db: &redb::Database,
session_id: u64,
turn_id: u32,
turn: &Turn,
) -> io::Result<()> {
let payload =
rmp_serde::to_vec_named(turn).map_err(|e| db_err(format!("codec encode turn: {e}")))?;
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut table = write_txn
.open_table(SESSION_TURNS)
.map_err(|e| db_err(format!("redb open turns: {e}")))?;
table
.insert((session_id, turn_id), payload.as_slice())
.map_err(|e| db_err(format!("redb insert turn: {e}")))?;
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit turn: {e}")))?;
Ok(())
}
pub fn read_turns(db: &redb::Database, session_id: u64) -> io::Result<Vec<(u32, Turn)>> {
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = read_txn
.open_table(SESSION_TURNS)
.map_err(|e| db_err(format!("redb open turns: {e}")))?;
let mut turns: Vec<(u32, Turn)> = Vec::new();
for result in table
.iter()
.map_err(|e| db_err(format!("redb iter turns: {e}")))?
{
let (key, value) = result.map_err(|e| db_err(format!("redb iter item: {e}")))?;
let (sid, idx) = key.value();
if sid == session_id {
match rmp_serde::from_slice::<Turn>(value.value()) {
Ok(turn) => turns.push((idx, turn)),
Err(e) => {
tracing::warn!(session_id, turn_id = idx, error = %e, "undecodable turn, skipping");
}
}
}
}
turns.sort_by_key(|(idx, _)| *idx);
Ok(turns)
}
pub fn write_turn_retry(
db: &redb::Database,
session_id: u64,
turn_id: u32,
turn: &Turn,
) -> io::Result<()> {
let mut attempts = 0;
loop {
match write_turn(db, session_id, turn_id, turn) {
Ok(()) => return Ok(()),
Err(_e) if attempts < 3 => {
attempts += 1;
std::thread::sleep(std::time::Duration::from_millis(1));
continue;
}
Err(e) => return Err(e),
}
}
}
pub fn delete_session_turns(db: &redb::Database, session_id: u64) -> io::Result<()> {
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut table = write_txn
.open_table(SESSION_TURNS)
.map_err(|e| db_err(format!("redb open turns: {e}")))?;
let keys_to_remove: Vec<(u64, u32)> = table
.iter()
.map_err(|e| db_err(format!("redb iter turns: {e}")))?
.filter_map(|result| match result {
Ok((key, _)) => {
if key.value().0 == session_id {
Some(key.value())
} else {
None
}
}
Err(e) => {
warn!("undecodable turn entry in session {session_id}: {e}");
None
}
})
.collect();
for key in keys_to_remove {
table
.remove(key)
.map_err(|e| db_err(format!("redb remove turn: {e}")))?;
}
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit delete turns: {e}")))?;
Ok(())
}
pub fn delete_session_turns_retry(db: &redb::Database, session_id: u64) -> io::Result<()> {
let mut attempts = 0;
loop {
match delete_session_turns(db, session_id) {
Ok(()) => return Ok(()),
Err(_e) if attempts < 3 => {
attempts += 1;
std::thread::sleep(std::time::Duration::from_millis(1));
continue;
}
Err(e) => return Err(e),
}
}
}
pub fn set_credential_blob(
db: &redb::Database,
service: &str,
blob: &[u8],
) -> Result<(), redb::Error> {
let write_txn = db.begin_write()?;
{
let mut table = write_txn.open_table(CREDENTIALS)?;
table.insert(service, blob)?;
}
write_txn.commit()?;
Ok(())
}
pub fn get_all_credential_blobs(
db: &redb::Database,
) -> Result<HashMap<String, Vec<u8>>, redb::Error> {
let read_txn = db.begin_read()?;
let table = match read_txn.open_table(CREDENTIALS) {
Ok(table) => table,
Err(redb::TableError::TableDoesNotExist(_)) => return Ok(HashMap::new()),
Err(e) => return Err(e.into()),
};
let mut map = HashMap::new();
for result in table.iter()? {
let (key, value) = result?;
map.insert(key.value().to_string(), value.value().to_vec());
}
Ok(map)
}
pub fn remove_credential_blob(db: &redb::Database, service: &str) -> Result<(), redb::Error> {
let write_txn = db.begin_write()?;
{
let mut table = write_txn.open_table(CREDENTIALS)?;
table.remove(service)?;
}
write_txn.commit()?;
Ok(())
}
pub fn kv_set(db: &redb::Database, session_id: u64, key: &str, value: &[u8]) -> io::Result<()> {
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut table = write_txn
.open_table(SESSION_KV)
.map_err(|e| db_err(format!("redb open session_kv: {e}")))?;
table
.insert((session_id, key.to_string()), value.to_vec())
.map_err(|e| db_err(format!("redb kv_set: {e}")))?;
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit kv_set: {e}")))?;
debug!("kv_set: session={} key=\"{}\" ok", session_id, key);
Ok(())
}
pub fn kv_get(db: &redb::Database, session_id: u64, key: &str) -> io::Result<Option<Vec<u8>>> {
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = read_txn
.open_table(SESSION_KV)
.map_err(|e| db_err(format!("redb open session_kv: {e}")))?;
match table
.get((session_id, key.to_string()))
.map_err(|e| db_err(format!("redb kv_get: {e}")))?
{
Some(guard) => Ok(Some(guard.value().to_vec())),
None => Ok(None),
}
}
pub fn kv_delete(db: &redb::Database, session_id: u64, key: &str) -> io::Result<bool> {
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
let removed = {
let mut table = write_txn
.open_table(SESSION_KV)
.map_err(|e| db_err(format!("redb open session_kv: {e}")))?;
table
.remove((session_id, key.to_string()))
.map_err(|e| db_err(format!("redb kv_delete: {e}")))?
.is_some()
};
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit kv_delete: {e}")))?;
debug!(
"kv_delete: session={} key=\"{}\" found={}",
session_id, key, removed
);
Ok(removed)
}
pub fn kv_delete_range(
db: &redb::Database,
session_id: u64,
start: &str,
end: Option<&str>,
) -> io::Result<u64> {
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
let count = {
let mut table = write_txn
.open_table(SESSION_KV)
.map_err(|e| db_err(format!("redb open session_kv: {e}")))?;
let range = match end {
Some(end) => {
let range_start = (session_id, start.to_string());
let range_end = (session_id, end.to_string());
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_delete_range: {e}")))?
}
None => {
let range_start = (session_id, start.to_string());
let range_end = (session_range_end(session_id), String::new());
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_delete_range: {e}")))?
}
};
let keys: Vec<(u64, String)> = range
.filter_map(|r| r.ok())
.map(|(k, _)| k.value())
.collect();
let count = keys.len() as u64;
for key in keys {
table
.remove(key)
.map_err(|e| db_err(format!("redb kv_delete_range remove: {e}")))?;
}
count
};
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit kv_delete_range: {e}")))?;
debug!(
"kv_delete_range: session={} start=\"{}\" end={:?} removed={}",
session_id, start, end, count
);
Ok(count)
}
pub fn kv_get_range(
db: &redb::Database,
session_id: u64,
start: &str,
end: Option<&str>,
) -> io::Result<Vec<(String, Vec<u8>)>> {
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = read_txn
.open_table(SESSION_KV)
.map_err(|e| db_err(format!("redb open session_kv: {e}")))?;
let range = match end {
Some(end) => {
let range_start = (session_id, start.to_string());
let range_end = (session_id, end.to_string());
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_get_range: {e}")))?
}
None => {
let range_start = (session_id, start.to_string());
let range_end = (session_range_end(session_id), String::new());
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_get_range: {e}")))?
}
};
let mut results = Vec::new();
for result in range {
let (key, value) = result.map_err(|e| db_err(format!("redb iter kv_get_range: {e}")))?;
results.push((key.value().1, value.value().to_vec()));
}
Ok(results)
}
pub fn kv_list(
db: &redb::Database,
session_id: u64,
start: Option<&str>,
end: Option<&str>,
) -> io::Result<Vec<String>> {
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = read_txn
.open_table(SESSION_KV)
.map_err(|e| db_err(format!("redb open session_kv: {e}")))?;
let range: KvRangeIter<'_> = match (start, end) {
(Some(start), Some(end)) => {
let range_start = (session_id, start.to_string());
let range_end = (session_id, end.to_string());
Box::new(
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_list: {e}")))?,
)
}
(Some(start), None) => {
let range_start = (session_id, start.to_string());
let range_end = (session_range_end(session_id), String::new());
Box::new(
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_list: {e}")))?,
)
}
(None, Some(end)) => {
let range_start = (session_id, String::new());
let range_end = (session_id, end.to_string());
Box::new(
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_list: {e}")))?,
)
}
(None, None) => {
let range_start = (session_id, String::new());
let range_end = (session_range_end(session_id), String::new());
Box::new(
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_list: {e}")))?,
)
}
};
let mut keys = Vec::new();
for result in range {
let (key, _) = result.map_err(|e| db_err(format!("redb iter kv_list: {e}")))?;
keys.push(key.value().1);
}
Ok(keys)
}
pub fn kv_count(db: &redb::Database, session_id: u64, prefix: Option<&str>) -> io::Result<u64> {
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = read_txn
.open_table(SESSION_KV)
.map_err(|e| db_err(format!("redb open session_kv: {e}")))?;
let range = match prefix {
Some(prefix) => {
let range_start = (session_id, prefix.to_string());
let mut end_bytes = prefix.as_bytes().to_vec();
end_bytes.push(0xFF);
let range_end_str = String::from_utf8_lossy(&end_bytes).into_owned();
let range_end = (session_id, range_end_str);
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_count: {e}")))?
}
None => {
let range_start = (session_id, String::new());
let range_end = (session_range_end(session_id), String::new());
table
.range::<(u64, String)>((range_start)..(range_end))
.map_err(|e| db_err(format!("redb range kv_count: {e}")))?
}
};
let mut count: u64 = 0;
for result in range {
result.map_err(|e| db_err(format!("redb iter kv_count: {e}")))?;
count += 1;
}
Ok(count)
}
pub fn write_session_retry(
db: &redb::Database,
session_id: u64,
record: &SessionRecord,
) -> io::Result<()> {
let mut attempts = 0;
loop {
match write_session(db, session_id, record) {
Ok(()) => return Ok(()),
Err(_e) if attempts < 3 => {
attempts += 1;
std::thread::sleep(std::time::Duration::from_millis(1));
continue;
}
Err(e) => return Err(e),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use choreo_proto::Turn;
fn next_session_id(db: &redb::Database) -> io::Result<u64> {
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
let current = {
let mut table = write_txn
.open_table(META)
.map_err(|e| db_err(format!("redb open meta: {e}")))?;
let current = table
.get("next_session_id")
.map_err(|e| db_err(format!("redb get meta: {e}")))?
.map(|g| g.value())
.unwrap_or(1);
table
.insert("next_session_id", current.wrapping_add(1))
.map_err(|e| db_err(format!("redb set meta: {e}")))?;
current
};
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit meta: {e}")))?;
Ok(current)
}
fn dummy_turn() -> Turn {
Turn {
created_at: choreo_proto::TimestampMs::now(),
undone: false,
error: None,
user_text: Some("hello".into()),
assistant_text: None,
assistant_reasoning: None,
tool_calls: Vec::new(),
token_usage: None,
tool_results: Vec::new(),
displayed_images: Vec::new(),
reasoning_artifact: None,
reasoning_producer: None,
}
}
#[test]
fn round_trip() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let id = next_session_id(&db).unwrap();
assert_eq!(id, 1);
let record = SessionRecord {
title: Some("test session".into()),
selected_model: Some("gpt-4".into()),
reasoning_effort: None,
parent_session_id: None,
working_dir: Some("/tmp".into()),
turn_count: 1,
created_at: 1234567890000,
last_modified: 1234567890000,
active_tool_groups: vec!["core".into(), "git".into()],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
write_session(&db, id, &record).unwrap();
let read = read_session(&db, id).unwrap().unwrap();
assert_eq!(read.title, record.title);
assert_eq!(read.turn_count, record.turn_count);
let all = read_all_sessions(&db).unwrap();
assert_eq!(all.len(), 1);
assert_eq!(all[0].0, id);
let turn = dummy_turn();
write_turn(&db, id, 0, &turn).unwrap();
let turns = read_turns(&db, id).unwrap();
assert_eq!(turns.len(), 1);
assert_eq!(turns[0].1, turn);
let id2 = next_session_id(&db).unwrap();
assert_eq!(id2, 2);
delete_session(&db, id).unwrap();
assert!(read_session(&db, id).unwrap().is_none());
assert!(read_turns(&db, id).unwrap().is_empty());
drop(db);
}
#[test]
fn session_record_last_response_id_round_trips() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let id = 1u64;
let record = SessionRecord {
title: Some("t".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1,
last_modified: 1,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: Some("resp_1".into()),
last_response_id_producer: Some(ReasoningProducer {
provider_slug: "openai".into(),
model: "gpt-5.4".into(),
}),
};
write_session(&db, id, &record).unwrap();
let read = read_session(&db, id).unwrap().unwrap();
assert_eq!(read.last_response_id.as_deref(), Some("resp_1"));
assert_eq!(
read.last_response_id_producer
.as_ref()
.map(|p| p.model.as_str()),
Some("gpt-5.4"),
"response id provenance must survive the write/read cycle",
);
assert_eq!(read.title.as_deref(), Some("t"));
}
#[test]
fn read_turns_skips_corrupt_entries() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let id = 1u64;
let valid_turn = dummy_turn();
write_turn(&db, id, 0, &valid_turn).unwrap();
{
let write_txn = db.begin_write().unwrap();
{
let mut table = write_txn.open_table(SESSION_TURNS).unwrap();
table
.insert((id, 1u32), b"not valid postcard data".as_slice())
.unwrap();
}
write_txn.commit().unwrap();
}
let valid_turn2 = dummy_turn();
write_turn(&db, id, 2, &valid_turn2).unwrap();
let turns = read_turns(&db, id).unwrap();
assert_eq!(turns.len(), 2, "corrupt turn should be skipped");
assert_eq!(turns[0].1, valid_turn);
assert_eq!(turns[1].1, valid_turn2);
}
#[test]
fn read_session_skips_corrupt_record_with_warning() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
{
let write_txn = db.begin_write().unwrap();
{
let mut table = write_txn.open_table(SESSIONS).unwrap();
table
.insert(42u64, b"not a session record".as_slice())
.unwrap();
}
write_txn.commit().unwrap();
}
assert!(
read_session(&db, 42).unwrap().is_none(),
"undecodable record must read as absent, not error"
);
assert!(read_session(&db, 99).unwrap().is_none());
}
#[test]
fn purge_removes_tombstoned_resurrected_record() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let record = SessionRecord {
title: Some("ghost".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
write_session(&db, 5, &record).unwrap();
mark_session_deleted(&db, 5).unwrap();
write_session(&db, 5, &record).unwrap();
let purged = purge_tombstoned_sessions(&db).unwrap();
assert_eq!(purged, 1, "the resurrected record must be purged");
assert!(
read_session(&db, 5).unwrap().is_none(),
"tombstoned session must not survive the purge"
);
assert_eq!(purge_tombstoned_sessions(&db).unwrap(), 0);
}
#[test]
fn clear_tombstone_prevents_purge_of_live_record() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let record = SessionRecord {
title: Some("live".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
write_session(&db, 6, &record).unwrap();
mark_session_deleted(&db, 6).unwrap();
clear_session_tombstone(&db, 6).unwrap();
let purged = purge_tombstoned_sessions(&db).unwrap();
assert_eq!(purged, 0);
assert!(read_session(&db, 6).unwrap().is_some());
}
#[test]
fn purge_empty_database_is_zero() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
assert_eq!(purge_tombstoned_sessions(&db).unwrap(), 0);
}
#[test]
fn run_migrations_stamps_fresh_database_v1() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
assert_eq!(
current_schema_version(&db).unwrap(),
0,
"a database created without open_db must be unversioned"
);
run_migrations(&db).unwrap();
assert_eq!(
current_schema_version(&db).unwrap(),
SCHEMA_VERSION,
"0 → 1 initialization must stamp the current schema version"
);
}
#[test]
fn production_migration_chain_matches_schema_version() {
let provided: Vec<u64> = MIGRATIONS.iter().map(|m| m.from).collect();
let expected: Vec<u64> = (1..SCHEMA_VERSION).collect();
assert_eq!(provided, expected);
}
#[test]
fn run_migrations_is_idempotent() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
run_migrations(&db).unwrap();
run_migrations(&db).unwrap();
assert_eq!(current_schema_version(&db).unwrap(), SCHEMA_VERSION);
}
#[test]
fn run_migrations_rejects_newer_schema_version() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
{
let write_txn = db.begin_write().unwrap();
{
let mut table = write_txn.open_table(META).unwrap();
table.insert(SCHEMA_VERSION_KEY, 5u64).unwrap();
}
write_txn.commit().unwrap();
}
let err = run_migrations(&db).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("newer") && msg.contains('5'),
"error must name the newer version: {msg}"
);
}
#[test]
fn legacy_unversioned_db_with_postcard_blobs_stamps_v1_and_skips() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let legacy_record = SessionRecord {
title: Some("legacy".into()),
selected_model: None,
reasoning_effort: None,
parent_session_id: None,
working_dir: None,
turn_count: 0,
created_at: 1000,
last_modified: 1000,
active_tool_groups: vec![],
context_config: ContextConfig::default(),
account_name: None,
last_response_id: None,
last_response_id_producer: None,
};
let legacy_blob = postcard::to_allocvec(&legacy_record).unwrap();
{
let write_txn = db.begin_write().unwrap();
{
let mut table = write_txn.open_table(SESSIONS).unwrap();
table.insert(42u64, legacy_blob.as_slice()).unwrap();
}
write_txn.commit().unwrap();
}
assert_eq!(current_schema_version(&db).unwrap(), 0);
run_migrations(&db).unwrap();
assert_eq!(current_schema_version(&db).unwrap(), SCHEMA_VERSION);
let all = read_all_sessions(&db).unwrap();
assert!(
all.is_empty(),
"legacy postcard blob must be skipped, not decoded or fatal"
);
}
#[test]
fn run_migrations_writes_no_backup_while_chain_is_empty() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("test.redb");
let db = redb::Database::create(&db_path).unwrap();
run_migrations(&db).unwrap();
assert!(
!db_path.with_file_name("test.redb.bak-v1").exists(),
"no backup artifact may be produced while the migration chain is empty"
);
}
fn dummy_migrate_1_to_2(db: &redb::Database) -> io::Result<()> {
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
let mut table = write_txn
.open_table(META)
.map_err(|e| db_err(format!("redb open meta: {e}")))?;
table
.insert("migrated", 1u64)
.map_err(|e| db_err(format!("redb set migrated marker: {e}")))?;
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit migrated marker: {e}")))?;
Ok(())
}
#[test]
fn run_migrations_applies_contiguous_chain_from_current_version() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
stamp_schema_version(&db, 1).unwrap();
run_migrations_to(
&db,
2,
&[Migration {
from: 1,
run: dummy_migrate_1_to_2,
}],
)
.unwrap();
assert_eq!(current_schema_version(&db).unwrap(), 2);
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(META).unwrap();
assert_eq!(
table.get("migrated").unwrap().unwrap().value(),
1,
"the dummy migration must have run"
);
}
}
#[test]
fn backup_db_file_names_backup_after_source_version() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("state.redb");
fs::write(&db_path, b"database contents").unwrap();
backup_db_file(&db_path, 1).unwrap();
assert!(db_path.with_file_name("state.redb.bak-v1").exists());
backup_db_file(&db_path, 2).unwrap();
assert!(db_path.with_file_name("state.redb.bak-v2").exists());
}
#[test]
fn run_migrations_rejects_non_contiguous_chain_before_writing() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
stamp_schema_version(&db, 1).unwrap();
let err = run_migrations_to(
&db,
2,
&[Migration {
from: 0,
run: dummy_migrate_1_to_2,
}],
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("not contiguous") && msg.contains('0') && msg.contains('1'),
"error must describe the chain mismatch: {msg}"
);
assert_eq!(current_schema_version(&db).unwrap(), 1);
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(META).unwrap();
assert!(
table.get("migrated").unwrap().is_none(),
"no migration may run when the chain is rejected"
);
}
}
#[test]
fn run_migrations_refuses_unversioned_db_when_target_above_initial() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let err = run_migrations_to(
&db,
2,
&[Migration {
from: 1,
run: dummy_migrate_1_to_2,
}],
)
.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("no schema version"),
"error must name the pre-release refusal: {msg}"
);
assert_eq!(current_schema_version(&db).unwrap(), 0);
{
let read_txn = db.begin_read().unwrap();
match read_txn.open_table(META) {
Ok(table) => assert!(
table.get("migrated").unwrap().is_none(),
"no migration may run when a pre-existing unversioned DB is refused"
),
Err(redb::TableError::TableDoesNotExist(_)) => {}
Err(e) => panic!("unexpected table error: {e}"),
}
}
}
}