use kcode_k1_chat_persistence_store::Projection;
pub use kcode_k1_chat_persistence_store::{
Batch, EventRecord, Record, SessionId, SessionLog, TxId,
};
use kcode_k1_peering::K1Peering;
use kcode_k1_txn_ordering::{K1TxnOrdering, Subsystem, SubsystemId};
use std::collections::HashMap;
use std::fs::{self, File, OpenOptions};
use std::io::Write;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex, MutexGuard};
const CURSOR: &str = "cursor";
const CURSOR_TEMP: &str = "cursor.tmp";
const SUBSYSTEM: &str = "k1-chat-persist";
pub struct K1ChatPersistence {
driver: Arc<Driver>,
peering: Arc<K1Peering>,
}
#[derive(Clone)]
pub struct Session {
id: SessionId,
driver: Arc<Driver>,
peering: Arc<K1Peering>,
}
struct Driver {
state: Mutex<State>,
}
struct State {
root: PathBuf,
projection: Projection,
cursor: Option<TxId>,
first_callback: bool,
pending: HashMap<SessionId, Pending>,
poison: Option<String>,
}
struct Pending {
payload: Vec<u8>,
evidence: Option<TxId>,
}
impl K1ChatPersistence {
pub fn open(
root: &Path,
ordering: Arc<K1TxnOrdering>,
peering: Arc<K1Peering>,
) -> Result<Self, String> {
fs::create_dir_all(root).map_err(|e| format!("cannot create persistence root: {e}"))?;
let temp = root.join(CURSOR_TEMP);
if remove_optional(&temp).map_err(|e| format!("cannot remove stale cursor temp: {e}"))? {
sync_directory(root).map_err(|e| format!("cannot sync stale-temp removal: {e}"))?;
}
let projection =
Projection::new(root).map_err(|e| format!("cannot open session projection: {e}"))?;
let cursor = read_cursor(&root.join(CURSOR))?;
let driver = Arc::new(Driver {
state: Mutex::new(State {
root: root.to_path_buf(),
projection,
cursor,
first_callback: true,
pending: HashMap::new(),
poison: None,
}),
});
ordering
.register_subsystem(subsystem_id()?, cursor, driver.clone())
.map_err(|e| format!("cannot register {SUBSYSTEM}: {e}"))?;
Ok(Self { driver, peering })
}
pub fn session(&self, id: SessionId) -> Result<(Session, SessionLog), String> {
let session = Session {
id,
driver: self.driver.clone(),
peering: self.peering.clone(),
};
let log = session.load()?;
Ok((session, log))
}
}
impl Session {
pub fn id(&self) -> SessionId {
self.id
}
pub fn load(&self) -> Result<SessionLog, String> {
let state = self.driver.lock()?;
ensure_ready(&state)?;
state
.projection
.load(self.id)
.map_err(|e| format!("cannot load session: {e}"))
}
pub fn persist(&self, records: Vec<Record>) -> Result<TxId, String> {
if records.is_empty() {
return Err("cannot persist an empty record batch".to_owned());
}
let payload = {
let mut state = self.driver.lock()?;
ensure_ready(&state)?;
if state.pending.contains_key(&self.id) {
return Err("this session already has a local persist in flight".to_owned());
}
let predecessor = state
.projection
.load(self.id)
.map_err(|e| format!("cannot load session tail: {e}"))?
.records
.last()
.cloned();
let payload = Batch::new(self.id, predecessor, records)
.and_then(|batch| batch.encode())
.map_err(|e| format!("cannot encode persistence batch: {e}"))?;
state.pending.insert(
self.id,
Pending {
payload: payload.clone(),
evidence: None,
},
);
payload
};
let submitted = self.peering.submit_txn(subsystem_id()?, &payload);
self.finish_submit(submitted)
}
fn finish_submit(&self, submitted: Result<TxId, String>) -> Result<TxId, String> {
let mut state = self.driver.lock()?;
ensure_ready(&state)?;
let pending = match state.pending.remove(&self.id) {
Some(value) => value,
None => {
return Err(poison(
&mut state,
"local callback evidence disappeared".into(),
));
}
};
match (submitted, pending.evidence) {
(Ok(returned), Some(seen)) if returned == seen => Ok(returned),
(Err(_), Some(seen)) => Ok(seen),
(Err(error), None) => Err(error),
(Ok(returned), Some(seen)) => Err(poison(
&mut state,
format!(
"callback transaction mismatch: submission returned {returned:?}, callback saw {seen:?}"
),
)),
(Ok(returned), None) => Err(poison(
&mut state,
format!("missing callback evidence for returned {returned:?}"),
)),
}
}
}
impl Driver {
fn lock(&self) -> Result<MutexGuard<'_, State>, String> {
self.state
.lock()
.map_err(|_| "persistence state mutex is poisoned".to_owned())
}
fn fault(&self, message: String) -> String {
match self.lock() {
Ok(mut state) => poison(&mut state, message),
Err(lock_error) => format!("{message}; {lock_error}"),
}
}
}
impl Subsystem for Driver {
fn submit_txn(&self, id: TxId, payload: &[u8]) -> Result<(), String> {
let batch = Batch::decode(payload)
.map_err(|e| self.fault(format!("cannot decode callback {id:?}: {e}")))?;
let mut state = self.lock()?;
ensure_ready(&state)?;
let reconcile = state.first_callback;
if let Err(error) = state.projection.apply(id, &batch, reconcile) {
return Err(poison(
&mut state,
format!("cannot apply callback {id:?}: {error}"),
));
}
if let Err(error) = replace_cursor(&state.root, id) {
return Err(poison(
&mut state,
format!("cannot persist callback cursor {id:?}: {error}"),
));
}
state.cursor = Some(id);
state.first_callback = false;
if let Some(pending) = state.pending.get_mut(&batch.session_id)
&& pending.payload == payload
&& pending.evidence.replace(id).is_some()
{
return Err(poison(
&mut state,
format!("duplicate local callback evidence for {id:?}"),
));
}
Ok(())
}
fn reorg(&self) -> Result<(), String> {
let mut state = self.lock()?;
let mut failures = Vec::new();
if let Err(error) = state.projection.discard_all() {
failures.push(format!("discard projection: {error}"));
}
for (label, path) in [
(CURSOR, state.root.join(CURSOR)),
(CURSOR_TEMP, state.root.join(CURSOR_TEMP)),
] {
if let Err(error) = remove_optional(&path) {
failures.push(format!("remove {label}: {error}"));
}
}
if let Err(error) = sync_directory(&state.root) {
failures.push(format!("sync persistence root: {error}"));
}
state.cursor.take();
state.pending.clear();
let mut message = "canonical reorganization faulted the persistence driver".to_owned();
if !failures.is_empty() {
message.push_str("; cleanup failures: ");
message.push_str(&failures.join("; "));
}
state.poison = Some(message.clone());
Err(message)
}
}
fn subsystem_id() -> Result<SubsystemId, String> {
SubsystemId::from_str(SUBSYSTEM)
}
fn ensure_ready(state: &State) -> Result<(), String> {
match &state.poison {
Some(reason) => Err(format!("persistence driver is faulted: {reason}")),
None => Ok(()),
}
}
fn poison(state: &mut State, message: String) -> String {
state.pending.clear();
if state.poison.is_none() {
state.poison = Some(message);
}
state.poison.clone().expect("poison was just installed")
}
fn read_cursor(path: &Path) -> Result<Option<TxId>, String> {
let bytes = match fs::read(path) {
Ok(bytes) => bytes,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
Err(error) => return Err(format!("cannot read cursor: {error}")),
};
let exact: [u8; 12] = bytes
.try_into()
.map_err(|_| "cursor must contain exactly 12 bytes".to_owned())?;
Ok(Some(TxId::from_bytes(exact)))
}
fn replace_cursor(root: &Path, id: TxId) -> Result<(), String> {
let temp = root.join(CURSOR_TEMP);
let mut file = OpenOptions::new()
.create(true)
.truncate(true)
.write(true)
.open(&temp)
.map_err(|e| format!("open cursor temp: {e}"))?;
file.write_all(id.as_bytes())
.map_err(|e| format!("write cursor temp: {e}"))?;
file.sync_all()
.map_err(|e| format!("sync cursor temp: {e}"))?;
fs::rename(&temp, root.join(CURSOR)).map_err(|e| format!("rename cursor temp: {e}"))?;
sync_directory(root).map_err(|e| format!("sync cursor parent: {e}"))
}
fn remove_optional(path: &Path) -> std::io::Result<bool> {
match fs::remove_file(path) {
Ok(()) => Ok(true),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(false),
Err(error) => Err(error),
}
}
fn sync_directory(path: &Path) -> std::io::Result<()> {
File::open(path)?.sync_all()
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU64, Ordering};
static NEXT: AtomicU64 = AtomicU64::new(0);
struct Roots(PathBuf);
impl Roots {
fn new() -> Self {
let path = std::env::temp_dir().join(format!(
"k1-chat-persistence-{}-{}",
std::process::id(),
NEXT.fetch_add(1, Ordering::Relaxed)
));
let _ = fs::remove_dir_all(&path);
Self(path)
}
}
impl Drop for Roots {
fn drop(&mut self) {
let _ = fs::remove_dir_all(&self.0);
}
}
fn event(index: u64) -> Record {
Record::Event(EventRecord::new(0, index, 0, "test".into(), "value".into()).unwrap())
}
fn generic_box() -> Record {
Batch::decode(
br#"{"version":2,"session_id":"090909090909090909090909","predecessor":null,"records":[{"kind":"box","id":1,"type":"Future Kind","contents":"opaque","hidden_type":"future/v9","hidden_contents":"hidden bytes"}]}"#,
)
.unwrap()
.records
.into_iter()
.next()
.unwrap()
}
fn stack(roots: &Roots) -> (Arc<K1TxnOrdering>, Arc<K1Peering>) {
let ordering = Arc::new(K1TxnOrdering::open(&roots.0.join("ordering")).unwrap());
let peering =
Arc::new(K1Peering::open(&roots.0.join("peering"), ordering.clone()).unwrap());
(ordering, peering)
}
#[test]
fn store_v3_generic_box_round_trips_through_kto() {
let roots = Roots::new();
let (ordering, peering) = stack(&roots);
let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
let (session, _) = app.session([9; 12]).unwrap();
session.persist(vec![generic_box()]).unwrap();
let Record::Box(value) = &session.load().unwrap().records[0] else {
panic!("expected generic box");
};
assert_eq!(
(
value.box_type(),
value.contents(),
value.hidden_type(),
value.hidden_contents()
),
("Future Kind", "opaque", "future/v9", "hidden bytes")
);
}
#[test]
fn two_sessions_cursor_correlation_and_cold_restart() {
let roots = Roots::new();
let (ordering, peering) = stack(&roots);
let app = K1ChatPersistence::open(
&roots.0.join("persistence"),
ordering.clone(),
peering.clone(),
)
.unwrap();
assert!(
K1ChatPersistence::open(
&roots.0.join("persistence"),
ordering.clone(),
peering.clone()
)
.is_err()
);
let (one, log) = app.session([1; 12]).unwrap();
let (two, _) = app.session([2; 12]).unwrap();
assert!(log.records.is_empty());
let first = one.persist(vec![event(1)]).unwrap();
two.persist(vec![event(1)]).unwrap();
assert_eq!(one.load().unwrap().records.len(), 1);
assert_eq!(two.load().unwrap().records.len(), 1);
let cursor = fs::read(roots.0.join("persistence/cursor")).unwrap();
assert_eq!(cursor.len(), 12);
assert_ne!(cursor.as_slice(), first.as_bytes());
drop(one);
drop(two);
drop(app);
drop(peering);
drop(ordering);
let (ordering, peering) = stack(&roots);
let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
let (one, log) = app.session([1; 12]).unwrap();
assert_eq!(log.records, vec![event(1)]);
let id = one.persist(vec![event(2)]).unwrap();
assert_eq!(
fs::read(roots.0.join("persistence/cursor")).unwrap(),
id.as_bytes()
);
assert_eq!(one.load().unwrap().records, vec![event(1), event(2)]);
}
#[test]
fn overlap_rejection_and_reorg_discard_fault() {
let roots = Roots::new();
let (ordering, peering) = stack(&roots);
let app = K1ChatPersistence::open(&roots.0.join("persistence"), ordering, peering).unwrap();
let (session, _) = app.session([3; 12]).unwrap();
session.persist(vec![event(1)]).unwrap();
app.driver.state.lock().unwrap().pending.insert(
session.id(),
Pending {
payload: Vec::new(),
evidence: None,
},
);
assert!(
session
.persist(vec![event(2)])
.unwrap_err()
.contains("in flight")
);
app.driver.state.lock().unwrap().pending.clear();
let error = app.driver.reorg().unwrap_err();
assert!(error.contains("reorganization"));
assert!(session.load().unwrap_err().contains("faulted"));
assert!(!roots.0.join("persistence/cursor").exists());
assert!(!roots.0.join("persistence/sessions").exists());
}
}