use std::sync::Arc;
use liminal::durability::DurableStore;
use liminal_protocol::lifecycle::{
ObserverProgressAdvanceDecision, ObserverProgressTrackDecision, ObserverRecoveryAggregate,
};
use liminal_protocol::wire::{ConversationId, DeliverySeq, ObserverEpoch};
use serde::{Deserialize, Serialize};
use super::log::OperationLogError;
const OBSERVER_STREAM_KEY: &str = "liminal:participant-observer-recovery";
const READ_BATCH_SIZE: usize = 64;
const SCHEMA_VERSION: u8 = 1;
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(rename_all = "snake_case", tag = "row")]
pub(super) enum ObserverRow {
Track {
conversation_id: ConversationId,
observer_progress: DeliverySeq,
},
Advance {
conversation_id: ConversationId,
observer_progress: DeliverySeq,
},
Arms {
arms: Vec<(ConversationId, ObserverEpoch)>,
},
}
#[derive(Clone, Debug, Deserialize, Serialize)]
struct StoredRow {
schema_version: u8,
row: ObserverRow,
}
#[derive(Debug)]
pub(super) struct RestoredObserver {
pub(super) aggregate: ObserverRecoveryAggregate,
pub(super) next_sequence: u64,
}
#[derive(Debug)]
pub(super) struct ObserverLog {
store: Arc<dyn DurableStore>,
}
impl ObserverLog {
pub(super) fn new(store: Arc<dyn DurableStore>) -> Self {
Self { store }
}
pub(super) async fn append(
&self,
row: &ObserverRow,
expected_sequence: u64,
) -> Result<(), OperationLogError> {
let payload = serde_json::to_vec(&StoredRow {
schema_version: SCHEMA_VERSION,
row: row.clone(),
})?;
let assigned = self
.store
.append(OBSERVER_STREAM_KEY, payload, expected_sequence)
.await?;
if assigned != expected_sequence {
return Err(OperationLogError::AssignedSequence {
expected: expected_sequence,
actual: assigned,
});
}
self.store.flush().await?;
Ok(())
}
pub(super) async fn restore(&self) -> Result<RestoredObserver, OperationLogError> {
let mut aggregate = ObserverRecoveryAggregate::new();
let mut sequence = 0_u64;
loop {
let entries = self
.store
.read_from(OBSERVER_STREAM_KEY, sequence, READ_BATCH_SIZE)
.await?;
if entries.is_empty() {
break;
}
let count = entries.len();
for entry in entries {
if entry.sequence != sequence {
return Err(OperationLogError::Sequence {
expected: sequence,
actual: entry.sequence,
});
}
let stored: StoredRow = serde_json::from_slice(&entry.payload)?;
if stored.schema_version != SCHEMA_VERSION {
return Err(OperationLogError::SchemaVersion(stored.schema_version));
}
aggregate = fold_row(aggregate, stored.row, entry.sequence)?;
sequence = sequence.checked_add(1).ok_or(OperationLogError::Sequence {
expected: u64::MAX,
actual: entry.sequence,
})?;
}
if count < READ_BATCH_SIZE {
break;
}
}
Ok(RestoredObserver {
aggregate,
next_sequence: sequence,
})
}
}
fn fold_row(
aggregate: ObserverRecoveryAggregate,
row: ObserverRow,
sequence: u64,
) -> Result<ObserverRecoveryAggregate, OperationLogError> {
match row {
ObserverRow::Track {
conversation_id,
observer_progress,
} => match aggregate.decide_track(conversation_id, observer_progress) {
ObserverProgressTrackDecision::Commit(transaction) => Ok(transaction.commit()),
ObserverProgressTrackDecision::Refuse { .. } => Err(refused(sequence)),
},
ObserverRow::Advance {
conversation_id,
observer_progress,
} => match aggregate.decide_progress_advance(conversation_id, observer_progress) {
ObserverProgressAdvanceDecision::Commit(transaction) => {
let (aggregate, _fired) = transaction.commit();
Ok(aggregate)
}
ObserverProgressAdvanceDecision::Refuse { .. } => Err(refused(sequence)),
},
ObserverRow::Arms { arms } => {
let progress = aggregate.progress_rows();
let mut armed = aggregate.armed_rows();
armed.extend(arms);
armed.sort_unstable();
armed.dedup();
ObserverRecoveryAggregate::restore(&progress, &armed).map_err(|_| refused(sequence))
}
}
}
const fn refused(sequence: u64) -> OperationLogError {
OperationLogError::CorruptRow { sequence }
}