use crate::error::Result;
use super::{StateRef, StateStore, open_connection, pg_sql};
#[derive(Debug, Clone)]
pub struct KeysetRangeRow {
pub range_index: i64,
pub lo: Option<String>,
pub hi: Option<String>,
pub done: bool,
}
pub struct KeysetRangePart {
pub file_name: String,
pub rows: i64,
pub bytes: i64,
}
impl StateStore {
pub fn persist_keyset_ranges(
&self,
export_name: &str,
run_id: &str,
ranges: &[(Option<String>, Option<String>)],
) -> Result<()> {
let now = chrono::Utc::now().to_rfc3339();
self.execute(
"DELETE FROM keyset_range WHERE export_name = ?1",
&[export_name.into()],
)?;
for (idx, (lo, hi)) in ranges.iter().enumerate() {
self.execute(
"INSERT INTO keyset_range \
(export_name, run_id, range_index, lo, hi, done, updated_at) \
VALUES (?1, ?2, ?3, ?4, ?5, 0, ?6)",
&[
export_name.into(),
run_id.into(),
(idx as i64).into(),
lo.clone().into(),
hi.clone().into(),
now.as_str().into(),
],
)?;
}
Ok(())
}
pub fn load_keyset_ranges(
&self,
export_name: &str,
run_id: &str,
) -> Result<Vec<KeysetRangeRow>> {
let sql = "SELECT range_index, lo, hi, done FROM keyset_range \
WHERE export_name = ?1 AND run_id = ?2 ORDER BY range_index";
self.query(sql, &[export_name.into(), run_id.into()], |r| {
KeysetRangeRow {
range_index: r.i64(0),
lo: r.opt_text(1),
hi: r.opt_text(2),
done: r.i64(3) != 0,
}
})
}
pub fn clear_keyset_ranges(&self, export_name: &str) -> Result<()> {
self.execute(
"DELETE FROM keyset_range WHERE export_name = ?1",
&[export_name.into()],
)?;
Ok(())
}
pub fn commit_keyset_range_at_ref(
state_ref: &StateRef,
run_id: &str,
export_name: &str,
range_index: i64,
parts: &[KeysetRangePart],
format: &str,
compression: Option<&str>,
) -> Result<()> {
let now = chrono::Utc::now().to_rfc3339();
let file_sql = "INSERT INTO file_log \
(run_id, export_name, file_name, row_count, bytes, format, compression, created_at) \
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)";
let done_sql = "UPDATE keyset_range SET done = 1, updated_at = ?1 \
WHERE export_name = ?2 AND range_index = ?3 AND run_id = ?4";
match state_ref {
StateRef::Sqlite(db_path) => {
let mut conn = open_connection(db_path)?;
let tx = conn.transaction()?;
for p in parts {
tx.execute(
file_sql,
rusqlite::params![
run_id,
export_name,
p.file_name,
p.rows,
p.bytes,
format,
compression,
now
],
)?;
}
tx.execute(
done_sql,
rusqlite::params![now, export_name, range_index, run_id],
)?;
tx.commit()?;
}
StateRef::Postgres(url) => {
let mut client = super::connect_pg(url)?;
let mut tx = client.transaction()?;
for p in parts {
tx.execute(
&pg_sql(file_sql),
&[
&run_id,
&export_name,
&p.file_name,
&p.rows,
&p.bytes,
&format,
&compression,
&now,
],
)?;
}
tx.execute(
&pg_sql(done_sql),
&[&now, &export_name, &range_index, &run_id],
)?;
tx.commit()?;
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn store() -> StateStore {
StateStore::open_in_memory().expect("in-memory store")
}
#[test]
fn persist_then_load_round_trips_ranges_all_not_done() {
let s = store();
let ranges = vec![
(None, Some("k0500".to_string())),
(Some("k0500".to_string()), Some("k1000".to_string())),
(Some("k1000".to_string()), None),
];
s.persist_keyset_ranges("exp", "run-1", &ranges).unwrap();
let loaded = s.load_keyset_ranges("exp", "run-1").unwrap();
assert_eq!(loaded.len(), 3);
assert_eq!(loaded[0].lo, None);
assert_eq!(loaded[0].hi.as_deref(), Some("k0500"));
assert_eq!(loaded[2].hi, None);
assert!(loaded.iter().all(|r| !r.done), "fresh ranges are not done");
}
#[test]
fn load_with_a_different_run_id_returns_nothing() {
let s = store();
s.persist_keyset_ranges("exp", "run-1", &[(None, None)])
.unwrap();
assert!(s.load_keyset_ranges("exp", "run-2").unwrap().is_empty());
}
#[test]
fn commit_range_marks_only_its_own_row_done_and_records_parts() {
let dir = tempfile::tempdir().unwrap();
let cfg = dir.path().join("rivet.yaml");
std::fs::write(&cfg, "# test").unwrap();
let s = StateStore::open(cfg.to_str().unwrap()).unwrap();
let ranges = vec![
(None, Some("k5".to_string())),
(Some("k5".to_string()), None),
];
s.persist_keyset_ranges("exp", "run-1", &ranges).unwrap();
StateStore::commit_keyset_range_at_ref(
s.state_ref(),
"run-1",
"exp",
1,
&[KeysetRangePart {
file_name: "exp_run-1_pk_w1_0.parquet".to_string(),
rows: 42,
bytes: 100,
}],
"parquet",
Some("zstd"),
)
.unwrap();
let loaded = s.load_keyset_ranges("exp", "run-1").unwrap();
assert!(!loaded[0].done, "range 0 untouched");
assert!(loaded[1].done, "range 1 committed → done");
let files = s.list_files_for_run("run-1").unwrap();
assert_eq!(files.len(), 1);
assert_eq!(files[0].file_name, "exp_run-1_pk_w1_0.parquet");
assert_eq!(files[0].row_count, 42);
}
#[test]
fn clear_removes_all_ranges_for_the_export() {
let s = store();
s.persist_keyset_ranges("exp", "run-1", &[(None, None)])
.unwrap();
s.clear_keyset_ranges("exp").unwrap();
assert!(s.load_keyset_ranges("exp", "run-1").unwrap().is_empty());
}
}