Skip to main content

ursula_stream/state_machine/
persist.rs

1//! Snapshot / restore / integrity serialization for the Raft state machine.
2
3use super::BucketStreamId;
4use super::ColdChunkRef;
5use super::ColdGcQueue;
6use super::HashMap;
7use super::HotBuffer;
8use super::HotPayloadSegment;
9use super::ObjectPayloadRef;
10use super::ProducerReceipt;
11use super::ProducerSnapshot;
12use super::ProducerState;
13use super::StreamColdState;
14use super::StreamErrorCode;
15use super::StreamIntegrity;
16use super::StreamMessageRecord;
17use super::StreamResponse;
18use super::StreamSlot;
19use super::StreamSnapshot;
20use super::StreamSnapshotEntry;
21use super::StreamSnapshotError;
22use super::StreamStateMachine;
23use super::compare_stream_ids;
24use super::normalize_stream_attrs;
25
26impl StreamStateMachine {
27    pub fn integrity_snapshot(
28        &self,
29        stream_id: &BucketStreamId,
30    ) -> Result<crate::integrity::StreamIntegritySnapshot, StreamResponse> {
31        let Some(slot) = self.stream_slot(stream_id) else {
32            return Err(StreamResponse::error(
33                StreamErrorCode::StreamNotFound,
34                format!("stream '{stream_id}' does not exist"),
35            ));
36        };
37        Ok(slot.integrity.snapshot(
38            self.earliest_retained_offset(stream_id),
39            slot.metadata.tail_offset,
40        ))
41    }
42
43    pub fn snapshot(&self) -> StreamSnapshot {
44        let mut buckets = self.buckets.iter().cloned().collect::<Vec<_>>();
45        buckets.sort();
46        let mut erased_buckets = self.erased_buckets.iter().cloned().collect::<Vec<_>>();
47        erased_buckets.sort();
48
49        let mut streams = self
50            .registry
51            .slots()
52            .map(|slot| {
53                let metadata = slot.metadata.clone();
54                let stream_id = metadata.stream_id.clone();
55                let tail_offset = metadata.tail_offset;
56                let payload = slot.hot_buffer.payload();
57                let producer_states = producer_snapshot(&slot.producers);
58                StreamSnapshotEntry {
59                    metadata,
60                    attrs: slot.attrs.clone(),
61                    hot_start_offset: self.hot_start_offset(&stream_id),
62                    payload,
63                    hot_segments: slot.hot_buffer.hot_segments(),
64                    cold_frontier_offset: self.cold_frontier_offset(
65                        &stream_id,
66                        self.earliest_retained_offset(&stream_id),
67                    ),
68                    cold_index_generation: slot.cold.cold_generation(),
69                    cold_chunks: slot.cold.cold_chunks().to_vec(),
70                    external_segments: slot.cold.external_segments().to_vec(),
71                    message_records: slot.message_records.clone(),
72                    record_index: slot.record_index.clone(),
73                    integrity: slot
74                        .integrity
75                        .snapshot(self.earliest_retained_offset(&stream_id), tail_offset),
76                    retained_offset: Some(slot.retained_offset),
77                    visible_snapshot: slot.visible_snapshot.clone(),
78                    producer_states,
79                }
80            })
81            .collect::<Vec<_>>();
82        streams.sort_by(|left, right| {
83            compare_stream_ids(&left.metadata.stream_id, &right.metadata.stream_id)
84        });
85
86        let mut shared_cold_object_owners = self
87            .shared_cold_object_owners
88            .iter()
89            .filter(|(path, _)| self.shared_cold_object_refs.contains_key(*path))
90            .map(|(path, owners)| {
91                let mut bucket_ids = owners.iter().cloned().collect::<Vec<_>>();
92                bucket_ids.sort();
93                crate::SharedColdObjectOwnersSnapshot {
94                    s3_path: path.clone(),
95                    bucket_ids,
96                }
97            })
98            .collect::<Vec<_>>();
99        shared_cold_object_owners.sort_by(|left, right| left.s3_path.cmp(&right.s3_path));
100
101        StreamSnapshot {
102            buckets,
103            erased_buckets,
104            streams,
105            pending_cold_gc: self.cold_gc.entries().cloned().collect(),
106            next_cold_gc_seq: self.cold_gc.next_seq(),
107            shared_cold_object_owners,
108            bucket_usage: self.bucket_usage_report(),
109            bucket_quotas: self.bucket_quota_report(),
110        }
111    }
112
113    /// Applies a whole-group state import (`StreamCommand::ImportSnapshot`).
114    ///
115    /// Restore-only by design: importing over live state would silently merge
116    /// two histories, so a non-empty group fails closed with
117    /// [`StreamErrorCode::ImportConflict`] and an invalid payload with
118    /// [`StreamErrorCode::ImportInvalid`].
119    pub(crate) fn import_snapshot(&mut self, snapshot: StreamSnapshot) -> StreamResponse {
120        if !self.buckets.is_empty() || !self.erased_buckets.is_empty() || !self.registry.is_empty()
121        {
122            return StreamResponse::error(
123                StreamErrorCode::ImportConflict,
124                format!(
125                    "group already holds {} bucket(s); snapshot import requires an empty group",
126                    self.buckets.len()
127                ),
128            );
129        }
130        let buckets = u64::try_from(snapshot.buckets.len()).unwrap_or(u64::MAX);
131        let streams = u64::try_from(snapshot.streams.len()).unwrap_or(u64::MAX);
132        match Self::restore(snapshot) {
133            Ok(restored) => {
134                *self = restored;
135                StreamResponse::SnapshotImported { buckets, streams }
136            }
137            Err(error) => StreamResponse::error(
138                StreamErrorCode::ImportInvalid,
139                format!("snapshot import failed validation: {error}"),
140            ),
141        }
142    }
143
144    pub fn restore(snapshot: StreamSnapshot) -> Result<Self, StreamSnapshotError> {
145        let mut machine = Self::default();
146        for bucket_id in &snapshot.buckets {
147            if !machine.buckets.insert(bucket_id.clone()) {
148                return Err(StreamSnapshotError::DuplicateBucket(bucket_id.clone()));
149            }
150        }
151        for bucket_id in &snapshot.erased_buckets {
152            if machine.buckets.contains(bucket_id) {
153                return Err(StreamSnapshotError::ActiveBucketErased(bucket_id.clone()));
154            }
155            if !machine.erased_buckets.insert(bucket_id.clone()) {
156                return Err(StreamSnapshotError::DuplicateErasedBucket(
157                    bucket_id.clone(),
158                ));
159            }
160        }
161
162        for entry in snapshot.streams {
163            let stream_id = entry.metadata.stream_id.clone();
164            if !machine.buckets.contains(&stream_id.bucket_id) {
165                return Err(StreamSnapshotError::MissingBucket(stream_id));
166            }
167            if let Some(snapshot) = entry.visible_snapshot.as_ref()
168                && snapshot.offset > entry.metadata.tail_offset
169            {
170                return Err(StreamSnapshotError::SnapshotOffsetOutOfRange {
171                    stream_id,
172                    snapshot_offset: snapshot.offset,
173                    tail_offset: entry.metadata.tail_offset,
174                });
175            }
176            let retained_offset = entry.retained_offset.unwrap_or_else(|| {
177                entry
178                    .visible_snapshot
179                    .as_ref()
180                    .map(|snapshot| snapshot.offset)
181                    .unwrap_or(0)
182            });
183            if retained_offset > entry.metadata.tail_offset {
184                return Err(StreamSnapshotError::SnapshotOffsetOutOfRange {
185                    stream_id,
186                    snapshot_offset: retained_offset,
187                    tail_offset: entry.metadata.tail_offset,
188                });
189            }
190            if let Some(record_index) = entry.record_index.as_ref()
191                && record_index
192                    .validate(retained_offset, entry.metadata.tail_offset)
193                    .is_err()
194            {
195                return Err(StreamSnapshotError::RecordBoundaryMismatch { stream_id });
196            }
197            let hot_segments = if entry.hot_segments.is_empty() && !entry.payload.is_empty() {
198                vec![HotPayloadSegment {
199                    start_offset: entry.hot_start_offset,
200                    end_offset: entry.metadata.tail_offset,
201                    payload_start: 0,
202                    payload_end: entry.payload.len(),
203                }]
204            } else {
205                entry.hot_segments
206            };
207            if !hot_segments_match_payload(&hot_segments, entry.payload.len())
208                || !payload_sources_cover_retained_suffix(
209                    entry.cold_frontier_offset,
210                    &entry.cold_chunks,
211                    &entry.external_segments,
212                    &hot_segments,
213                    retained_offset,
214                    entry.metadata.tail_offset,
215                )
216            {
217                return Err(StreamSnapshotError::PayloadLengthMismatch {
218                    stream_id,
219                    tail_offset: entry.metadata.tail_offset,
220                    payload_len: entry.payload.len(),
221                });
222            }
223            if !message_records_cover_retained_suffix(
224                &entry.message_records,
225                retained_offset,
226                entry.metadata.tail_offset,
227            ) {
228                return Err(StreamSnapshotError::MessageBoundaryMismatch { stream_id });
229            }
230            let integrity = StreamIntegrity::restore(entry.integrity).ok_or_else(|| {
231                StreamSnapshotError::IntegrityMismatch {
232                    stream_id: stream_id.clone(),
233                }
234            })?;
235            if machine.registry.contains_key(&stream_id) {
236                return Err(StreamSnapshotError::DuplicateStream(stream_id));
237            }
238            let producer_states = restore_producer_states(&stream_id, entry.producer_states)?;
239            let visible_snapshot = entry.visible_snapshot.map(|mut snapshot| {
240                if snapshot.digest.is_empty() {
241                    snapshot.digest =
242                        super::snapshot_digest(&snapshot.content_type, &snapshot.payload);
243                }
244                snapshot
245            });
246            let shared_cold_paths = entry
247                .cold_chunks
248                .iter()
249                .filter(|chunk| chunk.shared_object)
250                .map(|chunk| chunk.s3_path.clone())
251                .collect::<Vec<_>>();
252            let slot = StreamSlot {
253                metadata: entry.metadata,
254                attrs: normalize_stream_attrs(entry.attrs),
255                hot_buffer: HotBuffer::from_snapshot(entry.payload, &hot_segments),
256                cold: StreamColdState::restore(
257                    entry.cold_frontier_offset,
258                    entry.cold_index_generation,
259                    entry.cold_chunks,
260                    entry.external_segments,
261                ),
262                message_records: entry.message_records,
263                record_index: entry.record_index,
264                integrity,
265                retained_offset,
266                visible_snapshot,
267                producers: producer_states,
268            };
269            if machine.insert_stream_slot(slot).is_none() {
270                return Err(StreamSnapshotError::DuplicateStream(stream_id));
271            }
272            for path in shared_cold_paths {
273                machine.retain_shared_cold_object(&path, &stream_id.bucket_id);
274            }
275        }
276
277        for record in snapshot.shared_cold_object_owners {
278            if machine
279                .shared_cold_object_refs
280                .contains_key(&record.s3_path)
281            {
282                machine
283                    .shared_cold_object_owners
284                    .entry(record.s3_path)
285                    .or_default()
286                    .extend(record.bucket_ids);
287            }
288        }
289
290        machine.cold_gc =
291            ColdGcQueue::from_parts(snapshot.pending_cold_gc, snapshot.next_cold_gc_seq);
292
293        // Usage restore: gauges are recomputed from the restored slots so a
294        // snapshot can never carry gauge drift forward; only the monotonic
295        // counters are taken from the snapshot. Legacy snapshots without the
296        // field restart the monotonic counters from the recomputed gauges.
297        let mut recomputed: HashMap<String, super::BucketUsage> = HashMap::new();
298        for slot in machine.registry.slots() {
299            let usage = recomputed
300                .entry(slot.metadata.stream_id.bucket_id.clone())
301                .or_default();
302            usage.stream_count = usage.stream_count.saturating_add(1);
303            usage.retained_bytes = usage.retained_bytes.saturating_add(
304                slot.metadata
305                    .tail_offset
306                    .saturating_sub(slot.retained_offset),
307            );
308        }
309        for persisted in snapshot.bucket_usage {
310            let usage = recomputed.entry(persisted.bucket_id).or_default();
311            usage.committed_append_bytes = persisted.usage.committed_append_bytes;
312            usage.committed_records = persisted.usage.committed_records;
313            usage.committed_write_units = persisted.usage.committed_write_units;
314        }
315        for usage in recomputed.values_mut() {
316            if usage.committed_append_bytes == 0 {
317                usage.committed_append_bytes = usage.retained_bytes;
318            }
319        }
320        machine.bucket_usage = recomputed;
321
322        for persisted in snapshot.bucket_quotas {
323            if persisted.quota.is_unlimited() {
324                continue;
325            }
326            machine
327                .bucket_quotas
328                .insert(persisted.bucket_id, persisted.quota);
329        }
330
331        Ok(machine)
332    }
333}
334
335fn producer_snapshot(states: &HashMap<String, ProducerState>) -> Vec<ProducerSnapshot> {
336    let mut producer_states = states
337        .iter()
338        .map(|(producer_id, state)| ProducerSnapshot {
339            producer_id: producer_id.clone(),
340            producer_epoch: state.producer_epoch,
341            producer_seq: state.producer_seq,
342            last_start_offset: state.last_start_offset,
343            last_next_offset: state.last_next_offset,
344            last_closed: state.last_closed,
345            last_items: state.last_items.clone(),
346            receipts: state.receipts.clone(),
347        })
348        .collect::<Vec<_>>();
349    producer_states.sort_by(|left, right| left.producer_id.cmp(&right.producer_id));
350    producer_states
351}
352
353fn restore_producer_states(
354    stream_id: &BucketStreamId,
355    snapshots: Vec<ProducerSnapshot>,
356) -> Result<HashMap<String, ProducerState>, StreamSnapshotError> {
357    let mut states = HashMap::with_capacity(snapshots.len());
358    for snapshot in snapshots {
359        let receipts = if snapshot.receipts.is_empty() {
360            vec![ProducerReceipt {
361                producer_seq: snapshot.producer_seq,
362                start_offset: snapshot.last_start_offset,
363                next_offset: snapshot.last_next_offset,
364                closed: snapshot.last_closed,
365                items: snapshot.last_items.clone(),
366            }]
367        } else {
368            snapshot.receipts
369        };
370        if states
371            .insert(snapshot.producer_id.clone(), ProducerState {
372                producer_epoch: snapshot.producer_epoch,
373                producer_seq: snapshot.producer_seq,
374                last_start_offset: snapshot.last_start_offset,
375                last_next_offset: snapshot.last_next_offset,
376                last_closed: snapshot.last_closed,
377                last_items: snapshot.last_items,
378                receipts,
379            })
380            .is_some()
381        {
382            return Err(StreamSnapshotError::DuplicateProducer {
383                stream_id: stream_id.clone(),
384                producer_id: snapshot.producer_id,
385            });
386        }
387    }
388    Ok(states)
389}
390
391fn valid_cold_chunk_ref(chunk: &ColdChunkRef) -> bool {
392    let logical_len = chunk.end_offset.saturating_sub(chunk.start_offset);
393    chunk.end_offset > chunk.start_offset
394        && !chunk.s3_path.trim().is_empty()
395        && chunk
396            .object_offset
397            .checked_add(logical_len)
398            .is_some_and(|end| end <= chunk.object_size)
399}
400
401fn valid_object_payload_ref(object: &ObjectPayloadRef) -> bool {
402    let logical_len = object.end_offset.saturating_sub(object.start_offset);
403    object.end_offset > object.start_offset
404        && !object.s3_path.trim().is_empty()
405        && object
406            .object_offset
407            .checked_add(logical_len)
408            .is_some_and(|end| end <= object.object_size)
409}
410
411fn hot_segments_match_payload(segments: &[HotPayloadSegment], payload_len: usize) -> bool {
412    let mut expected_payload_start = 0;
413    for segment in segments {
414        if segment.end_offset <= segment.start_offset
415            || segment.payload_start != expected_payload_start
416            || segment.payload_end <= segment.payload_start
417            || segment.payload_end > payload_len
418        {
419            return false;
420        }
421        let Ok(logical_len) = usize::try_from(segment.end_offset - segment.start_offset) else {
422            return false;
423        };
424        if logical_len != segment.payload_end - segment.payload_start {
425            return false;
426        }
427        expected_payload_start = segment.payload_end;
428    }
429    expected_payload_start == payload_len
430}
431
432fn payload_sources_cover_retained_suffix(
433    cold_frontier_offset: u64,
434    cold_chunks: &[ColdChunkRef],
435    external_segments: &[ObjectPayloadRef],
436    hot_segments: &[HotPayloadSegment],
437    retained_offset: u64,
438    tail_offset: u64,
439) -> bool {
440    if tail_offset < retained_offset {
441        return false;
442    }
443    let mut ranges =
444        Vec::with_capacity(1 + cold_chunks.len() + external_segments.len() + hot_segments.len());
445    if cold_frontier_offset > retained_offset {
446        ranges.push((retained_offset, cold_frontier_offset));
447    }
448    for chunk in cold_chunks {
449        if !valid_cold_chunk_ref(chunk) {
450            return false;
451        }
452        ranges.push((chunk.start_offset, chunk.end_offset));
453    }
454    for object in external_segments {
455        if !valid_object_payload_ref(object) {
456            return false;
457        }
458        ranges.push((object.start_offset, object.end_offset));
459    }
460    for segment in hot_segments {
461        if segment.end_offset <= segment.start_offset {
462            return false;
463        }
464        ranges.push((segment.start_offset, segment.end_offset));
465    }
466    ranges.sort_unstable();
467
468    let mut expected_start = retained_offset;
469    for (start_offset, end_offset) in ranges {
470        if end_offset <= expected_start {
471            continue;
472        }
473        if start_offset > expected_start {
474            return false;
475        }
476        expected_start = end_offset;
477        if expected_start >= tail_offset {
478            return true;
479        }
480    }
481    expected_start == tail_offset
482}
483
484pub(super) fn message_records_cover_retained_suffix(
485    records: &[StreamMessageRecord],
486    retained_offset: u64,
487    tail_offset: u64,
488) -> bool {
489    let mut expected_start = retained_offset;
490    for record in records {
491        if record.start_offset != expected_start || record.end_offset <= record.start_offset {
492            return false;
493        }
494        expected_start = record.end_offset;
495    }
496    expected_start == tail_offset
497}