1use crate::metadata::TableMetadata;
11use crate::{Result, TableError};
12use cobble::{DbCoordinator, GlobalSnapshotManifest, ShardSnapshotMetadata};
13use std::collections::BTreeMap;
14use std::sync::{Arc, Mutex};
15
16pub struct TableSnapshotCommitter {
18 coordinator: Arc<DbCoordinator>,
19 total_buckets: u32,
20 max_pending_commits: usize,
21 state: Mutex<CommitterState>,
22}
23
24struct CommitterState {
25 latest_completed_commit_id: Option<u64>,
26 pending: BTreeMap<u64, PendingCommit>,
27}
28
29#[derive(Default)]
30struct PendingCommit {
31 shards: BTreeMap<String, ShardSnapshotMetadata>,
32 table_metadata: Option<BTreeMap<String, TableMetadata>>,
33 prepared_snapshot: Option<GlobalSnapshotManifest>,
34}
35
36impl TableSnapshotCommitter {
37 pub fn new(
38 coordinator: Arc<DbCoordinator>,
39 total_buckets: u32,
40 max_pending_commits: usize,
41 ) -> Result<Self> {
42 if total_buckets == 0 || total_buckets > u16::MAX as u32 + 1 {
43 return Err(cobble::Error::ConfigError(
44 "total_buckets must be in range 1..=65536".to_string(),
45 )
46 .into());
47 }
48 if max_pending_commits == 0 {
49 return Err(cobble::Error::ConfigError(
50 "max_pending_commits must be positive".to_string(),
51 )
52 .into());
53 }
54 Ok(Self {
55 coordinator,
56 total_buckets,
57 max_pending_commits,
58 state: Mutex::new(CommitterState {
59 latest_completed_commit_id: None,
60 pending: BTreeMap::new(),
61 }),
62 })
63 }
64
65 pub fn submit(
71 &self,
72 commit_id: u64,
73 mut snapshot: ShardSnapshotMetadata,
74 ) -> Result<Option<GlobalSnapshotManifest>> {
75 let mut state = self.lock_state()?;
76 if is_completed_or_superseded(&state, commit_id) {
77 return Ok(None);
78 }
79
80 normalize_ranges(self.total_buckets, &mut snapshot)?;
81 let table_metadata = crate::table::table_metadata_from_shard_snapshot(&snapshot)?;
82 if let std::collections::btree_map::Entry::Vacant(entry) = state.pending.entry(commit_id) {
83 entry.insert(PendingCommit::default());
84 if !retain_new_pending_commit(&mut state, commit_id, self.max_pending_commits) {
85 return Ok(None);
86 }
87 }
88 {
89 let pending = state
90 .pending
91 .get_mut(&commit_id)
92 .expect("commit remains in the pending window");
93 validate_table_metadata(pending, table_metadata)?;
94 insert_shard(pending, snapshot, commit_id)?;
95 if !has_exact_bucket_coverage(self.total_buckets, pending) {
96 return Ok(None);
97 }
98 }
99 self.commit_complete(&mut state, commit_id)
100 }
101
102 pub fn commit_batch(
109 &self,
110 commit_id: u64,
111 shard_snapshots: Vec<ShardSnapshotMetadata>,
112 ) -> Result<Option<GlobalSnapshotManifest>> {
113 let mut state = self.lock_state()?;
114 if is_completed_or_superseded(&state, commit_id) {
115 return Ok(None);
116 }
117
118 let mut batch = PendingCommit::default();
119 for mut snapshot in shard_snapshots {
120 normalize_ranges(self.total_buckets, &mut snapshot)?;
121 let table_metadata = crate::table::table_metadata_from_shard_snapshot(&snapshot)?;
122 validate_table_metadata(&mut batch, table_metadata)?;
123 insert_shard(&mut batch, snapshot, commit_id)?;
124 }
125 if !has_exact_bucket_coverage(self.total_buckets, &batch) {
126 return Err(coordination_error(format!(
127 "table snapshot commit {commit_id} does not cover all buckets"
128 )));
129 }
130 match state.pending.entry(commit_id) {
131 std::collections::btree_map::Entry::Occupied(mut entry) => {
132 let pending = entry.get();
133 if pending.table_metadata.as_ref() != batch.table_metadata.as_ref() {
134 return Err(coordination_error(
135 "table shard snapshots have incompatible captured table metadata",
136 ));
137 }
138 if pending.prepared_snapshot.is_some() && pending.shards != batch.shards {
139 return Err(coordination_error(format!(
140 "commit {commit_id} conflicts with its prepared snapshot"
141 )));
142 }
143 if pending.prepared_snapshot.is_none() {
144 entry.insert(batch);
145 }
146 }
147 std::collections::btree_map::Entry::Vacant(entry) => {
148 entry.insert(batch);
149 }
150 }
151 self.commit_complete(&mut state, commit_id)
152 }
153
154 fn lock_state(&self) -> Result<std::sync::MutexGuard<'_, CommitterState>> {
155 self.state
156 .lock()
157 .map_err(|_| coordination_error("table snapshot commit lock poisoned"))
158 }
159
160 fn commit_complete(
161 &self,
162 state: &mut CommitterState,
163 commit_id: u64,
164 ) -> Result<Option<GlobalSnapshotManifest>> {
165 let snapshot = {
166 let pending = state
167 .pending
168 .get_mut(&commit_id)
169 .expect("complete commit remains pending until publication succeeds");
170 if pending.prepared_snapshot.is_none() {
171 let inputs = canonical_inputs(pending);
172 pending.prepared_snapshot = Some(
173 self.coordinator
174 .take_global_snapshot(self.total_buckets, inputs)?,
175 );
176 }
177 self.materialize_prepared_snapshot(pending)?;
178 pending
179 .prepared_snapshot
180 .as_ref()
181 .expect("published commit retains its prepared snapshot")
182 .clone()
183 };
184 state.latest_completed_commit_id = Some(commit_id);
185 state.pending.retain(|id, _| *id > commit_id);
186 Ok(Some(snapshot))
187 }
188
189 fn materialize_prepared_snapshot(&self, pending: &mut PendingCommit) -> Result<()> {
190 let current = self.coordinator.load_current_global_snapshot()?;
191 if let Some(current) = ¤t {
192 let snapshot = pending
193 .prepared_snapshot
194 .as_ref()
195 .expect("complete commit has a prepared snapshot");
196 if current.id > snapshot.id {
197 pending.prepared_snapshot = Some(
200 self.coordinator
201 .take_global_snapshot(self.total_buckets, canonical_inputs(pending))?,
202 );
203 }
204 }
205 let snapshot = pending
206 .prepared_snapshot
207 .as_ref()
208 .expect("complete commit has a prepared snapshot");
209 if let Some(current) = current
210 && current.id == snapshot.id
211 {
212 if current != *snapshot {
213 return Err(coordination_error(format!(
214 "global snapshot {} conflicts with its prepared manifest",
215 snapshot.id
216 )));
217 }
218 return Ok(());
219 }
220 self.coordinator.materialize_global_snapshot(snapshot)?;
221 Ok(())
222 }
223}
224
225fn validate_table_metadata(
226 pending: &mut PendingCommit,
227 metadata: BTreeMap<String, TableMetadata>,
228) -> Result<()> {
229 if let Some(expected) = &pending.table_metadata {
230 if *expected != metadata {
233 return Err(coordination_error(
234 "table shard snapshots have incompatible captured table metadata",
235 ));
236 }
237 } else {
238 pending.table_metadata = Some(metadata);
239 }
240 Ok(())
241}
242
243fn is_completed_or_superseded(state: &CommitterState, commit_id: u64) -> bool {
244 state
245 .latest_completed_commit_id
246 .is_some_and(|latest| commit_id <= latest)
247}
248
249fn retain_new_pending_commit(
250 state: &mut CommitterState,
251 commit_id: u64,
252 max_pending_commits: usize,
253) -> bool {
254 if state.pending.len() <= max_pending_commits {
255 return true;
256 }
257 let oldest = *state
258 .pending
259 .first_key_value()
260 .expect("pending contains the inserted commit")
261 .0;
262 state.pending.remove(&oldest);
263 oldest != commit_id
264}
265
266fn insert_shard(
267 pending: &mut PendingCommit,
268 snapshot: ShardSnapshotMetadata,
269 commit_id: u64,
270) -> Result<()> {
271 match pending.shards.get(&snapshot.db_id) {
272 Some(existing) if existing != &snapshot => Err(coordination_error(format!(
273 "shard {} submitted conflicting snapshots for commit {commit_id}",
274 snapshot.db_id
275 ))),
276 Some(_) => Ok(()),
277 None => {
278 reject_overlapping_ranges(pending, &snapshot)?;
279 pending.shards.insert(snapshot.db_id.clone(), snapshot);
280 Ok(())
281 }
282 }
283}
284
285fn normalize_ranges(total_buckets: u32, input: &mut ShardSnapshotMetadata) -> Result<()> {
286 if input.ranges.is_empty() {
287 return Err(coordination_error(format!(
288 "shard snapshot ranges must not be empty for {}",
289 input.db_id
290 )));
291 }
292 input
293 .ranges
294 .sort_by_key(|range| (*range.start(), *range.end()));
295 let mut previous_end = None;
296 for range in &input.ranges {
297 let start = u32::from(*range.start());
298 let end = u32::from(*range.end());
299 if start > end || end >= total_buckets {
300 return Err(coordination_error(format!(
301 "invalid shard snapshot range {start}..={end} for {total_buckets} buckets"
302 )));
303 }
304 if previous_end.is_some_and(|previous| start <= previous) {
305 return Err(coordination_error(format!(
306 "shard snapshot ranges overlap at bucket {start}"
307 )));
308 }
309 previous_end = Some(end);
310 }
311 Ok(())
312}
313
314fn reject_overlapping_ranges(pending: &PendingCommit, input: &ShardSnapshotMetadata) -> Result<()> {
315 for (existing_db_id, existing) in &pending.shards {
316 let mut left = existing.ranges.iter().peekable();
317 let mut right = input.ranges.iter().peekable();
318 while let (Some(existing), Some(candidate)) = (left.peek(), right.peek()) {
319 if existing.end() < candidate.start() {
320 left.next();
321 } else if candidate.end() < existing.start() {
322 right.next();
323 } else {
324 return Err(coordination_error(format!(
325 "shard snapshot ranges overlap between {existing_db_id} and {}",
326 input.db_id
327 )));
328 }
329 }
330 }
331 Ok(())
332}
333
334fn has_exact_bucket_coverage(total_buckets: u32, pending: &PendingCommit) -> bool {
335 let mut ranges = pending
336 .shards
337 .values()
338 .flat_map(|input| input.ranges.iter())
339 .collect::<Vec<_>>();
340 ranges.sort_by_key(|range| (*range.start(), *range.end()));
341 let mut expected = 0;
342 for range in ranges {
343 if u32::from(*range.start()) != expected {
344 return false;
345 }
346 expected = u32::from(*range.end()) + 1;
347 }
348 expected == total_buckets
349}
350
351fn canonical_inputs(pending: &PendingCommit) -> Vec<ShardSnapshotMetadata> {
352 let mut inputs = pending.shards.values().cloned().collect::<Vec<_>>();
353 inputs.sort_by(|left, right| {
354 left.ranges
355 .iter()
356 .map(|range| (*range.start(), *range.end()))
357 .cmp(
358 right
359 .ranges
360 .iter()
361 .map(|range| (*range.start(), *range.end())),
362 )
363 .then_with(|| left.db_id.cmp(&right.db_id))
364 });
365 inputs
366}
367
368fn coordination_error(message: impl Into<String>) -> TableError {
369 cobble::Error::CoordinationError(message.into()).into()
370}