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::ProducerSnapshot;
11use super::ProducerState;
12use super::StreamColdState;
13use super::StreamErrorCode;
14use super::StreamIntegrity;
15use super::StreamMessageRecord;
16use super::StreamResponse;
17use super::StreamSlot;
18use super::StreamSnapshot;
19use super::StreamSnapshotEntry;
20use super::StreamSnapshotError;
21use super::StreamStateMachine;
22use super::compare_stream_ids;
23use super::is_soft_deleted;
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        if is_soft_deleted(&slot.metadata) {
38            return Err(StreamResponse::error(
39                StreamErrorCode::StreamGone,
40                format!("stream '{stream_id}' is gone"),
41            ));
42        }
43        Ok(slot.integrity.snapshot(
44            self.earliest_retained_offset(stream_id),
45            slot.metadata.tail_offset,
46        ))
47    }
48
49    pub fn snapshot(&self) -> StreamSnapshot {
50        let mut buckets = self.buckets.iter().cloned().collect::<Vec<_>>();
51        buckets.sort();
52
53        let mut streams = self
54            .registry
55            .slots()
56            .map(|slot| {
57                let metadata = slot.metadata.clone();
58                let stream_id = metadata.stream_id.clone();
59                let tail_offset = metadata.tail_offset;
60                let payload = slot.hot_buffer.payload();
61                let producer_states = producer_snapshot(&slot.producers);
62                StreamSnapshotEntry {
63                    metadata,
64                    attrs: slot.attrs.clone(),
65                    hot_start_offset: self.hot_start_offset(&stream_id),
66                    payload,
67                    hot_segments: slot.hot_buffer.hot_segments(),
68                    cold_frontier_offset: self.cold_frontier_offset(
69                        &stream_id,
70                        self.earliest_retained_offset(&stream_id),
71                    ),
72                    cold_index_generation: slot.cold.cold_generation(),
73                    cold_chunks: slot.cold.cold_chunks().to_vec(),
74                    external_segments: slot.cold.external_segments().to_vec(),
75                    message_records: slot.message_records.clone(),
76                    integrity: slot
77                        .integrity
78                        .snapshot(self.earliest_retained_offset(&stream_id), tail_offset),
79                    visible_snapshot: slot.visible_snapshot.clone(),
80                    producer_states,
81                }
82            })
83            .collect::<Vec<_>>();
84        streams.sort_by(|left, right| {
85            compare_stream_ids(&left.metadata.stream_id, &right.metadata.stream_id)
86        });
87
88        StreamSnapshot {
89            buckets,
90            streams,
91            pending_cold_gc: self.cold_gc.entries().cloned().collect(),
92            next_cold_gc_seq: self.cold_gc.next_seq(),
93        }
94    }
95
96    pub fn restore(snapshot: StreamSnapshot) -> Result<Self, StreamSnapshotError> {
97        let mut machine = Self::default();
98        for bucket_id in snapshot.buckets {
99            if !machine.buckets.insert(bucket_id.clone()) {
100                return Err(StreamSnapshotError::DuplicateBucket(bucket_id));
101            }
102        }
103
104        for entry in snapshot.streams {
105            let stream_id = entry.metadata.stream_id.clone();
106            if !machine.buckets.contains(&stream_id.bucket_id) {
107                return Err(StreamSnapshotError::MissingBucket(stream_id));
108            }
109            if let Some(snapshot) = entry.visible_snapshot.as_ref()
110                && snapshot.offset > entry.metadata.tail_offset
111            {
112                return Err(StreamSnapshotError::SnapshotOffsetOutOfRange {
113                    stream_id,
114                    snapshot_offset: snapshot.offset,
115                    tail_offset: entry.metadata.tail_offset,
116                });
117            }
118            let retained_offset = entry
119                .visible_snapshot
120                .as_ref()
121                .map(|snapshot| snapshot.offset)
122                .unwrap_or(0);
123            let hot_segments = if entry.hot_segments.is_empty() && !entry.payload.is_empty() {
124                vec![HotPayloadSegment {
125                    start_offset: entry.hot_start_offset,
126                    end_offset: entry.metadata.tail_offset,
127                    payload_start: 0,
128                    payload_end: entry.payload.len(),
129                }]
130            } else {
131                entry.hot_segments
132            };
133            if !hot_segments_match_payload(&hot_segments, entry.payload.len())
134                || !payload_sources_cover_retained_suffix(
135                    entry.cold_frontier_offset,
136                    &entry.cold_chunks,
137                    &entry.external_segments,
138                    &hot_segments,
139                    retained_offset,
140                    entry.metadata.tail_offset,
141                )
142            {
143                return Err(StreamSnapshotError::PayloadLengthMismatch {
144                    stream_id,
145                    tail_offset: entry.metadata.tail_offset,
146                    payload_len: entry.payload.len(),
147                });
148            }
149            if !message_records_cover_retained_suffix(
150                &entry.message_records,
151                retained_offset,
152                entry.metadata.tail_offset,
153            ) {
154                return Err(StreamSnapshotError::MessageBoundaryMismatch { stream_id });
155            }
156            let integrity = StreamIntegrity::restore(entry.integrity).ok_or_else(|| {
157                StreamSnapshotError::IntegrityMismatch {
158                    stream_id: stream_id.clone(),
159                }
160            })?;
161            if machine.registry.contains_key(&stream_id) {
162                return Err(StreamSnapshotError::DuplicateStream(stream_id));
163            }
164            let producer_states = restore_producer_states(&stream_id, entry.producer_states)?;
165            let visible_snapshot = entry.visible_snapshot;
166            let slot = StreamSlot {
167                metadata: entry.metadata,
168                attrs: normalize_stream_attrs(entry.attrs),
169                hot_buffer: HotBuffer::from_snapshot(entry.payload, &hot_segments),
170                cold: StreamColdState::restore(
171                    entry.cold_frontier_offset,
172                    entry.cold_index_generation,
173                    entry.cold_chunks,
174                    entry.external_segments,
175                ),
176                message_records: entry.message_records,
177                integrity,
178                visible_snapshot,
179                producers: producer_states,
180            };
181            if machine.insert_stream_slot(slot).is_none() {
182                return Err(StreamSnapshotError::DuplicateStream(stream_id));
183            }
184        }
185
186        machine.cold_gc =
187            ColdGcQueue::from_parts(snapshot.pending_cold_gc, snapshot.next_cold_gc_seq);
188
189        Ok(machine)
190    }
191}
192
193fn producer_snapshot(states: &HashMap<String, ProducerState>) -> Vec<ProducerSnapshot> {
194    let mut producer_states = states
195        .iter()
196        .map(|(producer_id, state)| ProducerSnapshot {
197            producer_id: producer_id.clone(),
198            producer_epoch: state.producer_epoch,
199            producer_seq: state.producer_seq,
200            last_start_offset: state.last_start_offset,
201            last_next_offset: state.last_next_offset,
202            last_closed: state.last_closed,
203            last_items: state.last_items.clone(),
204        })
205        .collect::<Vec<_>>();
206    producer_states.sort_by(|left, right| left.producer_id.cmp(&right.producer_id));
207    producer_states
208}
209
210fn restore_producer_states(
211    stream_id: &BucketStreamId,
212    snapshots: Vec<ProducerSnapshot>,
213) -> Result<HashMap<String, ProducerState>, StreamSnapshotError> {
214    let mut states = HashMap::with_capacity(snapshots.len());
215    for snapshot in snapshots {
216        if states
217            .insert(snapshot.producer_id.clone(), ProducerState {
218                producer_epoch: snapshot.producer_epoch,
219                producer_seq: snapshot.producer_seq,
220                last_start_offset: snapshot.last_start_offset,
221                last_next_offset: snapshot.last_next_offset,
222                last_closed: snapshot.last_closed,
223                last_items: snapshot.last_items,
224            })
225            .is_some()
226        {
227            return Err(StreamSnapshotError::DuplicateProducer {
228                stream_id: stream_id.clone(),
229                producer_id: snapshot.producer_id,
230            });
231        }
232    }
233    Ok(states)
234}
235
236fn valid_cold_chunk_ref(chunk: &ColdChunkRef) -> bool {
237    chunk.end_offset > chunk.start_offset
238        && !chunk.s3_path.trim().is_empty()
239        && chunk.object_size >= chunk.end_offset - chunk.start_offset
240}
241
242fn valid_object_payload_ref(object: &ObjectPayloadRef) -> bool {
243    object.end_offset > object.start_offset
244        && !object.s3_path.trim().is_empty()
245        && object.object_size >= object.end_offset - object.start_offset
246}
247
248fn hot_segments_match_payload(segments: &[HotPayloadSegment], payload_len: usize) -> bool {
249    let mut expected_payload_start = 0;
250    for segment in segments {
251        if segment.end_offset <= segment.start_offset
252            || segment.payload_start != expected_payload_start
253            || segment.payload_end <= segment.payload_start
254            || segment.payload_end > payload_len
255        {
256            return false;
257        }
258        let Ok(logical_len) = usize::try_from(segment.end_offset - segment.start_offset) else {
259            return false;
260        };
261        if logical_len != segment.payload_end - segment.payload_start {
262            return false;
263        }
264        expected_payload_start = segment.payload_end;
265    }
266    expected_payload_start == payload_len
267}
268
269fn payload_sources_cover_retained_suffix(
270    cold_frontier_offset: u64,
271    cold_chunks: &[ColdChunkRef],
272    external_segments: &[ObjectPayloadRef],
273    hot_segments: &[HotPayloadSegment],
274    retained_offset: u64,
275    tail_offset: u64,
276) -> bool {
277    if tail_offset < retained_offset {
278        return false;
279    }
280    let mut ranges =
281        Vec::with_capacity(1 + cold_chunks.len() + external_segments.len() + hot_segments.len());
282    if cold_frontier_offset > retained_offset {
283        ranges.push((retained_offset, cold_frontier_offset));
284    }
285    for chunk in cold_chunks {
286        if !valid_cold_chunk_ref(chunk) {
287            return false;
288        }
289        ranges.push((chunk.start_offset, chunk.end_offset));
290    }
291    for object in external_segments {
292        if !valid_object_payload_ref(object) {
293            return false;
294        }
295        ranges.push((object.start_offset, object.end_offset));
296    }
297    for segment in hot_segments {
298        if segment.end_offset <= segment.start_offset {
299            return false;
300        }
301        ranges.push((segment.start_offset, segment.end_offset));
302    }
303    ranges.sort_unstable();
304
305    let mut expected_start = retained_offset;
306    for (start_offset, end_offset) in ranges {
307        if end_offset <= expected_start {
308            continue;
309        }
310        if start_offset > expected_start {
311            return false;
312        }
313        expected_start = end_offset;
314        if expected_start >= tail_offset {
315            return true;
316        }
317    }
318    expected_start == tail_offset
319}
320
321pub(super) fn message_records_cover_retained_suffix(
322    records: &[StreamMessageRecord],
323    retained_offset: u64,
324    tail_offset: u64,
325) -> bool {
326    let mut expected_start = retained_offset;
327    for record in records {
328        if record.start_offset != expected_start || record.end_offset <= record.start_offset {
329            return false;
330        }
331        expected_start = record.end_offset;
332    }
333    expected_start == tail_offset
334}