Skip to main content

holos_tda/distributed/
store.rs

1use std::fs;
2use std::path::{Path, PathBuf};
3
4use crate::{CertificateLimits, RelativeInterfaceCertificate};
5
6use super::compose::{
7    canonical_vertices, check_progress_identity, check_progress_prefix, check_progress_shape,
8    combined_vertices, compose_accumulator, compose_certificates, decode_certificate,
9    encode_certificate, job_id, require_first_shard, require_separator,
10};
11use super::fs::{atomic_replace, atomic_write, read_bounded, sync_directory, sync_parent};
12use super::model::{
13    ArtifactId, CommitPlan, DistributedInterfaceCommit, DistributedInterfaceError,
14    DistributedInterfaceManifest, DistributedInterfaceWork, DurableInterfaceStore, FoldState,
15    PreparedCommit, Progress,
16};
17use super::wire::{PROGRESS_MAGIC, Reader, VERSION, decode_ids, encode_ids, put_usize};
18
19mod verification;
20
21impl DurableInterfaceStore {
22    /// Open or create a store rooted at `path`.
23    pub fn open(path: impl AsRef<Path>) -> Result<Self, DistributedInterfaceError> {
24        let root = path.as_ref().to_path_buf();
25        fs::create_dir_all(root.join("objects"))
26            .and_then(|_| fs::create_dir_all(root.join("jobs")))
27            .and_then(|_| fs::create_dir_all(root.join("manifests")))
28            .map_err(|error| DistributedInterfaceError::new(format!("create store: {error}")))?;
29        Ok(Self { root })
30    }
31
32    /// Store exact bytes and return their content identifier.
33    pub fn put(&self, bytes: &[u8]) -> Result<(ArtifactId, bool), DistributedInterfaceError> {
34        let id = ArtifactId::for_bytes(bytes);
35        let path = self.object_path(id);
36        if path.is_file() {
37            let existing = read_bounded(&path, bytes.len())?;
38            if existing != bytes {
39                return Err(DistributedInterfaceError::new(
40                    "stored object differs under the same content identifier",
41                ));
42            }
43            return Ok((id, false));
44        }
45        let parent = path.parent().expect("object path has a parent");
46        fs::create_dir_all(parent).map_err(|error| {
47            DistributedInterfaceError::new(format!("create object path: {error}"))
48        })?;
49        sync_parent(parent)?;
50        sync_directory(parent)?;
51        atomic_write(&path, bytes)?;
52        Ok((id, true))
53    }
54
55    /// Read an object and check its identifier.
56    pub fn get(
57        &self,
58        id: ArtifactId,
59        maximum_bytes: usize,
60    ) -> Result<Vec<u8>, DistributedInterfaceError> {
61        let bytes = read_bounded(&self.object_path(id), maximum_bytes)?;
62        if ArtifactId::for_bytes(&bytes) != id {
63            return Err(DistributedInterfaceError::new(
64                "stored object fails its content identifier",
65            ));
66        }
67        Ok(bytes)
68    }
69
70    /// Whether an object path exists. This does not read or verify the bytes.
71    pub fn contains(&self, id: ArtifactId) -> bool {
72        self.object_path(id).is_file()
73    }
74
75    /// Compose shard artifacts and publish one atomic manifest.
76    ///
77    /// Shards are ordered. Every intermediate fold protects the union of
78    /// `separator_vertices` and `output_protected_vertices`. A retry resumes
79    /// from the last durable fold for the same content-bound job.
80    pub fn commit(
81        &self,
82        shard_artifacts: &[Vec<u8>],
83        separator_vertices: &[usize],
84        output_protected_vertices: &[usize],
85        limits: CertificateLimits,
86    ) -> Result<DistributedInterfaceCommit, DistributedInterfaceError> {
87        let mut shard_ids = Vec::with_capacity(shard_artifacts.len());
88        let mut work = DistributedInterfaceWork {
89            shards: shard_artifacts.len(),
90            ..DistributedInterfaceWork::default()
91        };
92        for bytes in shard_artifacts {
93            let (id, written) = self.put(bytes)?;
94            shard_ids.push(id);
95            if written {
96                work.bytes_written += bytes.len();
97            }
98        }
99        self.commit_stored_inner(
100            &shard_ids,
101            separator_vertices,
102            output_protected_vertices,
103            limits,
104            work,
105        )
106    }
107
108    /// Compose shards already present in this store.
109    ///
110    /// The coordinator loads at most one shard and one accumulator artifact
111    /// for each fold. The identifiers are ordered and bind the job.
112    pub fn commit_stored(
113        &self,
114        shard_ids: &[ArtifactId],
115        separator_vertices: &[usize],
116        output_protected_vertices: &[usize],
117        limits: CertificateLimits,
118    ) -> Result<DistributedInterfaceCommit, DistributedInterfaceError> {
119        self.commit_stored_inner(
120            shard_ids,
121            separator_vertices,
122            output_protected_vertices,
123            limits,
124            DistributedInterfaceWork {
125                shards: shard_ids.len(),
126                ..DistributedInterfaceWork::default()
127            },
128        )
129    }
130
131    fn commit_stored_inner(
132        &self,
133        shard_ids: &[ArtifactId],
134        separator_vertices: &[usize],
135        output_protected_vertices: &[usize],
136        limits: CertificateLimits,
137        mut work: DistributedInterfaceWork,
138    ) -> Result<DistributedInterfaceCommit, DistributedInterfaceError> {
139        let prepared = self.prepare_commit(
140            shard_ids,
141            separator_vertices,
142            output_protected_vertices,
143            limits,
144            &mut work,
145        )?;
146        if let Some(manifest) = self.read_manifest(prepared.plan.job, limits.max_bytes)? {
147            return self.load_commit(manifest, work, limits);
148        }
149        let mut state =
150            self.resume_or_start(&prepared.plan, prepared.first_id, prepared.first, &mut work)?;
151        self.compute_folds(&prepared.plan, &mut state, &mut work)?;
152        self.finish_protection(&prepared.plan, &mut state, &mut work)?;
153        self.publish_commit(&prepared.plan, state, work)
154    }
155
156    fn prepare_commit<'a>(
157        &self,
158        shard_ids: &'a [ArtifactId],
159        separator_vertices: &[usize],
160        output_protected_vertices: &[usize],
161        limits: CertificateLimits,
162        work: &mut DistributedInterfaceWork,
163    ) -> Result<PreparedCommit<'a>, DistributedInterfaceError> {
164        let first_id = require_first_shard(shard_ids)?;
165        let separator_vertices = canonical_vertices(separator_vertices)?;
166        let output_protected_vertices = canonical_vertices(output_protected_vertices)?;
167        let first_bytes = self.get(first_id, limits.max_bytes)?;
168        work.bytes_read += first_bytes.len();
169        work.peak_artifact_bytes = first_bytes.len();
170        work.shards_loaded += 1;
171        let first = decode_certificate(&first_bytes, limits)?;
172        require_separator(&first, &separator_vertices)?;
173        let plan = CommitPlan {
174            job: job_id(
175                first.max_dim(),
176                first.modulus(),
177                &separator_vertices,
178                &output_protected_vertices,
179                shard_ids,
180            ),
181            shard_ids,
182            intermediate_protected: combined_vertices(
183                &separator_vertices,
184                &output_protected_vertices,
185            ),
186            separator_vertices,
187            output_protected_vertices,
188            limits,
189        };
190        Ok(PreparedCommit {
191            plan,
192            first_id,
193            first,
194        })
195    }
196
197    fn resume_or_start(
198        &self,
199        plan: &CommitPlan<'_>,
200        first_id: ArtifactId,
201        first: RelativeInterfaceCertificate,
202        work: &mut DistributedInterfaceWork,
203    ) -> Result<FoldState, DistributedInterfaceError> {
204        let Some(progress) = self.read_progress(plan.job, plan.limits.max_bytes)? else {
205            return self.start_folds(plan, first_id, first);
206        };
207        check_progress_prefix(progress.prefix, plan.shard_ids.len())?;
208        let bytes = self.get(progress.accumulator, plan.limits.max_bytes)?;
209        work.bytes_read += bytes.len();
210        work.folds_reused = progress.prefix;
211        Ok(FoldState {
212            accumulator: decode_certificate(&bytes, plan.limits)?,
213            folds: progress.folds,
214            next: progress.prefix,
215            accumulator_bytes: bytes,
216        })
217    }
218
219    fn start_folds(
220        &self,
221        plan: &CommitPlan<'_>,
222        first_id: ArtifactId,
223        first: RelativeInterfaceCertificate,
224    ) -> Result<FoldState, DistributedInterfaceError> {
225        let folds = vec![first_id];
226        self.write_progress(
227            plan.job,
228            &Progress {
229                prefix: 1,
230                accumulator: first_id,
231                folds: folds.clone(),
232            },
233        )?;
234        Ok(FoldState {
235            accumulator_bytes: encode_certificate(&first, plan.limits)?,
236            accumulator: first,
237            folds,
238            next: 1,
239        })
240    }
241
242    fn compute_folds(
243        &self,
244        plan: &CommitPlan<'_>,
245        state: &mut FoldState,
246        work: &mut DistributedInterfaceWork,
247    ) -> Result<(), DistributedInterfaceError> {
248        for (position, shard_id) in plan.shard_ids.iter().enumerate().skip(state.next) {
249            self.compute_one_fold(plan, state, work, position, *shard_id)?;
250        }
251        Ok(())
252    }
253
254    fn compute_one_fold(
255        &self,
256        plan: &CommitPlan<'_>,
257        state: &mut FoldState,
258        work: &mut DistributedInterfaceWork,
259        position: usize,
260        shard_id: ArtifactId,
261    ) -> Result<(), DistributedInterfaceError> {
262        let (child, child_bytes) = self.load_child(plan, shard_id, work)?;
263        compose_accumulator(plan, state, child, child_bytes, work)?;
264        let fold = self.store_fold(&state.accumulator_bytes, work)?;
265        state.folds.push(fold);
266        state.next = position + 1;
267        work.folds_computed += 1;
268        self.write_progress(
269            plan.job,
270            &Progress {
271                prefix: state.next,
272                accumulator: fold,
273                folds: state.folds.clone(),
274            },
275        )
276    }
277
278    fn load_child(
279        &self,
280        plan: &CommitPlan<'_>,
281        shard_id: ArtifactId,
282        work: &mut DistributedInterfaceWork,
283    ) -> Result<(RelativeInterfaceCertificate, usize), DistributedInterfaceError> {
284        let bytes = self.get(shard_id, plan.limits.max_bytes)?;
285        work.bytes_read += bytes.len();
286        work.shards_loaded += 1;
287        let child = decode_certificate(&bytes, plan.limits)?;
288        require_separator(&child, &plan.separator_vertices)?;
289        Ok((child, bytes.len()))
290    }
291
292    fn store_fold(
293        &self,
294        bytes: &[u8],
295        work: &mut DistributedInterfaceWork,
296    ) -> Result<ArtifactId, DistributedInterfaceError> {
297        let (fold, written) = self.put(bytes)?;
298        if written {
299            work.bytes_written += bytes.len();
300        }
301        Ok(fold)
302    }
303
304    fn finish_protection(
305        &self,
306        plan: &CommitPlan<'_>,
307        state: &mut FoldState,
308        work: &mut DistributedInterfaceWork,
309    ) -> Result<(), DistributedInterfaceError> {
310        if state.accumulator.protected_vertices() == plan.output_protected_vertices {
311            return Ok(());
312        }
313        state.accumulator = compose_certificates(
314            &[&state.accumulator],
315            &plan.output_protected_vertices,
316            plan.limits,
317        )?;
318        state.accumulator_bytes = encode_certificate(&state.accumulator, plan.limits)?;
319        self.store_fold(&state.accumulator_bytes, work)?;
320        Ok(())
321    }
322
323    fn publish_commit(
324        &self,
325        plan: &CommitPlan<'_>,
326        state: FoldState,
327        work: DistributedInterfaceWork,
328    ) -> Result<DistributedInterfaceCommit, DistributedInterfaceError> {
329        let manifest = DistributedInterfaceManifest {
330            job: plan.job,
331            max_dim: state.accumulator.max_dim(),
332            modulus: state.accumulator.modulus(),
333            separator_vertices: plan.separator_vertices.clone(),
334            output_protected_vertices: plan.output_protected_vertices.clone(),
335            shards: plan.shard_ids.to_vec(),
336            folds: state.folds,
337            result: ArtifactId::for_bytes(&state.accumulator_bytes),
338        };
339        atomic_write(&self.manifest_path(plan.job), &manifest.encode()?)?;
340        Ok(DistributedInterfaceCommit {
341            manifest,
342            certificate: state.accumulator,
343            work,
344        })
345    }
346
347    /// Read a committed manifest by job identifier.
348    pub fn manifest(
349        &self,
350        job: ArtifactId,
351        maximum_bytes: usize,
352    ) -> Result<DistributedInterfaceManifest, DistributedInterfaceError> {
353        self.read_manifest(job, maximum_bytes)?
354            .ok_or_else(|| DistributedInterfaceError::new("distributed manifest is absent"))
355    }
356
357    fn load_commit(
358        &self,
359        manifest: DistributedInterfaceManifest,
360        mut work: DistributedInterfaceWork,
361        limits: CertificateLimits,
362    ) -> Result<DistributedInterfaceCommit, DistributedInterfaceError> {
363        let bytes = self.get(manifest.result, limits.max_bytes)?;
364        work.bytes_read += bytes.len();
365        work.folds_reused = manifest.folds.len();
366        let certificate = RelativeInterfaceCertificate::decode(&bytes, limits)
367            .map_err(|error| DistributedInterfaceError::new(error.to_string()))?;
368        Ok(DistributedInterfaceCommit {
369            manifest,
370            certificate,
371            work,
372        })
373    }
374
375    pub(super) fn object_path(&self, id: ArtifactId) -> PathBuf {
376        let hex = id.to_hex();
377        self.root.join("objects").join(&hex[..2]).join(&hex[2..])
378    }
379
380    pub(super) fn manifest_path(&self, id: ArtifactId) -> PathBuf {
381        self.root.join("manifests").join(format!("{id}.hdm"))
382    }
383
384    fn progress_path(&self, id: ArtifactId) -> PathBuf {
385        self.root.join("jobs").join(format!("{id}.work"))
386    }
387
388    fn read_manifest(
389        &self,
390        id: ArtifactId,
391        maximum_bytes: usize,
392    ) -> Result<Option<DistributedInterfaceManifest>, DistributedInterfaceError> {
393        let path = self.manifest_path(id);
394        if !path.is_file() {
395            return Ok(None);
396        }
397        let bytes = read_bounded(&path, maximum_bytes)?;
398        DistributedInterfaceManifest::decode(&bytes, maximum_bytes).map(Some)
399    }
400
401    pub(super) fn write_progress(
402        &self,
403        job: ArtifactId,
404        progress: &Progress,
405    ) -> Result<(), DistributedInterfaceError> {
406        let mut output = Vec::new();
407        output.extend_from_slice(PROGRESS_MAGIC);
408        output.extend_from_slice(&VERSION.to_be_bytes());
409        output.extend_from_slice(job.as_bytes());
410        put_usize(&mut output, progress.prefix)?;
411        output.extend_from_slice(progress.accumulator.as_bytes());
412        encode_ids(&mut output, &progress.folds)?;
413        atomic_replace(&self.progress_path(job), &output)
414    }
415
416    fn read_progress(
417        &self,
418        job: ArtifactId,
419        maximum_bytes: usize,
420    ) -> Result<Option<Progress>, DistributedInterfaceError> {
421        let path = self.progress_path(job);
422        if !path.is_file() {
423            return Ok(None);
424        }
425        let bytes = read_bounded(&path, maximum_bytes)?;
426        let mut reader = Reader::new(&bytes);
427        check_progress_identity(&mut reader, job)?;
428        let prefix = reader.usize()?;
429        let accumulator = ArtifactId(reader.array32()?);
430        let folds = decode_ids(&mut reader)?;
431        check_progress_shape(reader.remaining(), prefix, accumulator, &folds)?;
432        Ok(Some(Progress {
433            prefix,
434            accumulator,
435            folds,
436        }))
437    }
438}