1use std::collections::{BTreeMap, BTreeSet};
8use std::fmt;
9use std::path::{Path, PathBuf};
10
11use henad_core::explore::plan::Shard;
12
13use crate::output::manifest::{
14 BuildRole, Manifest, ManifestError, ManifestMode, ManifestStatus, ManifestTimestamps, RecordedBuild, ResultCounts,
15 now_unix_ms, rfc3339,
16};
17use crate::output::read::{ReadError, RunRecord, RunsCsv, SeriesScan, SeriesSegment, merge_series};
18use crate::output::runs_csv::header_line;
19use crate::output::{MANIFEST_FILE, OutputDir, OutputError, RUNS_FILE, SERIES_FILE, table_paths};
20use crate::progress::{Progress, ProgressEvent};
21use crate::sweep::SweepWarning;
22
23pub const MAX_LISTED_RUNS: usize = 5;
25
26#[derive(Debug, Clone, PartialEq, Eq)]
28pub struct MergeReport {
29 pub counts: ResultCounts,
31 pub missing: u64,
33 pub output_dir: PathBuf,
35}
36
37struct ShardInput {
39 dir: PathBuf,
40 manifest: Manifest,
41 shard: Shard,
42 runs: RunsCsv,
43 series: SeriesScan,
44 run_ids: BTreeSet<u64>,
46}
47
48pub fn merge(
62 shard_dirs: &[PathBuf],
63 output_dir: &Path,
64 progress: &mut dyn Progress,
65) -> Result<MergeReport, MergeError> {
66 let mut shards = shard_dirs
67 .iter()
68 .map(|dir| read_input(dir))
69 .collect::<Result<Vec<_>, _>>()?;
70 shards.sort_by_key(|input| input.shard.index());
71 let Some(first) = shards.first() else {
72 return Err(MergeError::NoInputs);
73 };
74 check_shards(&shards)?;
75 for warning in build_warnings(&shards) {
76 progress.report(&ProgressEvent::Warned(&warning));
77 }
78 let (runs_header, series_header) = common_headers(&shards)?;
79 let plan_runs = first.manifest.plan.runs;
80 let (records, counts) = merged_records(&shards, plan_runs, first.manifest.plan.replicates)?;
81
82 let dir = OutputDir::create(output_dir).map_err(MergeError::Output)?;
83 let mut manifest = merged_manifest(&shards, shard_dirs);
84 dir.write_manifest(&manifest).map_err(MergeError::Output)?;
85 if let Err(error) = write_tables(&dir, &shards, &records, &runs_header, &series_header) {
86 manifest.fail(now_unix_ms());
87 drop(dir.write_manifest(&manifest));
89 return Err(MergeError::Output(error));
90 }
91
92 let missing = plan_runs - records.len() as u64;
93 if missing > 0 {
94 let first_missing = (0..plan_runs)
95 .filter(|run_id| !records.contains_key(run_id))
96 .take(MAX_LISTED_RUNS)
97 .collect();
98 progress.report(&ProgressEvent::Warned(&SweepWarning::MissingRuns {
99 count: missing,
100 first: first_missing,
101 }));
102 }
103 let status = if missing == 0 {
104 ManifestStatus::Complete
105 } else {
106 ManifestStatus::Incomplete
107 };
108 let finished = now_unix_ms();
110 manifest.status = status;
111 manifest.results = Some(counts);
112 manifest.timestamps.finished_unix_ms = Some(finished);
113 manifest.timestamps.finished = Some(rfc3339(finished));
114 dir.write_manifest(&manifest).map_err(MergeError::Output)?;
115 Ok(MergeReport {
116 counts,
117 missing,
118 output_dir: output_dir.to_owned(),
119 })
120}
121
122fn write_tables(
125 dir: &OutputDir,
126 shards: &[ShardInput],
127 records: &BTreeMap<u64, &RunRecord>,
128 runs_header: &str,
129 series_header: &str,
130) -> Result<(), OutputError> {
131 let segments: Vec<SeriesSegment> = shards
132 .iter()
133 .enumerate()
134 .flat_map(|(input_index, input)| {
135 input.series.segments.iter().map(move |range| SeriesSegment {
136 path: input.series.path.clone(),
137 range: range.clone(),
138 input_index,
139 })
140 })
141 .collect();
142 dir.replace_tables(
143 |dest| {
144 dest.write_all(runs_header.as_bytes())?;
145 records
146 .values()
147 .try_for_each(|record| dest.write_all(record.text.as_bytes()))
148 },
149 |dest| {
150 merge_series(dest, series_header, &segments, |input_index, run_id| {
151 shards[input_index].run_ids.contains(&run_id).then_some(run_id)
152 })
153 },
154 )?;
155 dir.write_summary()
156}
157
158fn common_headers(shards: &[ShardInput]) -> Result<(String, String), MergeError> {
160 let runs_header = shards
161 .iter()
162 .find_map(|input| input.runs.header.as_ref())
163 .map(|header| header_line(header));
164 let series_header = shards.iter().find_map(|input| input.series.header.clone());
165 let (Some(runs_header), Some(series_header)) = (runs_header, series_header) else {
166 return Err(MergeError::NoTables);
167 };
168 for input in shards {
169 let differs = |file| MergeError::ColumnsDiffer {
170 dir: input.dir.clone(),
171 file,
172 };
173 if input
174 .runs
175 .header
176 .as_ref()
177 .is_some_and(|header| header_line(header) != runs_header)
178 {
179 return Err(differs(RUNS_FILE));
180 }
181 if input
182 .series
183 .header
184 .as_ref()
185 .is_some_and(|header| *header != series_header)
186 {
187 return Err(differs(SERIES_FILE));
188 }
189 }
190 Ok((runs_header, series_header))
191}
192
193fn merged_records(
197 shards: &[ShardInput],
198 plan_runs: u64,
199 replicates: u64,
200) -> Result<(BTreeMap<u64, &RunRecord>, ResultCounts), MergeError> {
201 let mut records = BTreeMap::new();
202 let mut counts = ResultCounts::default();
203 for input in shards {
204 for record in &input.runs.records {
205 let numbered = record
206 .config_id
207 .checked_mul(replicates)
208 .and_then(|first| first.checked_add(record.rep));
209 if record.rep >= replicates || numbered != Some(record.run_id) {
210 return Err(MergeError::RunIdDiffers {
211 dir: input.dir.clone(),
212 run_id: record.run_id,
213 config_id: record.config_id,
214 rep: record.rep,
215 });
216 }
217 if !input.shard.contains(record.run_id) || record.run_id >= plan_runs {
218 return Err(MergeError::OutsideShard {
219 dir: input.dir.clone(),
220 run_id: record.run_id,
221 });
222 }
223 if records.insert(record.run_id, record).is_some() {
224 return Err(MergeError::DuplicateRun {
225 dir: input.dir.clone(),
226 run_id: record.run_id,
227 });
228 }
229 counts.count(record.status);
230 }
231 }
232 Ok((records, counts))
233}
234
235fn read_input(dir: &Path) -> Result<ShardInput, MergeError> {
237 let manifest = Manifest::read(&dir.join(MANIFEST_FILE)).map_err(MergeError::Manifest)?;
238 if manifest.mode != ManifestMode::Sweep {
239 return Err(MergeError::NotASweep { dir: dir.to_owned() });
240 }
241 let shard = manifest
242 .shard
243 .to_shard()
244 .ok_or_else(|| MergeError::BadShard { dir: dir.to_owned() })?;
245 let (runs_path, series_path) = table_paths(dir);
246 let runs = RunsCsv::read(&runs_path).map_err(MergeError::Table)?;
247 let run_ids: BTreeSet<u64> = runs.records.iter().map(|record| record.run_id).collect();
248 let series = SeriesScan::read(&series_path, |run_id| run_ids.contains(&run_id)).map_err(MergeError::Table)?;
249 Ok(ShardInput {
250 dir: dir.to_owned(),
251 manifest,
252 shard,
253 runs,
254 series,
255 run_ids,
256 })
257}
258
259fn check_shards(shards: &[ShardInput]) -> Result<(), MergeError> {
261 let Some(first) = shards.first() else {
262 return Err(MergeError::NoInputs);
263 };
264 for pair in shards.windows(2) {
265 if pair[0].shard == pair[1].shard {
266 return Err(MergeError::SameShard {
267 shard: pair[0].shard,
268 first: pair[0].dir.clone(),
269 second: pair[1].dir.clone(),
270 });
271 }
272 }
273 for input in shards {
274 let (plan, expected) = (&input.manifest.plan, &first.manifest.plan);
275 if plan.plan_hash != expected.plan_hash
276 || input.manifest.model.schema_hash != first.manifest.model.schema_hash
277 || plan.replicates != expected.replicates
278 || plan.runs != expected.runs
279 {
280 return Err(MergeError::PlanDiffers {
281 dir: input.dir.clone(),
282 first: first.dir.clone(),
283 });
284 }
285 if input.shard.count() != first.shard.count() {
286 return Err(MergeError::ShardCountDiffers {
287 dir: input.dir.clone(),
288 first: first.dir.clone(),
289 });
290 }
291 }
292 Ok(())
293}
294
295fn build_warnings(shards: &[ShardInput]) -> Vec<SweepWarning> {
302 let mut warnings = Vec::new();
303 for role in [BuildRole::Engine, BuildRole::Model] {
304 let mut listed = shards.iter().map(|input| input.manifest.recorded_builds(role));
305 let Some(reference) = listed.by_ref().find(|builds| !builds.is_empty()) else {
306 continue;
307 };
308 let mut others: Vec<RecordedBuild> = Vec::new();
309 for build in listed.flatten() {
310 if !others.iter().any(|known| known.reads_as(&build)) {
311 others.push(build);
312 }
313 }
314 for other in others {
315 if reference.iter().any(|build| build.same_build(&other)) {
316 continue;
317 }
318 let recorded = reference
319 .iter()
320 .find(|build| build.reads_as(&other))
321 .unwrap_or(&reference[0]);
322 warnings.push(SweepWarning::BuildChanged {
323 role,
324 recorded: Box::new(recorded.clone()),
325 current: Box::new(other),
326 between_shards: true,
327 });
328 }
329 }
330 warnings
331}
332
333fn merged_manifest(shards: &[ShardInput], shard_dirs: &[PathBuf]) -> Manifest {
341 let mut manifest = shards[0].manifest.clone();
342 manifest.status = ManifestStatus::Running;
343 manifest.shard = Shard::WHOLE.into();
344 manifest.engine = RecordedBuild::engine();
345 manifest.model.replays_exactly = shards.iter().all(|input| input.manifest.model.replays_exactly);
346 manifest.sessions = shards
347 .iter()
348 .flat_map(|input| {
349 let mut recorded = input.manifest.clone();
350 recorded.record_session_engines();
351 if matches!(recorded.status, ManifestStatus::Running | ManifestStatus::Failed)
353 && let Some(last) = recorded.sessions.last_mut()
354 {
355 last.ran = (input.runs.records.len() as u64).saturating_sub(last.skipped);
356 }
357 recorded.sessions
358 })
359 .collect();
360 let started = shards
361 .iter()
362 .map(|input| input.manifest.timestamps.started_unix_ms)
363 .min()
364 .unwrap_or_default();
365 manifest.timestamps = ManifestTimestamps::started_at(started);
366 manifest.results = None;
367 manifest.merged_shards = Some(shard_dirs.iter().map(|dir| dir.display().to_string()).collect());
368 manifest
369}
370
371#[derive(Debug)]
373pub enum MergeError {
374 NoInputs,
376 Manifest(ManifestError),
378 Table(ReadError),
380 BadShard {
382 dir: PathBuf,
384 },
385 NoTables,
387 PlanDiffers {
389 dir: PathBuf,
391 first: PathBuf,
393 },
394 ShardCountDiffers {
396 dir: PathBuf,
398 first: PathBuf,
400 },
401 SameShard {
403 shard: Shard,
405 first: PathBuf,
407 second: PathBuf,
409 },
410 ColumnsDiffer {
412 dir: PathBuf,
414 file: &'static str,
416 },
417 RunIdDiffers {
420 dir: PathBuf,
422 run_id: u64,
424 config_id: u64,
426 rep: u64,
428 },
429 OutsideShard {
431 dir: PathBuf,
433 run_id: u64,
435 },
436 DuplicateRun {
438 dir: PathBuf,
440 run_id: u64,
442 },
443 Output(OutputError),
445 NotASweep {
447 dir: PathBuf,
449 },
450}
451
452impl fmt::Display for MergeError {
453 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
454 match self {
455 Self::NoInputs => f.write_str("a merge needs at least one shard directory"),
456 Self::Manifest(_) => f.write_str("cannot read the manifest of a shard"),
457 Self::Table(_) => f.write_str("cannot read the tables of a shard"),
458 Self::BadShard { dir } => write!(f, "'{}' records an invalid shard", dir.display()),
459 Self::NoTables => f.write_str("no shard holds a runs.csv and a series.csv with a header"),
460 Self::PlanDiffers { dir, first } => write!(
461 f,
462 "'{}' has a different plan, model schema or replicate count from '{}'",
463 dir.display(),
464 first.display()
465 ),
466 Self::ShardCountDiffers { dir, first } => write!(
467 f,
468 "'{}' has a different shard count from '{}'",
469 dir.display(),
470 first.display()
471 ),
472 Self::SameShard { shard, first, second } => write!(
473 f,
474 "'{}' and '{}' both hold shard {shard}",
475 first.display(),
476 second.display()
477 ),
478 Self::ColumnsDiffer { dir, file } => {
479 write!(
480 f,
481 "{file} of '{}' has different columns from the other shards",
482 dir.display()
483 )
484 }
485 Self::RunIdDiffers {
486 dir,
487 run_id,
488 config_id,
489 rep,
490 } => write!(
491 f,
492 "'{}' holds replicate {rep} of config {config_id} as run {run_id}, and its manifest gives a \
493 different run id. Resume that shard before merging it",
494 dir.display()
495 ),
496 Self::OutsideShard { dir, run_id } => {
497 write!(f, "'{}' holds run {run_id}, outside its shard", dir.display())
498 }
499 Self::DuplicateRun { dir, run_id } => {
500 write!(f, "'{}' lists run {run_id} twice", dir.display())
501 }
502 Self::Output(_) => f.write_str("cannot write the merged results"),
503 Self::NotASweep { dir } => write!(f, "'{}' holds a search. A search has no shards to merge", dir.display()),
504 }
505 }
506}
507
508impl std::error::Error for MergeError {
509 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
510 match self {
511 Self::Manifest(error) => Some(error),
512 Self::Table(error) => Some(error),
513 Self::Output(error) => Some(error),
514 Self::NoInputs
515 | Self::BadShard { .. }
516 | Self::NoTables
517 | Self::PlanDiffers { .. }
518 | Self::ShardCountDiffers { .. }
519 | Self::SameShard { .. }
520 | Self::ColumnsDiffer { .. }
521 | Self::RunIdDiffers { .. }
522 | Self::OutsideShard { .. }
523 | Self::DuplicateRun { .. }
524 | Self::NotASweep { .. } => None,
525 }
526 }
527}
528
529#[cfg(test)]
530mod tests {
531 use std::fs;
532 use std::path::{Path, PathBuf};
533
534 use henad_core::explore::plan::Shard;
535 use henad_core::explore::spec::SweepSpec;
536
537 use super::{MergeError, merge};
538 use crate::exec::Concurrency;
539 use crate::output::RUNS_FILE;
540 use crate::output::manifest::{BuildRole, RecordedBuild};
541 use crate::progress::NoProgress;
542 use crate::sweep::SweepOptions;
543 use crate::tests::support::{
544 Recorder, ScratchDir, entry, manifest, other_engine, provenance, rewrite_manifest, sweep, sweep_options,
545 sweep_with,
546 };
547
548 fn life_spec() -> SweepSpec {
550 let mut spec = SweepSpec::new("game_of_life");
551 spec.fixed = vec![
552 ("grid_width".to_owned(), "8".to_owned()),
553 ("grid_height".to_owned(), "8".to_owned()),
554 ];
555 spec.run.steps = 2;
556 spec.run.replicates = 4;
557 spec
558 }
559
560 fn shard_under(dir: &Path, index: u64, engine: RecordedBuild) {
562 let options = SweepOptions {
563 provenance: provenance().with_engine(engine),
564 shard: Shard::new(index, 2).expect("a valid shard"),
565 ..sweep_options(false)
566 };
567 let life = entry("game_of_life", None);
568 sweep_with(&life, None, &life_spec(), dir, &options, &mut Recorder::default()).expect("the shard runs");
569 }
570
571 fn shard_dirs(scratch: &ScratchDir) -> Vec<PathBuf> {
573 (0..2)
574 .map(|index| scratch.path().join(format!("shard-{index}")))
575 .collect()
576 }
577
578 #[test]
579 fn a_build_the_lowest_shard_ran_is_never_warned_about() {
580 let scratch = ScratchDir::new("merge-reference-builds");
581 let dirs = shard_dirs(&scratch);
582 shard_under(&dirs[0], 0, RecordedBuild::engine());
583 shard_under(&dirs[1], 1, RecordedBuild::engine());
584 rewrite_manifest(&dirs[0], |recorded| {
586 let mut resumed = recorded.sessions[0].clone();
587 (resumed.skipped, resumed.ran, resumed.engine) = (1, 1, Some(other_engine()));
588 recorded.sessions[0].ran = 1;
589 recorded.sessions.push(resumed);
590 });
591 let mut progress = Recorder::default();
592 merge(&dirs, &scratch.path().join("merged"), &mut progress).expect("the shards merge");
593 assert_eq!(progress.warnings, [], "shard 1 ran a build shard 0 ran");
594
595 let reversed = ScratchDir::new("merge-reference-other");
596 let dirs = shard_dirs(&reversed);
597 shard_under(&dirs[0], 0, RecordedBuild::engine());
598 shard_under(&dirs[1], 1, other_engine());
599 let mut progress = Recorder::default();
600 merge(&dirs, &reversed.path().join("merged"), &mut progress).expect("the shards merge");
601 assert_eq!(progress.warnings.len(), 1, "{:?}", progress.warnings);
602 }
603
604 #[test]
605 fn a_merge_credits_a_shard_session_that_never_ended() {
606 let scratch = ScratchDir::new("merge-credits");
607 let dirs = shard_dirs(&scratch);
608 shard_under(&dirs[0], 0, RecordedBuild::engine());
609 shard_under(&dirs[1], 1, RecordedBuild::engine());
610 rewrite_manifest(&dirs[1], |recorded| {
611 recorded.status = crate::output::manifest::ManifestStatus::Running;
612 recorded.results = None;
613 recorded.sessions[0].ran = 0;
614 recorded.model.replays_exactly = false;
615 });
616 let merged = scratch.path().join("merged");
617 merge(&dirs, &merged, &mut NoProgress).expect("the shards merge");
618 let recorded = manifest(&merged);
619 let ran: Vec<u64> = recorded.sessions.iter().map(|session| session.ran).collect();
620 assert_eq!(ran, [2, 2], "the killed shard's session wrote both of its runs");
621 assert_eq!(recorded.recorded_builds(BuildRole::Engine), [RecordedBuild::engine()]);
622 assert!(
623 !recorded.model.replays_exactly,
624 "a model replays exactly only when every shard's does"
625 );
626 }
627
628 #[test]
629 fn a_shard_listing_a_run_twice_is_refused() {
630 let scratch = ScratchDir::new("merge-duplicate");
631 let shard_dir = scratch.path().join("shard");
632 let mut spec = SweepSpec::new("game_of_life");
633 spec.fixed = vec![
634 ("grid_width".to_owned(), "8".to_owned()),
635 ("grid_height".to_owned(), "8".to_owned()),
636 ];
637 spec.run.steps = 2;
638 spec.run.replicates = 2;
639 sweep(&entry("game_of_life", None), None, &spec, &shard_dir, Concurrency::Auto);
640 let runs_path = shard_dir.join(RUNS_FILE);
641 let runs = fs::read_to_string(&runs_path).expect("runs.csv is written");
642 let last_row = runs.lines().last().expect("runs.csv holds a run");
643 fs::write(&runs_path, format!("{runs}{last_row}\n")).expect("runs.csv is written again");
644
645 let error = merge(
646 std::slice::from_ref(&shard_dir),
647 &scratch.path().join("merged"),
648 &mut NoProgress,
649 )
650 .expect_err("a run listed twice");
651 assert!(matches!(error, MergeError::DuplicateRun { run_id: 1, .. }), "{error:?}");
652 assert_eq!(
653 error.to_string(),
654 format!("'{}' lists run 1 twice", shard_dir.display())
655 );
656 }
657}