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 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 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 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 pub fn contains(&self, id: ArtifactId) -> bool {
72 self.object_path(id).is_file()
73 }
74
75 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 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 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}