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