Skip to main content

oxirs_stream/fault_tolerance/
checkpoint.rs

1//! # Chandy-Lamport-style checkpointing for stream operators
2//!
3//! Implements a marker-based snapshot algorithm in the spirit of the original
4//! Chandy-Lamport protocol, adapted to stream processing topologies.
5//!
6//! ## Protocol summary
7//!
8//! 1. The coordinator broadcasts a [`Marker`] with a unique
9//!    [`CheckpointId`] to every *source* operator.
10//! 2. When an operator receives a marker on **input edge `e`**:
11//!    a. If this is the first marker observed in this checkpoint, the operator
12//!    (i) snapshots its local state, (ii) emits the marker on every output
13//!    edge, (iii) starts recording the messages that arrive on every other
14//!    input edge.
15//!    b. Otherwise, the operator stops recording on `e` and writes the
16//!    recorded prefix as part of the snapshot.
17//! 3. When the operator has received the marker on every input edge, the
18//!    snapshot is complete.
19//! 4. Snapshots are persisted via [`CheckpointStore`] (the production
20//!    implementation forwards them to the cluster snapshot store).
21//!
22//! ## Marker propagation
23//!
24//! [`MarkerPropagator`] tracks the per-operator marker state and surfaces a
25//! [`MarkerPropagatorEvent`] every time a checkpoint completes for an
26//! operator. A [`CheckpointController`] builds on top of the propagator to
27//! coordinate a global checkpoint.
28
29use std::collections::{HashMap, HashSet};
30use std::sync::atomic::{AtomicU64, Ordering};
31use std::sync::Arc;
32use std::time::{Duration, SystemTime};
33
34use async_trait::async_trait;
35use parking_lot::RwLock;
36use serde::{Deserialize, Serialize};
37use thiserror::Error;
38use tracing::{debug, info};
39
40// ─── Identifiers ────────────────────────────────────────────────────────────
41
42/// Unique identifier of a checkpoint round.
43pub type CheckpointId = u64;
44
45/// Identifier of an operator participating in a checkpoint.
46pub type OperatorId = String;
47
48/// Identifier of an input edge feeding into an operator.
49pub type InputEdgeId = String;
50
51// ─── Marker ─────────────────────────────────────────────────────────────────
52
53/// Chandy-Lamport marker.
54///
55/// Markers carry the [`CheckpointId`] and a wall-clock timestamp. They flow
56/// in-band with regular events on every operator edge.
57#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
58pub struct Marker {
59    pub checkpoint_id: CheckpointId,
60    pub emitted_at_ms: u64,
61}
62
63impl Marker {
64    /// Build a marker with the current wall-clock time.
65    pub fn new(checkpoint_id: CheckpointId) -> Self {
66        Self {
67            checkpoint_id,
68            emitted_at_ms: SystemTime::now()
69                .duration_since(SystemTime::UNIX_EPOCH)
70                .map(|d| d.as_millis() as u64)
71                .unwrap_or(0),
72        }
73    }
74}
75
76// ─── Snapshot ──────────────────────────────────────────────────────────────
77
78/// State snapshot produced by a single operator during a checkpoint round.
79#[derive(Debug, Clone, Serialize, Deserialize)]
80pub struct OperatorSnapshot {
81    pub operator_id: OperatorId,
82    pub checkpoint_id: CheckpointId,
83    /// Opaque serialised state.
84    pub state_blob: Vec<u8>,
85    /// Pre-snapshot in-flight events recorded per input edge after the marker
86    /// arrived on the *first* edge but before it arrived on others.
87    pub channel_logs: HashMap<InputEdgeId, Vec<Vec<u8>>>,
88    /// Wall-clock time at which the snapshot completed (Unix ms).
89    pub completed_at_ms: u64,
90}
91
92// ─── Errors ─────────────────────────────────────────────────────────────────
93
94/// Errors raised by the checkpointing pipeline.
95#[derive(Debug, Error)]
96pub enum CheckpointError {
97    /// Marker for an unknown operator.
98    #[error("unknown operator: {0}")]
99    UnknownOperator(OperatorId),
100    /// Marker for an unknown input edge on an operator.
101    #[error("unknown edge {edge} on operator {op}")]
102    UnknownEdge { op: OperatorId, edge: InputEdgeId },
103    /// Internal store error.
104    #[error("store error: {0}")]
105    Store(String),
106}
107
108/// Convenience alias.
109pub type CheckpointResult<T> = std::result::Result<T, CheckpointError>;
110
111// ─── CheckpointStore ───────────────────────────────────────────────────────
112
113/// Where completed snapshots are persisted.
114#[async_trait]
115pub trait CheckpointStore: Send + Sync {
116    /// Persist a single operator snapshot.
117    async fn put(&self, snap: OperatorSnapshot) -> CheckpointResult<()>;
118
119    /// Load a snapshot by `(operator, checkpoint)` pair, if any.
120    async fn get(
121        &self,
122        operator: &OperatorId,
123        checkpoint: CheckpointId,
124    ) -> CheckpointResult<Option<OperatorSnapshot>>;
125
126    /// Latest committed checkpoint id, if any. Used by the recovery path to
127    /// pick the most recent global snapshot.
128    async fn latest(&self) -> CheckpointResult<Option<CheckpointId>>;
129}
130
131/// In-memory snapshot store used in tests and as a default for embedded
132/// deployments. Production deployments swap this for an
133/// `oxirs-cluster`-backed snapshot store.
134pub struct InMemoryCheckpointStore {
135    inner: RwLock<HashMap<(OperatorId, CheckpointId), OperatorSnapshot>>,
136    latest: RwLock<Option<CheckpointId>>,
137}
138
139impl Default for InMemoryCheckpointStore {
140    fn default() -> Self {
141        Self::new()
142    }
143}
144
145impl InMemoryCheckpointStore {
146    /// Build an empty store.
147    pub fn new() -> Self {
148        Self {
149            inner: RwLock::new(HashMap::new()),
150            latest: RwLock::new(None),
151        }
152    }
153
154    /// Total number of snapshots currently in the store.
155    pub fn len(&self) -> usize {
156        self.inner.read().len()
157    }
158
159    /// True when no snapshots are stored.
160    pub fn is_empty(&self) -> bool {
161        self.inner.read().is_empty()
162    }
163}
164
165#[async_trait]
166impl CheckpointStore for InMemoryCheckpointStore {
167    async fn put(&self, snap: OperatorSnapshot) -> CheckpointResult<()> {
168        let cp = snap.checkpoint_id;
169        self.inner
170            .write()
171            .insert((snap.operator_id.clone(), cp), snap);
172        let mut latest = self.latest.write();
173        if latest.map_or(true, |old| cp > old) {
174            *latest = Some(cp);
175        }
176        Ok(())
177    }
178
179    async fn get(
180        &self,
181        operator: &OperatorId,
182        checkpoint: CheckpointId,
183    ) -> CheckpointResult<Option<OperatorSnapshot>> {
184        Ok(self
185            .inner
186            .read()
187            .get(&(operator.clone(), checkpoint))
188            .cloned())
189    }
190
191    async fn latest(&self) -> CheckpointResult<Option<CheckpointId>> {
192        Ok(*self.latest.read())
193    }
194}
195
196// ─── MarkerPropagator ──────────────────────────────────────────────────────
197
198/// Per-operator state used by the marker propagator.
199#[derive(Debug, Default)]
200struct OperatorMarkerState {
201    /// Set of input edges that have observed the marker so far.
202    seen_on: HashSet<InputEdgeId>,
203    /// Set of *all* configured input edges.
204    expected: HashSet<InputEdgeId>,
205    /// Recorded pre-snapshot messages keyed by input edge.
206    channel_logs: HashMap<InputEdgeId, Vec<Vec<u8>>>,
207}
208
209impl OperatorMarkerState {
210    fn new(expected: HashSet<InputEdgeId>) -> Self {
211        Self {
212            seen_on: HashSet::new(),
213            expected,
214            channel_logs: HashMap::new(),
215        }
216    }
217}
218
219/// Event emitted by [`MarkerPropagator::on_marker`].
220#[derive(Debug, Clone, PartialEq, Eq)]
221pub enum MarkerPropagatorEvent {
222    /// First marker observed for this checkpoint id; the operator should
223    /// snapshot its state and emit the marker on every output edge.
224    StartSnapshot,
225    /// Subsequent marker on a different edge; the recording for that edge is
226    /// closed.
227    EdgeClosed,
228    /// Marker observed on every input edge: the operator's checkpoint round is
229    /// complete.
230    Completed,
231}
232
233/// Tracks marker arrival per (operator, edge) pair.
234pub struct MarkerPropagator {
235    state: RwLock<HashMap<(OperatorId, CheckpointId), OperatorMarkerState>>,
236    expected_edges: RwLock<HashMap<OperatorId, HashSet<InputEdgeId>>>,
237}
238
239impl Default for MarkerPropagator {
240    fn default() -> Self {
241        Self::new()
242    }
243}
244
245impl MarkerPropagator {
246    /// Build an empty propagator.
247    pub fn new() -> Self {
248        Self {
249            state: RwLock::new(HashMap::new()),
250            expected_edges: RwLock::new(HashMap::new()),
251        }
252    }
253
254    /// Register an operator with the set of input edges it expects to see
255    /// markers on.
256    pub fn register_operator(
257        &self,
258        op: OperatorId,
259        edges: impl IntoIterator<Item = impl Into<InputEdgeId>>,
260    ) {
261        let edges: HashSet<InputEdgeId> = edges.into_iter().map(Into::into).collect();
262        self.expected_edges.write().insert(op, edges);
263    }
264
265    /// Returns true if the operator is registered.
266    pub fn is_registered(&self, op: &OperatorId) -> bool {
267        self.expected_edges.read().contains_key(op)
268    }
269
270    /// Record an in-flight event arriving on `edge` for an operator that has
271    /// already started recording for the given checkpoint. Events are stored
272    /// in the order they arrive.
273    pub fn record_inflight(
274        &self,
275        op: &OperatorId,
276        checkpoint: CheckpointId,
277        edge: &InputEdgeId,
278        payload: Vec<u8>,
279    ) -> CheckpointResult<bool> {
280        let mut states = self.state.write();
281        let st = match states.get_mut(&(op.clone(), checkpoint)) {
282            Some(s) => s,
283            None => return Ok(false),
284        };
285        if st.seen_on.is_empty() {
286            // No marker yet — nothing to record.
287            return Ok(false);
288        }
289        if st.seen_on.contains(edge) {
290            // Marker already seen on this edge: do not record.
291            return Ok(false);
292        }
293        st.channel_logs
294            .entry(edge.clone())
295            .or_default()
296            .push(payload);
297        Ok(true)
298    }
299
300    /// Process a marker arrival on `edge` for `operator`.
301    pub fn on_marker(
302        &self,
303        op: &OperatorId,
304        edge: &InputEdgeId,
305        marker: &Marker,
306    ) -> CheckpointResult<MarkerPropagatorEvent> {
307        let expected = {
308            let edges = self.expected_edges.read();
309            edges
310                .get(op)
311                .ok_or_else(|| CheckpointError::UnknownOperator(op.clone()))?
312                .clone()
313        };
314        if !expected.contains(edge) {
315            return Err(CheckpointError::UnknownEdge {
316                op: op.clone(),
317                edge: edge.clone(),
318            });
319        }
320        let mut states = self.state.write();
321        let entry = states
322            .entry((op.clone(), marker.checkpoint_id))
323            .or_insert_with(|| OperatorMarkerState::new(expected.clone()));
324
325        let event = if entry.seen_on.is_empty() {
326            entry.seen_on.insert(edge.clone());
327            MarkerPropagatorEvent::StartSnapshot
328        } else if entry.seen_on.contains(edge) {
329            MarkerPropagatorEvent::EdgeClosed
330        } else {
331            entry.seen_on.insert(edge.clone());
332            MarkerPropagatorEvent::EdgeClosed
333        };
334
335        let completed = entry.seen_on == entry.expected;
336        if completed {
337            Ok(MarkerPropagatorEvent::Completed)
338        } else {
339            Ok(event)
340        }
341    }
342
343    /// Drain the recorded channel logs for an operator's checkpoint round.
344    pub fn drain_channel_logs(
345        &self,
346        op: &OperatorId,
347        checkpoint: CheckpointId,
348    ) -> HashMap<InputEdgeId, Vec<Vec<u8>>> {
349        let mut states = self.state.write();
350        match states.remove(&(op.clone(), checkpoint)) {
351            Some(s) => s.channel_logs,
352            None => HashMap::new(),
353        }
354    }
355
356    /// Reset all per-checkpoint state for an operator.
357    ///
358    /// Clears every checkpoint round currently tracked for `op`. The
359    /// operator registration (its expected input edges) is preserved so the
360    /// next checkpoint round can resume immediately.
361    pub fn reset(&self, op: &OperatorId) {
362        let mut states = self.state.write();
363        states.retain(|(o, _), _| o != op);
364    }
365}
366
367// ─── CheckpointController ──────────────────────────────────────────────────
368
369/// Configuration for [`CheckpointController`].
370#[derive(Debug, Clone, Serialize, Deserialize)]
371pub struct CheckpointControllerConfig {
372    /// How often to issue a new checkpoint.
373    pub interval: Duration,
374    /// Maximum time an operator has to commit its snapshot before it is
375    /// considered failed (the controller will issue a fresh checkpoint).
376    pub timeout: Duration,
377}
378
379impl Default for CheckpointControllerConfig {
380    fn default() -> Self {
381        Self {
382            interval: Duration::from_secs(30),
383            timeout: Duration::from_secs(60),
384        }
385    }
386}
387
388/// Per-checkpoint completion progress.
389#[derive(Debug, Clone, Default, Serialize, Deserialize)]
390pub struct CheckpointProgress {
391    pub checkpoint_id: CheckpointId,
392    /// Set of operators that have committed a snapshot for this round.
393    pub committed: HashSet<OperatorId>,
394    /// Total expected operators.
395    pub expected: HashSet<OperatorId>,
396    pub started_at_ms: u64,
397    pub completed_at_ms: Option<u64>,
398}
399
400impl CheckpointProgress {
401    pub fn is_complete(&self) -> bool {
402        !self.expected.is_empty() && self.committed == self.expected
403    }
404}
405
406/// Coordinator that drives checkpoint rounds across the operator topology.
407pub struct CheckpointController {
408    config: CheckpointControllerConfig,
409    propagator: Arc<MarkerPropagator>,
410    store: Arc<dyn CheckpointStore>,
411    next_id: AtomicU64,
412    /// Set of operators participating.
413    operators: RwLock<HashSet<OperatorId>>,
414    /// Round-by-round progress.
415    progress: RwLock<HashMap<CheckpointId, CheckpointProgress>>,
416}
417
418impl CheckpointController {
419    /// Build a controller.
420    pub fn new(
421        config: CheckpointControllerConfig,
422        propagator: Arc<MarkerPropagator>,
423        store: Arc<dyn CheckpointStore>,
424    ) -> Self {
425        Self {
426            config,
427            propagator,
428            store,
429            next_id: AtomicU64::new(1),
430            operators: RwLock::new(HashSet::new()),
431            progress: RwLock::new(HashMap::new()),
432        }
433    }
434
435    /// Configuration accessor.
436    pub fn config(&self) -> &CheckpointControllerConfig {
437        &self.config
438    }
439
440    /// Marker propagator handle (so operator code can re-use the same
441    /// propagator).
442    pub fn propagator(&self) -> &Arc<MarkerPropagator> {
443        &self.propagator
444    }
445
446    /// Register an operator with the controller.
447    pub fn register_operator(&self, op: OperatorId) {
448        self.operators.write().insert(op);
449    }
450
451    /// Open a new checkpoint round.
452    pub fn open(&self) -> Marker {
453        let id = self.next_id.fetch_add(1, Ordering::Relaxed);
454        let marker = Marker::new(id);
455        let expected = self.operators.read().clone();
456        self.progress.write().insert(
457            id,
458            CheckpointProgress {
459                checkpoint_id: id,
460                committed: HashSet::new(),
461                expected,
462                started_at_ms: marker.emitted_at_ms,
463                completed_at_ms: None,
464            },
465        );
466        debug!(checkpoint_id = id, "checkpoint controller: opened");
467        marker
468    }
469
470    /// Acknowledge a snapshot from an operator.
471    pub async fn commit_snapshot(&self, snapshot: OperatorSnapshot) -> CheckpointResult<bool> {
472        let cp = snapshot.checkpoint_id;
473        let op = snapshot.operator_id.clone();
474        self.store.put(snapshot).await?;
475        let mut progress = self.progress.write();
476        if let Some(p) = progress.get_mut(&cp) {
477            p.committed.insert(op.clone());
478            if p.is_complete() && p.completed_at_ms.is_none() {
479                p.completed_at_ms = Some(now_ms());
480                info!(checkpoint_id = cp, "checkpoint complete");
481                return Ok(true);
482            }
483        }
484        Ok(false)
485    }
486
487    /// Fetch the current progress snapshot for a checkpoint.
488    pub fn progress(&self, cp: CheckpointId) -> Option<CheckpointProgress> {
489        self.progress.read().get(&cp).cloned()
490    }
491
492    /// Drop progress for an old checkpoint (used when retiring rounds).
493    pub fn forget(&self, cp: CheckpointId) {
494        self.progress.write().remove(&cp);
495    }
496
497    /// Latest checkpoint id known to the underlying store.
498    pub async fn latest_committed(&self) -> CheckpointResult<Option<CheckpointId>> {
499        self.store.latest().await
500    }
501
502    /// Number of rounds opened by the controller so far.
503    pub fn opened_rounds(&self) -> u64 {
504        self.next_id.load(Ordering::Relaxed).saturating_sub(1)
505    }
506
507    /// Snapshot store reference for callers that need to load on recovery.
508    pub fn store(&self) -> &Arc<dyn CheckpointStore> {
509        &self.store
510    }
511}
512
513fn now_ms() -> u64 {
514    SystemTime::now()
515        .duration_since(SystemTime::UNIX_EPOCH)
516        .map(|d| d.as_millis() as u64)
517        .unwrap_or(0)
518}
519
520// ─── Tests ──────────────────────────────────────────────────────────────────
521
522#[cfg(test)]
523mod tests {
524    use super::*;
525
526    fn op(name: &str) -> OperatorId {
527        name.to_string()
528    }
529
530    fn edge(name: &str) -> InputEdgeId {
531        name.to_string()
532    }
533
534    #[tokio::test]
535    async fn marker_propagator_completes_on_all_edges() {
536        let prop = MarkerPropagator::new();
537        prop.register_operator(op("op1"), ["e1", "e2"]);
538        let marker = Marker::new(1);
539
540        let ev1 = prop
541            .on_marker(&op("op1"), &edge("e1"), &marker)
542            .expect("ok");
543        assert_eq!(ev1, MarkerPropagatorEvent::StartSnapshot);
544        let ev2 = prop
545            .on_marker(&op("op1"), &edge("e2"), &marker)
546            .expect("ok");
547        assert_eq!(ev2, MarkerPropagatorEvent::Completed);
548    }
549
550    #[test]
551    fn marker_propagator_records_inflight_only_after_first_marker() {
552        let prop = MarkerPropagator::new();
553        prop.register_operator(op("op1"), ["e1", "e2"]);
554        let marker = Marker::new(2);
555        // No marker yet → no recording.
556        let ok = prop
557            .record_inflight(&op("op1"), 2, &edge("e2"), b"early".to_vec())
558            .expect("ok");
559        assert!(!ok);
560        // First marker on e1.
561        let _ = prop
562            .on_marker(&op("op1"), &edge("e1"), &marker)
563            .expect("ok");
564        // Now record on e2.
565        let ok = prop
566            .record_inflight(&op("op1"), 2, &edge("e2"), b"after-marker".to_vec())
567            .expect("ok");
568        assert!(ok);
569        // Marker on e1 already seen → recording on e1 is suppressed.
570        let ok = prop
571            .record_inflight(&op("op1"), 2, &edge("e1"), b"x".to_vec())
572            .expect("ok");
573        assert!(!ok);
574    }
575
576    #[test]
577    fn marker_propagator_reports_unknown_operator() {
578        let prop = MarkerPropagator::new();
579        let err = prop
580            .on_marker(&op("ghost"), &edge("e1"), &Marker::new(1))
581            .expect_err("should fail");
582        assert!(matches!(err, CheckpointError::UnknownOperator(_)));
583    }
584
585    #[test]
586    fn marker_propagator_reports_unknown_edge() {
587        let prop = MarkerPropagator::new();
588        prop.register_operator(op("op1"), ["e1"]);
589        let err = prop
590            .on_marker(&op("op1"), &edge("e2"), &Marker::new(1))
591            .expect_err("should fail");
592        assert!(matches!(err, CheckpointError::UnknownEdge { .. }));
593    }
594
595    #[tokio::test]
596    async fn controller_drives_full_round() {
597        let propagator = Arc::new(MarkerPropagator::new());
598        let store = Arc::new(InMemoryCheckpointStore::new());
599        let controller = CheckpointController::new(
600            CheckpointControllerConfig::default(),
601            propagator.clone(),
602            store.clone(),
603        );
604        controller.register_operator(op("op-a"));
605        controller.register_operator(op("op-b"));
606        propagator.register_operator(op("op-a"), ["src"]);
607        propagator.register_operator(op("op-b"), ["a"]);
608        let marker = controller.open();
609
610        propagator
611            .on_marker(&op("op-a"), &edge("src"), &marker)
612            .expect("ok");
613        let snap_a = OperatorSnapshot {
614            operator_id: op("op-a"),
615            checkpoint_id: marker.checkpoint_id,
616            state_blob: vec![1, 2, 3],
617            channel_logs: HashMap::new(),
618            completed_at_ms: now_ms(),
619        };
620        let done = controller.commit_snapshot(snap_a).await.expect("ok");
621        assert!(!done);
622
623        propagator
624            .on_marker(&op("op-b"), &edge("a"), &marker)
625            .expect("ok");
626        let snap_b = OperatorSnapshot {
627            operator_id: op("op-b"),
628            checkpoint_id: marker.checkpoint_id,
629            state_blob: vec![9, 9],
630            channel_logs: HashMap::new(),
631            completed_at_ms: now_ms(),
632        };
633        let done = controller.commit_snapshot(snap_b).await.expect("ok");
634        assert!(done);
635
636        let prog = controller.progress(marker.checkpoint_id).expect("progress");
637        assert!(prog.is_complete());
638        let latest = controller.latest_committed().await.expect("ok");
639        assert_eq!(latest, Some(marker.checkpoint_id));
640    }
641
642    #[tokio::test]
643    async fn store_round_trip() {
644        let store = InMemoryCheckpointStore::new();
645        let snap = OperatorSnapshot {
646            operator_id: op("op1"),
647            checkpoint_id: 7,
648            state_blob: vec![1],
649            channel_logs: HashMap::new(),
650            completed_at_ms: now_ms(),
651        };
652        store.put(snap.clone()).await.expect("put");
653        let back = store.get(&snap.operator_id, 7).await.expect("get");
654        assert_eq!(back.expect("hit").operator_id, op("op1"));
655        let latest = store.latest().await.expect("latest");
656        assert_eq!(latest, Some(7));
657        assert_eq!(store.len(), 1);
658    }
659
660    #[test]
661    fn controller_opens_unique_ids() {
662        let propagator = Arc::new(MarkerPropagator::new());
663        let store: Arc<dyn CheckpointStore> = Arc::new(InMemoryCheckpointStore::new());
664        let controller =
665            CheckpointController::new(CheckpointControllerConfig::default(), propagator, store);
666        let m1 = controller.open();
667        let m2 = controller.open();
668        assert_ne!(m1.checkpoint_id, m2.checkpoint_id);
669        assert_eq!(controller.opened_rounds(), 2);
670    }
671}