use std::collections::BTreeMap;
use std::path::Path;
use std::sync::Arc;
use anyhow::{Result, anyhow};
use znippy_common::GUNNAR_REFS_MODULE;
use znippy_common::arrow::array::{Array, StringArray, StringBuilder, UInt64Array, UInt64Builder};
use znippy_common::arrow::datatypes::{DataType, Field, Schema};
use znippy_common::arrow::record_batch::RecordBatch;
use crate::pushlog::{PushLog, PushLogScan, read_sealed};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RefUpdate {
pub name: String,
pub target: Option<String>,
pub peeled: Option<String>,
pub symref_target: Option<String>,
}
impl RefUpdate {
pub fn set(name: impl Into<String>, target: impl Into<String>) -> Self {
Self {
name: name.into(),
target: Some(target.into()),
peeled: None,
symref_target: None,
}
}
pub fn delete(name: impl Into<String>) -> Self {
Self { name: name.into(), target: None, peeled: None, symref_target: None }
}
pub fn symbolic(name: impl Into<String>, points_to: impl Into<String>) -> Self {
Self {
name: name.into(),
target: None,
peeled: None,
symref_target: Some(points_to.into()),
}
}
pub fn with_peeled(mut self, peeled: impl Into<String>) -> Self {
self.peeled = Some(peeled.into());
self
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RefState {
pub target: Option<String>,
pub peeled: Option<String>,
pub symref_target: Option<String>,
pub push_seq: u64,
pub updated_ms: u64,
}
pub fn refs_schema() -> Arc<Schema> {
Arc::new(Schema::new(vec![
Field::new("name", DataType::Utf8, false),
Field::new("target", DataType::Utf8, true),
Field::new("peeled", DataType::Utf8, true),
Field::new("symref_target", DataType::Utf8, true),
Field::new("push_seq", DataType::UInt64, false),
Field::new("updated_ms", DataType::UInt64, false),
]))
}
pub fn build_push_batch(updates: &[RefUpdate], push_seq: u64, updated_ms: u64) -> Result<RecordBatch> {
let n = updates.len();
let mut name = StringBuilder::with_capacity(n, n * 32);
let mut target = StringBuilder::with_capacity(n, n * 64);
let mut peeled = StringBuilder::with_capacity(n, n * 64);
let mut symref = StringBuilder::with_capacity(n, n * 32);
let mut seq = UInt64Builder::with_capacity(n);
let mut ms = UInt64Builder::with_capacity(n);
for u in updates {
name.append_value(&u.name);
match &u.target {
Some(t) => target.append_value(t),
None => target.append_null(),
}
match &u.peeled {
Some(t) => peeled.append_value(t),
None => peeled.append_null(),
}
match &u.symref_target {
Some(t) => symref.append_value(t),
None => symref.append_null(),
}
seq.append_value(push_seq);
ms.append_value(updated_ms);
}
RecordBatch::try_new(
refs_schema(),
vec![
Arc::new(name.finish()),
Arc::new(target.finish()),
Arc::new(peeled.finish()),
Arc::new(symref.finish()),
Arc::new(seq.finish()),
Arc::new(ms.finish()),
],
)
.map_err(|e| anyhow!("refs push batch: {e}"))
}
pub struct RefLog {
log: PushLog,
}
impl RefLog {
pub fn new(path: impl Into<std::path::PathBuf>) -> Self {
Self { log: PushLog::new(path, refs_schema()) }
}
pub fn next_push_seq(&self) -> Result<u64> {
let scan = self.log.scan()?;
Ok(max_push_seq(&scan.pushes).map_or(0, |m| m + 1))
}
pub fn push(&self, updates: &[RefUpdate]) -> Result<u64> {
let seq = self.next_push_seq()?;
let ms = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
let batch = build_push_batch(updates, seq, ms)?;
self.log.append(&batch)?;
Ok(seq)
}
pub fn append_batch(&self, batch: &RecordBatch) -> Result<u64> {
self.log.append(batch)
}
pub fn scan(&self) -> Result<PushLogScan> {
self.log.scan()
}
pub fn compact(&self) -> Result<crate::pushlog::CompactionReport> {
self.log.compact()
}
pub fn maybe_compact(
&self,
policy: crate::pushlog::CompactionPolicy,
) -> Result<Option<crate::pushlog::CompactionReport>> {
self.log.maybe_compact(policy)
}
pub fn current(&self) -> Result<BTreeMap<String, RefState>> {
Ok(fold(&self.log.scan()?.pushes)?)
}
pub fn seal_section(&self) -> Result<znippy_common::ReservedSection> {
self.log.seal_section(GUNNAR_REFS_MODULE)
}
}
fn max_push_seq(batches: &[RecordBatch]) -> Option<u64> {
let mut max = None;
for b in batches {
let seq = b.column_by_name("push_seq")?.as_any().downcast_ref::<UInt64Array>()?;
for i in 0..seq.len() {
max = Some(max.map_or(seq.value(i), |m: u64| m.max(seq.value(i))));
}
}
max
}
pub fn fold(batches: &[RecordBatch]) -> Result<BTreeMap<String, RefState>> {
let mut rows: Vec<(u64, usize, RefState, String)> = Vec::new();
for (bi, b) in batches.iter().enumerate() {
let name = col::<StringArray>(b, "name")?;
let target = col::<StringArray>(b, "target")?;
let peeled = col::<StringArray>(b, "peeled")?;
let symref = col::<StringArray>(b, "symref_target")?;
let seq = col::<UInt64Array>(b, "push_seq")?;
let ms = col::<UInt64Array>(b, "updated_ms")?;
for i in 0..b.num_rows() {
rows.push((
seq.value(i),
bi,
RefState {
target: (!target.is_null(i)).then(|| target.value(i).to_string()),
peeled: (!peeled.is_null(i)).then(|| peeled.value(i).to_string()),
symref_target: (!symref.is_null(i)).then(|| symref.value(i).to_string()),
push_seq: seq.value(i),
updated_ms: ms.value(i),
},
name.value(i).to_string(),
));
}
}
rows.sort_by_key(|(seq, bi, _, _)| (*seq, *bi));
let mut out: BTreeMap<String, RefState> = BTreeMap::new();
for (_, _, state, name) in rows {
if state.target.is_none() && state.symref_target.is_none() {
out.remove(&name);
} else {
out.insert(name, state);
}
}
Ok(out)
}
pub fn read_refs(archive: &Path) -> Result<Option<BTreeMap<String, RefState>>> {
match read_sealed(archive, GUNNAR_REFS_MODULE)? {
Some(batches) => Ok(Some(fold(&batches)?)),
None => Ok(None),
}
}
fn col<'a, T: Array + 'static>(b: &'a RecordBatch, name: &str) -> Result<&'a T> {
b.column_by_name(name)
.ok_or_else(|| anyhow!("refs: no `{name}` column"))?
.as_any()
.downcast_ref::<T>()
.ok_or_else(|| anyhow!("refs: `{name}` has an unexpected type"))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::pushlog::truncate_for_test;
fn tmpdir(tag: &str) -> std::path::PathBuf {
let ns = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let d = std::env::temp_dir().join(format!("znippy_refs_{tag}_{ns}"));
std::fs::create_dir_all(&d).unwrap();
d
}
fn oid(c: char) -> String {
std::iter::repeat_n(c, 40).collect()
}
#[test]
fn last_writer_wins_and_a_null_target_deletes() {
let dir = tmpdir("fold");
let log = RefLog::new(dir.join("refs.log"));
log.push(&[
RefUpdate::set("refs/heads/main", oid('a')),
RefUpdate::set("refs/heads/topic", oid('b')),
])
.unwrap();
log.push(&[RefUpdate::set("refs/heads/main", oid('c'))]).unwrap();
log.push(&[RefUpdate::delete("refs/heads/topic")]).unwrap();
let refs = log.current().unwrap();
assert_eq!(refs["refs/heads/main"].target, Some(oid('c')), "second push must win");
assert!(!refs.contains_key("refs/heads/topic"), "a null target deletes the ref");
assert_eq!(refs.len(), 1);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn a_crash_mid_push_leaves_no_partial_ref_update() {
let dir = tmpdir("atomic");
let path = dir.join("refs.log");
let log = RefLog::new(&path);
log.push(&[RefUpdate::set("refs/heads/main", oid('a'))]).unwrap();
let before = std::fs::metadata(&path).unwrap().len();
log.push(&[
RefUpdate::set("refs/heads/main", oid('9')),
RefUpdate::set("refs/heads/a", oid('1')),
RefUpdate::set("refs/heads/b", oid('2')),
])
.unwrap();
let after = std::fs::metadata(&path).unwrap().len();
let intact = std::fs::read(&path).unwrap();
for cut in (before + 1)..after {
std::fs::write(&path, &intact).unwrap();
truncate_for_test(&path, cut).unwrap();
let refs = log.current().unwrap();
assert_eq!(
refs.len(),
1,
"cut at {cut}: a torn push must not publish ANY of its refs (got {refs:?})"
);
assert_eq!(
refs["refs/heads/main"].target,
Some(oid('a')),
"cut at {cut}: main must still be the pre-push value"
);
assert!(!refs.contains_key("refs/heads/a"), "cut at {cut}: leaked a partial ref");
assert!(!refs.contains_key("refs/heads/b"), "cut at {cut}: leaked a partial ref");
}
std::fs::write(&path, &intact).unwrap();
let refs = log.current().unwrap();
assert_eq!(refs.len(), 3, "the complete push publishes all three refs");
assert_eq!(refs["refs/heads/main"].target, Some(oid('9')));
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn push_seq_is_recovered_from_the_log_not_from_memory() {
let dir = tmpdir("seq");
let path = dir.join("refs.log");
let a = RefLog::new(&path);
assert_eq!(a.push(&[RefUpdate::set("refs/heads/main", oid('a'))]).unwrap(), 0);
assert_eq!(a.push(&[RefUpdate::set("refs/heads/main", oid('b'))]).unwrap(), 1);
let b = RefLog::new(&path);
assert_eq!(
b.next_push_seq().unwrap(),
2,
"a restarted writer must continue the sequence, not restart it"
);
assert_eq!(b.push(&[RefUpdate::set("refs/heads/main", oid('c'))]).unwrap(), 2);
assert_eq!(b.current().unwrap()["refs/heads/main"].target, Some(oid('c')));
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn ordering_is_by_push_seq_not_by_timestamp_or_position() {
let newer = build_push_batch(&[RefUpdate::set("refs/heads/main", oid('c'))], 7, 1000).unwrap();
let older = build_push_batch(&[RefUpdate::set("refs/heads/main", oid('a'))], 3, 9999).unwrap();
let refs = fold(&[newer, older]).unwrap();
assert_eq!(
refs["refs/heads/main"].target,
Some(oid('c')),
"push_seq 7 must beat push_seq 3 regardless of order or clock"
);
assert_eq!(refs["refs/heads/main"].push_seq, 7);
}
#[test]
fn symbolic_and_peeled_refs_round_trip() {
let dir = tmpdir("sym");
let log = RefLog::new(dir.join("refs.log"));
log.push(&[
RefUpdate::symbolic("HEAD", "refs/heads/main"),
RefUpdate::set("refs/tags/v1", oid('t')).with_peeled(oid('e')),
])
.unwrap();
let refs = log.current().unwrap();
assert_eq!(refs["HEAD"].symref_target.as_deref(), Some("refs/heads/main"));
assert!(refs["HEAD"].target.is_none(), "a symref has no direct target");
assert_eq!(refs["refs/tags/v1"].peeled, Some(oid('e')));
std::fs::remove_dir_all(&dir).ok();
}
}