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