1use 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
40pub type CheckpointId = u64;
44
45pub type OperatorId = String;
47
48pub type InputEdgeId = String;
50
51#[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 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#[derive(Debug, Clone, Serialize, Deserialize)]
80pub struct OperatorSnapshot {
81 pub operator_id: OperatorId,
82 pub checkpoint_id: CheckpointId,
83 pub state_blob: Vec<u8>,
85 pub channel_logs: HashMap<InputEdgeId, Vec<Vec<u8>>>,
88 pub completed_at_ms: u64,
90}
91
92#[derive(Debug, Error)]
96pub enum CheckpointError {
97 #[error("unknown operator: {0}")]
99 UnknownOperator(OperatorId),
100 #[error("unknown edge {edge} on operator {op}")]
102 UnknownEdge { op: OperatorId, edge: InputEdgeId },
103 #[error("store error: {0}")]
105 Store(String),
106}
107
108pub type CheckpointResult<T> = std::result::Result<T, CheckpointError>;
110
111#[async_trait]
115pub trait CheckpointStore: Send + Sync {
116 async fn put(&self, snap: OperatorSnapshot) -> CheckpointResult<()>;
118
119 async fn get(
121 &self,
122 operator: &OperatorId,
123 checkpoint: CheckpointId,
124 ) -> CheckpointResult<Option<OperatorSnapshot>>;
125
126 async fn latest(&self) -> CheckpointResult<Option<CheckpointId>>;
129}
130
131pub 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 pub fn new() -> Self {
148 Self {
149 inner: RwLock::new(HashMap::new()),
150 latest: RwLock::new(None),
151 }
152 }
153
154 pub fn len(&self) -> usize {
156 self.inner.read().len()
157 }
158
159 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#[derive(Debug, Default)]
200struct OperatorMarkerState {
201 seen_on: HashSet<InputEdgeId>,
203 expected: HashSet<InputEdgeId>,
205 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#[derive(Debug, Clone, PartialEq, Eq)]
221pub enum MarkerPropagatorEvent {
222 StartSnapshot,
225 EdgeClosed,
228 Completed,
231}
232
233pub 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 pub fn new() -> Self {
248 Self {
249 state: RwLock::new(HashMap::new()),
250 expected_edges: RwLock::new(HashMap::new()),
251 }
252 }
253
254 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 pub fn is_registered(&self, op: &OperatorId) -> bool {
267 self.expected_edges.read().contains_key(op)
268 }
269
270 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 return Ok(false);
288 }
289 if st.seen_on.contains(edge) {
290 return Ok(false);
292 }
293 st.channel_logs
294 .entry(edge.clone())
295 .or_default()
296 .push(payload);
297 Ok(true)
298 }
299
300 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 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 pub fn reset(&self, op: &OperatorId) {
362 let mut states = self.state.write();
363 states.retain(|(o, _), _| o != op);
364 }
365}
366
367#[derive(Debug, Clone, Serialize, Deserialize)]
371pub struct CheckpointControllerConfig {
372 pub interval: Duration,
374 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#[derive(Debug, Clone, Default, Serialize, Deserialize)]
390pub struct CheckpointProgress {
391 pub checkpoint_id: CheckpointId,
392 pub committed: HashSet<OperatorId>,
394 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
406pub struct CheckpointController {
408 config: CheckpointControllerConfig,
409 propagator: Arc<MarkerPropagator>,
410 store: Arc<dyn CheckpointStore>,
411 next_id: AtomicU64,
412 operators: RwLock<HashSet<OperatorId>>,
414 progress: RwLock<HashMap<CheckpointId, CheckpointProgress>>,
416}
417
418impl CheckpointController {
419 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 pub fn config(&self) -> &CheckpointControllerConfig {
437 &self.config
438 }
439
440 pub fn propagator(&self) -> &Arc<MarkerPropagator> {
443 &self.propagator
444 }
445
446 pub fn register_operator(&self, op: OperatorId) {
448 self.operators.write().insert(op);
449 }
450
451 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 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 pub fn progress(&self, cp: CheckpointId) -> Option<CheckpointProgress> {
489 self.progress.read().get(&cp).cloned()
490 }
491
492 pub fn forget(&self, cp: CheckpointId) {
494 self.progress.write().remove(&cp);
495 }
496
497 pub async fn latest_committed(&self) -> CheckpointResult<Option<CheckpointId>> {
499 self.store.latest().await
500 }
501
502 pub fn opened_rounds(&self) -> u64 {
504 self.next_id.load(Ordering::Relaxed).saturating_sub(1)
505 }
506
507 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#[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 let ok = prop
557 .record_inflight(&op("op1"), 2, &edge("e2"), b"early".to_vec())
558 .expect("ok");
559 assert!(!ok);
560 let _ = prop
562 .on_marker(&op("op1"), &edge("e1"), &marker)
563 .expect("ok");
564 let ok = prop
566 .record_inflight(&op("op1"), 2, &edge("e2"), b"after-marker".to_vec())
567 .expect("ok");
568 assert!(ok);
569 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}