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};
mod codec;
use codec::{ZSTD_FRAME_MAGIC, zstd_decode, zstd_encode};
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 CATALOG_STATE: TableDefinition<&str, &[u8]> = TableDefinition::new("catalog_state");
const CATALOG_LAST_ATTEMPT_KEY: &str = "last_attempt_ms";
const CATALOG_ETAG_KEY: &str = "etag";
const SESSION_KV: TableDefinition<(u64, String), Vec<u8>> = TableDefinition::new("session_kv");
const SESSION_ATTACHMENTS: TableDefinition<(u64, u32, String), &[u8]> =
TableDefinition::new("session_attachments");
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 = 2;
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] = &[Migration {
from: 1,
run: migrate_turn_values_to_zstd,
}];
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))
}
pub fn schema_version(db: &redb::Database) -> io::Result<u64> {
current_schema_version(db)
}
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_path_for(path: &std::path::Path, from: u64) -> std::path::PathBuf {
let file_name = path
.file_name()
.map(|name| name.to_string_lossy().into_owned())
.unwrap_or_else(|| "state.redb".to_string());
path.with_file_name(format!("{file_name}.bak-v{from}"))
}
pub fn backup_database(path: &std::path::Path, from_version: u64) -> io::Result<()> {
let backup_path = backup_path_for(path, from_version);
fs::copy(path, &backup_path)?;
info!(
from = %path.display(),
to = %backup_path.display(),
"backed up database before applying migrations (taken before the database lock)"
);
Ok(())
}
fn backup_db_file(path: &std::path::Path, from: u64) -> io::Result<()> {
let backup_path = backup_path_for(path, from);
if backup_path.exists() {
info!(
path = %backup_path.display(),
"pre-migration backup already exists (created before the database lock was taken); \
skipping the copy and leaving it untouched"
);
return Ok(());
}
fs::copy(path, &backup_path)?;
info!(
from = %path.display(),
to = %backup_path.display(),
"backed up database before applying migrations"
);
Ok(())
}
pub(crate) fn migration_backup_version(db: &redb::Database) -> io::Result<Option<u64>> {
let current = schema_version(db)?;
debug!(current, "checking whether a pre-migration backup is needed");
if current == 0 || current >= SCHEMA_VERSION || MIGRATIONS.is_empty() {
return Ok(None);
}
let expected: Vec<u64> = (1..SCHEMA_VERSION).collect();
let provided: Vec<u64> = MIGRATIONS.iter().map(|m| m.from).collect();
if provided != expected {
return Ok(None);
}
Ok(Some(current))
}
pub fn run_migrations(db: &redb::Database) -> io::Result<()> {
run_migrations_at(db, &db_path()?)
}
pub fn run_migrations_at(db: &redb::Database, path: &std::path::Path) -> io::Result<()> {
run_migrations_to(db, SCHEMA_VERSION, MIGRATIONS, path)
}
fn run_migrations_to(
db: &redb::Database,
target: u64,
migrations: &[Migration],
db_path: &std::path::Path,
) -> 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() {
backup_db_file(db_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> {
open_db_at(&db_path()?)
}
pub fn open_db_at(path: &std::path::Path) -> io::Result<redb::Database> {
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)
}
fn delete_session_attachments(
write_txn: &redb::WriteTransaction,
session_id: u64,
) -> io::Result<()> {
let mut att_table = write_txn
.open_table(SESSION_ATTACHMENTS)
.map_err(|e| db_err(format!("redb open session_attachments: {e}")))?;
let att_keys: Vec<(u64, u32, String)> = att_table
.range::<(u64, u32, String)>(
(session_id, 0u32, String::new())..(session_range_end(session_id), 0u32, String::new()),
)
.map_err(|e| db_err(format!("redb range session_attachments: {e}")))?
.filter_map(|result| result.ok())
.map(|(k, _)| k.value())
.collect();
for key in att_keys {
att_table
.remove(key)
.map_err(|e| db_err(format!("redb remove session_attachment: {e}")))?;
}
Ok(())
}
fn delete_turn_attachments(
write_txn: &redb::WriteTransaction,
session_id: u64,
turn_id: u32,
) -> io::Result<()> {
let mut att_table = write_txn
.open_table(SESSION_ATTACHMENTS)
.map_err(|e| db_err(format!("redb open session_attachments: {e}")))?;
let att_keys: Vec<(u64, u32, String)> = att_table
.range::<(u64, u32, String)>(
(session_id, turn_id, String::new())
..(session_id, turn_id.saturating_add(1), String::new()),
)
.map_err(|e| db_err(format!("redb range session_attachments: {e}")))?
.filter_map(|result| result.ok())
.map(|(k, _)| k.value())
.collect();
for key in att_keys {
att_table
.remove(key)
.map_err(|e| db_err(format!("redb remove session_attachment: {e}")))?;
}
Ok(())
}
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}")))?;
}
}
delete_session_attachments(&write_txn, session_id)?;
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)
}
fn migrate_turn_values_to_zstd(db: &redb::Database) -> io::Result<()> {
info!("applying 1→2 migration: re-encoding session_turns values with zstd");
let mut keys: Vec<(u64, u32)> = Vec::new();
let mut skipped = 0usize;
{
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = match read_txn.open_table(SESSION_TURNS) {
Err(redb::TableError::TableDoesNotExist(_)) => return Ok(()),
Ok(t) => t,
Err(e) => return Err(db_err(format!("redb open turns (migration): {e}"))),
};
let iter = table
.iter()
.map_err(|e| db_err(format!("redb iter turns (migration): {e}")))?;
for result in iter {
let (key, value) =
result.map_err(|e| db_err(format!("redb iter item (migration): {e}")))?;
let (sid, idx) = key.value();
if value.value().starts_with(&ZSTD_FRAME_MAGIC) {
debug!(
session_id = sid,
turn_id = idx,
"turn already zstd-compressed; skipping"
);
skipped += 1;
continue;
}
keys.push((sid, idx));
}
}
if keys.is_empty() {
info!(skipped, "no raw session_turns values to re-encode");
return Ok(());
}
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn (migration): {e}")))?;
{
let mut table = write_txn
.open_table(SESSION_TURNS)
.map_err(|e| db_err(format!("redb open turns (migration): {e}")))?;
for key in &keys {
let raw = table
.get(*key)
.map_err(|e| db_err(format!("redb get turn (migration): {e}")))?
.ok_or_else(|| db_err(format!("turn vanished during migration: {key:?}")))
.map(|g| g.value().to_vec())?;
let compressed = zstd_encode(&raw);
table
.insert(*key, compressed.as_slice())
.map_err(|e| db_err(format!("redb insert turn (migration): {e}")))?;
}
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit turn migration: {e}")))?;
info!(
encoded = keys.len(),
skipped, "re-encoded session_turns values with zstd"
);
Ok(())
}
pub fn write_turn(
db: &redb::Database,
session_id: u64,
turn_id: u32,
turn: &Turn,
) -> io::Result<()> {
let mut storage = turn.clone();
for img in &mut storage.displayed_images {
img.data.clear();
}
for tr in &mut storage.tool_results {
if let Some(image) = &mut tr.image {
image.data.clear();
}
}
let payload =
rmp_serde::to_vec_named(&storage).map_err(|e| db_err(format!("codec encode turn: {e}")))?;
let compressed = zstd_encode(&payload);
let write_txn = db
.begin_write()
.map_err(|e| db_err(format!("redb write txn: {e}")))?;
{
delete_turn_attachments(&write_txn, session_id, turn_id)?;
let mut attachments = write_txn
.open_table(SESSION_ATTACHMENTS)
.map_err(|e| db_err(format!("redb open session_attachments: {e}")))?;
for (i, img) in turn.displayed_images.iter().enumerate() {
if img.data.is_empty() {
continue; }
let slot = format!("d{i}");
attachments
.insert((session_id, turn_id, slot), img.data.as_slice())
.map_err(|e| db_err(format!("redb insert display attachment: {e}")))?;
}
for tr in &turn.tool_results {
if let Some(image) = &tr.image
&& !image.data.is_empty()
{
let slot = format!("r{}", tr.call_id);
attachments
.insert((session_id, turn_id, slot), image.data.as_slice())
.map_err(|e| db_err(format!("redb insert result attachment: {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), compressed.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 attachments = match read_txn.open_table(SESSION_ATTACHMENTS) {
Ok(t) => Some(t),
Err(redb::TableError::TableDoesNotExist(_)) => None,
Err(e) => return Err(db_err(format!("redb open session_attachments: {e}"))),
};
let mut turns: Vec<(u32, Turn)> = Vec::new();
let iter = table
.range::<(u64, u32)>((session_id, 0u32)..(session_range_end(session_id), 0u32))
.map_err(|e| db_err(format!("redb range turns: {e}")))?;
for result in iter {
let (key, value) = result.map_err(|e| db_err(format!("redb iter item: {e}")))?;
let (_, idx) = key.value();
match zstd_decode(value.value())
.and_then(|buf| rmp_serde::from_slice::<Turn>(&buf).map_err(io::Error::other))
{
Ok(mut turn) => {
if let Some(attachments) = &attachments {
for (i, img) in turn.displayed_images.iter_mut().enumerate() {
if img.data.is_empty() {
let slot = format!("d{i}");
if let Some(guard) = attachments
.get((session_id, idx, slot))
.map_err(|e| db_err(format!("redb get display attachment: {e}")))?
{
img.data = guard.value().to_vec();
}
}
}
for tr in turn.tool_results.iter_mut() {
if let Some(image) = &mut tr.image
&& image.data.is_empty()
{
let slot = format!("r{}", tr.call_id);
if let Some(guard) = attachments
.get((session_id, idx, slot))
.map_err(|e| db_err(format!("redb get result attachment: {e}")))?
{
image.data = guard.value().to_vec();
}
}
}
}
debug!(
session_id,
turn_id = idx,
"re-attached turn image attachments"
);
turns.push((idx, turn));
}
Err(e) => {
tracing::warn!(session_id, turn_id = idx, error = %e, "undecodable turn, skipping");
}
}
}
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}")))?;
}
}
delete_session_attachments(&write_txn, session_id)?;
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(())
}
fn catalog_state_get(db: &redb::Database, 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 = match read_txn.open_table(CATALOG_STATE) {
Ok(table) => table,
Err(redb::TableError::TableDoesNotExist(_)) => return Ok(None),
Err(e) => return Err(db_err(format!("redb open catalog_state: {e}"))),
};
match table
.get(key)
.map_err(|e| db_err(format!("redb get catalog_state {key}: {e}")))?
{
Some(guard) => Ok(Some(guard.value().to_vec())),
None => Ok(None),
}
}
fn catalog_state_write(db: &redb::Database, key: &str, value: Option<&[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(CATALOG_STATE)
.map_err(|e| db_err(format!("redb open catalog_state: {e}")))?;
match value {
Some(value) => {
table
.insert(key, value)
.map_err(|e| db_err(format!("redb set catalog_state {key}: {e}")))?;
}
None => {
table
.remove(key)
.map_err(|e| db_err(format!("redb remove catalog_state {key}: {e}")))?;
}
}
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit catalog_state {key}: {e}")))?;
Ok(())
}
pub fn set_catalog_last_attempt_ms(db: &redb::Database, ms: u64) -> io::Result<()> {
catalog_state_write(db, CATALOG_LAST_ATTEMPT_KEY, Some(&ms.to_le_bytes()))
}
pub fn get_catalog_last_attempt_ms(db: &redb::Database) -> io::Result<Option<u64>> {
let Some(bytes) = catalog_state_get(db, CATALOG_LAST_ATTEMPT_KEY)? else {
return Ok(None);
};
match <[u8; 8]>::try_from(bytes.as_slice()) {
Ok(bytes) => Ok(Some(u64::from_le_bytes(bytes))),
Err(_) => {
warn!("catalog last_attempt_ms has an invalid length; treating as absent");
Ok(None)
}
}
}
pub fn set_catalog_etag(db: &redb::Database, etag: Option<&str>) -> io::Result<()> {
catalog_state_write(db, CATALOG_ETAG_KEY, etag.map(str::as_bytes))
}
pub fn get_catalog_etag(db: &redb::Database) -> io::Result<Option<String>> {
let Some(bytes) = catalog_state_get(db, CATALOG_ETAG_KEY)? else {
return Ok(None);
};
let trimmed = String::from_utf8_lossy(&bytes).trim().to_string();
if trimmed.is_empty() {
Ok(None)
} else {
Ok(Some(trimmed))
}
}
const KEYSTORE: TableDefinition<&str, &[u8]> = TableDefinition::new("keystore");
const KEYSTORE_BINDING_KEY: &str = "binding";
pub fn get_keystore_binding(db: &redb::Database) -> io::Result<Option<[u8; 32]>> {
let read_txn = db
.begin_read()
.map_err(|e| db_err(format!("redb read txn: {e}")))?;
let table = match read_txn.open_table(KEYSTORE) {
Ok(table) => table,
Err(redb::TableError::TableDoesNotExist(_)) => return Ok(None),
Err(e) => return Err(db_err(format!("redb open keystore: {e}"))),
};
let Some(guard) = table
.get(KEYSTORE_BINDING_KEY)
.map_err(|e| db_err(format!("redb get keystore binding: {e}")))?
else {
return Ok(None);
};
match <[u8; 32]>::try_from(guard.value()) {
Ok(key) => Ok(Some(key)),
Err(_) => {
warn!(
stored_len = guard.value().len(),
"keystore binding has an invalid length; treating as unbound"
);
Ok(None)
}
}
}
pub fn set_keystore_binding(db: &redb::Database, public_key: &[u8; 32]) -> 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(KEYSTORE)
.map_err(|e| db_err(format!("redb open keystore: {e}")))?;
table
.insert(KEYSTORE_BINDING_KEY, public_key.as_slice())
.map_err(|e| db_err(format!("redb set keystore binding: {e}")))?;
}
write_txn
.commit()
.map_err(|e| db_err(format!("redb commit keystore binding: {e}")))?;
info!("persisted keystore binding (TOFU adoption)");
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::{
DisplayedImageRecord, ImageMetadata, ImageReference, ToolResultRecord, 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 read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_TURNS).unwrap();
let guard = table.get((id, 0u32)).unwrap().unwrap();
let v = guard.value();
assert!(
v.starts_with(&ZSTD_FRAME_MAGIC),
"turns must be stored as zstd frames in schema 2"
);
}
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 turn_image_bytes_split_into_attachments_and_reattached() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let id = 1u64;
let mut turn = dummy_turn();
turn.displayed_images = vec![DisplayedImageRecord {
metadata: ImageMetadata {
mime_type: "image/png".into(),
width: 32,
height: 32,
byte_len: 5,
alt: None,
},
data: b"\x89PNG\r".to_vec(),
tool_call_id: Some("call_disp".into()),
}];
turn.tool_results = vec![ToolResultRecord {
call_id: "call_v".into(),
name: "read_image".into(),
content: "image".into(),
is_error: false,
invocation_description: "read_image".into(),
image: Some(ImageReference {
path: "/tmp/foo.png".into(),
mime_type: "image/png".into(),
width: 16,
height: 16,
data: b"\x89PNG-vision".to_vec(),
}),
}];
write_turn(&db, id, 0, &turn).unwrap();
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_TURNS).unwrap();
let guard = table.get((id, 0u32)).unwrap().unwrap();
let blob = guard.value();
assert!(blob.starts_with(&ZSTD_FRAME_MAGIC));
let decoded: Turn = rmp_serde::from_slice(&zstd_decode(blob).unwrap()).unwrap();
assert_eq!(decoded.displayed_images[0].data, Vec::<u8>::new());
assert_eq!(
decoded.tool_results[0].image.as_ref().unwrap().data,
Vec::<u8>::new()
);
}
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_ATTACHMENTS).unwrap();
let d0 = table.get((id, 0u32, "d0".to_string())).unwrap().unwrap();
assert_eq!(d0.value(), b"\x89PNG\r");
let rv = table
.get((id, 0u32, "rcall_v".to_string()))
.unwrap()
.unwrap();
assert_eq!(rv.value(), b"\x89PNG-vision");
}
let turns = read_turns(&db, id).unwrap();
assert_eq!(turns.len(), 1);
assert_eq!(turns[0].1, turn);
delete_session_turns(&db, id).unwrap();
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_ATTACHMENTS).unwrap();
assert!(
table.get((id, 0u32, "d0".to_string())).unwrap().is_none(),
"delete_session_turns must remove the display attachment"
);
assert!(
table
.get((id, 0u32, "rcall_v".to_string()))
.unwrap()
.is_none(),
"delete_session_turns must remove the result attachment"
);
}
write_turn(&db, id, 0, &turn).unwrap();
delete_session(&db, id).unwrap();
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_ATTACHMENTS).unwrap();
assert!(
table.get((id, 0u32, "d0".to_string())).unwrap().is_none(),
"delete_session must remove the display attachment"
);
assert!(
table
.get((id, 0u32, "rcall_v".to_string()))
.unwrap()
.is_none(),
"delete_session must remove the result attachment"
);
}
drop(db);
}
#[test]
fn turn_rewrite_clears_stale_attachment_slots() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let id = 1u64;
let mut turn = dummy_turn();
turn.displayed_images = vec![DisplayedImageRecord {
metadata: ImageMetadata {
mime_type: "image/png".into(),
width: 1,
height: 1,
byte_len: 4,
alt: None,
},
data: b"AAAA".to_vec(),
tool_call_id: Some("c0".into()),
}];
turn.tool_results = vec![ToolResultRecord {
call_id: "call_v".into(),
name: "read_image".into(),
content: "image".into(),
is_error: false,
invocation_description: "read_image".into(),
image: Some(ImageReference {
path: "/tmp/a.png".into(),
mime_type: "image/png".into(),
width: 1,
height: 1,
data: b"BBBB".to_vec(),
}),
}];
write_turn(&db, id, 0, &turn).unwrap();
let mut shifted = turn.clone();
shifted.displayed_images[0].data.clear();
shifted.displayed_images.push(DisplayedImageRecord {
metadata: ImageMetadata {
mime_type: "image/png".into(),
width: 2,
height: 2,
byte_len: 4,
alt: None,
},
data: b"CCCC".to_vec(),
tool_call_id: Some("c1".into()),
});
shifted.tool_results[0].image.as_mut().unwrap().data.clear();
write_turn(&db, id, 0, &shifted).unwrap();
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_ATTACHMENTS).unwrap();
assert!(
table.get((id, 0u32, "d0".to_string())).unwrap().is_none(),
"stale d0 from the first write must be cleared on rewrite"
);
assert!(
table
.get((id, 0u32, "rcall_v".to_string()))
.unwrap()
.is_none(),
"stale rcall_v from the first write must be cleared on rewrite"
);
let d1 = table.get((id, 0u32, "d1".to_string())).unwrap().unwrap();
assert_eq!(d1.value(), b"CCCC");
}
let turns = read_turns(&db, id).unwrap();
assert_eq!(turns.len(), 1);
let read = &turns[0].1;
assert_eq!(read.displayed_images[0].data, Vec::<u8>::new());
assert_eq!(read.displayed_images[1].data, b"CCCC");
assert_eq!(
read.tool_results[0].image.as_ref().unwrap().data,
Vec::<u8>::new()
);
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 a zstd frame".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 catalog_last_attempt_ms_round_trips() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
assert_eq!(get_catalog_last_attempt_ms(&db).unwrap(), None);
set_catalog_last_attempt_ms(&db, 1_700_000_123_456).unwrap();
assert_eq!(
get_catalog_last_attempt_ms(&db).unwrap(),
Some(1_700_000_123_456)
);
set_catalog_last_attempt_ms(&db, 1_700_000_500_000).unwrap();
assert_eq!(
get_catalog_last_attempt_ms(&db).unwrap(),
Some(1_700_000_500_000)
);
}
#[test]
fn catalog_last_attempt_ms_corrupt_length_treated_as_absent() {
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(CATALOG_STATE).unwrap();
table
.insert(CATALOG_LAST_ATTEMPT_KEY, b"too short".as_slice())
.unwrap();
}
write_txn.commit().unwrap();
}
assert_eq!(get_catalog_last_attempt_ms(&db).unwrap(), None);
}
#[test]
fn keystore_binding_adopts_and_round_trips() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
assert_eq!(get_keystore_binding(&db).unwrap(), None);
let pubkey = [7u8; 32];
set_keystore_binding(&db, &pubkey).unwrap();
assert_eq!(get_keystore_binding(&db).unwrap(), Some(pubkey));
}
#[test]
fn keystore_binding_corrupt_length_treated_as_absent() {
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(KEYSTORE).unwrap();
table
.insert(KEYSTORE_BINDING_KEY, b"short".as_slice())
.unwrap();
}
write_txn.commit().unwrap();
}
assert_eq!(get_keystore_binding(&db).unwrap(), None);
}
#[test]
fn catalog_etag_round_trips_and_clears() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
assert_eq!(get_catalog_etag(&db).unwrap(), None);
set_catalog_etag(&db, Some("\"v1\"")).unwrap();
assert_eq!(get_catalog_etag(&db).unwrap().as_deref(), Some("\"v1\""));
set_catalog_etag(&db, Some("W/\"v2\"")).unwrap();
assert_eq!(get_catalog_etag(&db).unwrap().as_deref(), Some("W/\"v2\""));
set_catalog_etag(&db, None).unwrap();
assert_eq!(get_catalog_etag(&db).unwrap(), None);
}
#[test]
fn catalog_etag_blank_value_reads_as_absent() {
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(CATALOG_STATE).unwrap();
table.insert(CATALOG_ETAG_KEY, b" ".as_slice()).unwrap();
}
write_txn.commit().unwrap();
}
assert_eq!(get_catalog_etag(&db).unwrap(), None);
}
#[test]
fn migrate_turn_values_to_zstd_rewrites_legacy_rows() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let sid = 1u64;
let turns: Vec<Turn> = (0..3).map(|_| dummy_turn()).collect();
{
let write_txn = db.begin_write().unwrap();
{
let mut table = write_txn.open_table(SESSION_TURNS).unwrap();
for (i, turn) in turns.iter().enumerate() {
let raw = rmp_serde::to_vec_named(turn).unwrap();
table.insert((sid, i as u32), raw.as_slice()).unwrap();
}
}
write_txn.commit().unwrap();
}
assert_eq!(read_turns(&db, sid).unwrap().len(), 0);
migrate_turn_values_to_zstd(&db).unwrap();
let decoded = read_turns(&db, sid).unwrap();
assert_eq!(decoded.len(), 3);
for (i, (idx, turn)) in decoded.iter().enumerate() {
assert_eq!(*idx as usize, i);
assert_eq!(turn, &turns[i]);
}
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_TURNS).unwrap();
for i in 0..3 {
let guard = table.get((sid, i)).unwrap().unwrap();
let v = guard.value();
assert!(
v.starts_with(&ZSTD_FRAME_MAGIC),
"row {i} must be stored as a zstd frame"
);
}
}
}
#[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 migrate_turn_values_to_zstd_is_idempotent() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
let sid = 1u64;
let turn = dummy_turn();
{
let raw = rmp_serde::to_vec_named(&turn).unwrap();
let write_txn = db.begin_write().unwrap();
{
let mut table = write_txn.open_table(SESSION_TURNS).unwrap();
table.insert((sid, 0u32), raw.as_slice()).unwrap();
}
write_txn.commit().unwrap();
}
migrate_turn_values_to_zstd(&db).unwrap();
let after_first: Vec<u8> = {
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_TURNS).unwrap();
table.get((sid, 0)).unwrap().unwrap().value().to_vec()
};
assert!(after_first.starts_with(&ZSTD_FRAME_MAGIC));
migrate_turn_values_to_zstd(&db).unwrap();
{
let read_txn = db.begin_read().unwrap();
let table = read_txn.open_table(SESSION_TURNS).unwrap();
let guard = table.get((sid, 0)).unwrap().unwrap();
let v = guard.value();
assert_eq!(
v,
after_first.as_slice(),
"re-run must not rewrite an already-compressed row"
);
}
assert_eq!(read_turns(&db, sid).unwrap()[0].1, turn);
}
#[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 run_migrations_to_backs_up_and_stamps_for_first_real_migration() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("state.redb");
let db = redb::Database::create(&db_path).unwrap();
stamp_schema_version(&db, 1).unwrap();
run_migrations_to(&db, SCHEMA_VERSION, MIGRATIONS, &db_path).unwrap();
assert!(
db_path.with_file_name("state.redb.bak-v1").exists(),
"the 1→2 migration must back up the source-version file"
);
assert_eq!(
current_schema_version(&db).unwrap(),
SCHEMA_VERSION,
"the migration must reach the current schema version"
);
}
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,
}],
&dir.path().join("test.redb"),
)
.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,
}],
&dir.path().join("test.redb"),
)
.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 backup_database_produces_versioned_name_and_identical_content() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("state.redb");
fs::write(&db_path, b"database bytes v1").unwrap();
backup_database(&db_path, 1).unwrap();
let backup_path = db_path.with_file_name("state.redb.bak-v1");
assert!(backup_path.exists());
assert_eq!(fs::read(&backup_path).unwrap(), b"database bytes v1");
}
#[test]
fn run_migrations_to_does_not_overwrite_pre_existing_backup() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("state.redb");
let db = redb::Database::create(&db_path).unwrap();
stamp_schema_version(&db, 1).unwrap();
let backup_path = db_path.with_file_name("state.redb.bak-v1");
fs::write(&backup_path, b"pre-lock sentinel").unwrap();
run_migrations_to(&db, SCHEMA_VERSION, MIGRATIONS, &db_path).unwrap();
assert_eq!(
fs::read(&backup_path).unwrap(),
b"pre-lock sentinel",
"the pre-existing backup must not be overwritten by the runner"
);
assert_eq!(current_schema_version(&db).unwrap(), SCHEMA_VERSION);
}
#[test]
fn schema_version_accessor_returns_stamped_version() {
let dir = tempfile::tempdir().unwrap();
let db = redb::Database::create(dir.path().join("test.redb")).unwrap();
stamp_schema_version(&db, 1).unwrap();
assert_eq!(schema_version(&db).unwrap(), 1);
}
#[test]
fn migration_backup_version_matches_run_migrations_to_backup_conditions() {
let dir = tempfile::tempdir().unwrap();
let make_db_at = |version: u64| {
let db =
redb::Database::create(dir.path().join(format!("test-{version}.redb"))).unwrap();
if version > 0 {
stamp_schema_version(&db, version).unwrap();
}
db
};
assert_eq!(migration_backup_version(&make_db_at(1)).unwrap(), Some(1));
assert_eq!(migration_backup_version(&make_db_at(0)).unwrap(), None);
assert_eq!(
migration_backup_version(&make_db_at(SCHEMA_VERSION)).unwrap(),
None
);
assert_eq!(
migration_backup_version(&make_db_at(SCHEMA_VERSION + 1)).unwrap(),
None
);
}
#[test]
fn production_orchestration_pre_backs_up_then_migrates() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("state.redb");
let opened = redb::Database::create(&db_path).unwrap();
stamp_schema_version(&opened, 1).unwrap();
let version = migration_backup_version(&opened)
.unwrap()
.expect("a pending v1 migration must be reported as needing a backup");
drop(opened);
backup_database(&db_path, version).unwrap();
let backup_path = db_path.with_file_name("state.redb.bak-v1");
let backup_before = fs::read(&backup_path).unwrap();
let reopened = redb::Database::open(&db_path).unwrap();
run_migrations_to(&reopened, SCHEMA_VERSION, MIGRATIONS, &db_path).unwrap();
assert_eq!(current_schema_version(&reopened).unwrap(), SCHEMA_VERSION);
assert_eq!(
fs::read(&backup_path).unwrap(),
backup_before,
"the pre-lock backup must survive the migration untouched"
);
}
#[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,
}],
&dir.path().join("test.redb"),
)
.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}"),
}
}
}
}