1use std::{
2 collections::BTreeMap,
3 fs::File,
4 io::BufReader,
5 path::{Path, PathBuf},
6};
7
8use ndarray::Array2;
9use serde::{Deserialize, Serialize};
10
11use crate::{
12 Complex64, Result,
13 array_serde::Array2Data,
14 configuration::SimulationConfiguration,
15 error::Error,
16 measurements::{ImageSet, MeasurementSpec, MeasurementStack},
17 reconstruction::ReconstructionProblem,
18};
19
20pub const DATASET_FORMAT_VERSION: u32 = 1;
22
23#[derive(Clone, Debug, Serialize, Deserialize)]
25#[serde(deny_unknown_fields)]
26pub struct DatasetManifest {
27 pub format_version: u32,
29 pub measurement_manifest: PathBuf,
31 pub configuration: PathBuf,
33 #[serde(default)]
35 pub ground_truth_object: Option<PathBuf>,
36 #[serde(default)]
38 pub valid_object_mask: Option<PathBuf>,
39 #[serde(default)]
40 pub provenance: BTreeMap<String, String>,
42 #[serde(default)]
43 pub measurement_units: Option<String>,
45}
46
47#[derive(Clone, Debug)]
49pub struct Dataset {
50 source_path: Option<PathBuf>,
51 measurements: MeasurementStack,
52 configuration: SimulationConfiguration,
53 ground_truth_object: Option<Array2<Complex64>>,
54 valid_object_mask: Option<Array2<u8>>,
55 provenance: BTreeMap<String, String>,
56 measurement_units: Option<String>,
57}
58
59impl Dataset {
60 pub fn new(
62 measurements: MeasurementStack,
63 configuration: SimulationConfiguration,
64 ) -> Result<Self> {
65 measurements.validate()?;
66 configuration.validate()?;
67 ReconstructionProblem::new(
68 measurements.clone(),
69 configuration.compiled_models.reconstruction_model.clone(),
70 )?;
71 Ok(Self {
72 source_path: None,
73 measurements,
74 configuration,
75 ground_truth_object: None,
76 valid_object_mask: None,
77 provenance: BTreeMap::new(),
78 measurement_units: None,
79 })
80 }
81
82 pub fn measurements(&self) -> &MeasurementStack {
84 &self.measurements
85 }
86
87 pub fn source_path(&self) -> Option<&Path> {
89 self.source_path.as_deref()
90 }
91
92 pub fn configuration(&self) -> &SimulationConfiguration {
94 &self.configuration
95 }
96
97 pub fn ground_truth_object(&self) -> Option<&Array2<Complex64>> {
99 self.ground_truth_object.as_ref()
100 }
101
102 pub fn valid_object_mask(&self) -> Option<&Array2<u8>> {
104 self.valid_object_mask.as_ref()
105 }
106
107 pub fn provenance(&self) -> &BTreeMap<String, String> {
109 &self.provenance
110 }
111
112 pub fn measurement_units(&self) -> Option<&str> {
114 self.measurement_units.as_deref()
115 }
116
117 pub fn reconstruction_problem(&self) -> Result<ReconstructionProblem<MeasurementStack>> {
119 ReconstructionProblem::new(
120 self.measurements.clone(),
121 self.configuration
122 .compiled_models
123 .reconstruction_model
124 .clone(),
125 )
126 }
127
128 pub fn subset(&self) -> super::DatasetSubsetBuilder<'_> {
130 super::DatasetSubsetBuilder::new(self)
131 }
132
133 fn with_metadata(
134 mut self,
135 ground_truth_object: Option<Array2<Complex64>>,
136 valid_object_mask: Option<Array2<u8>>,
137 provenance: BTreeMap<String, String>,
138 measurement_units: Option<String>,
139 ) -> Result<Self> {
140 validate_dataset_metadata(
141 &self.configuration,
142 ground_truth_object.as_ref(),
143 valid_object_mask.as_ref(),
144 &provenance,
145 measurement_units.as_deref(),
146 )?;
147 self.ground_truth_object = ground_truth_object;
148 self.valid_object_mask = valid_object_mask;
149 self.provenance = provenance;
150 self.measurement_units = measurement_units;
151 Ok(self)
152 }
153}
154
155#[derive(Clone, Debug)]
157pub struct DatasetLoader {
158 root: PathBuf,
159}
160
161impl DatasetLoader {
162 pub fn new(root: impl Into<PathBuf>) -> Result<Self> {
164 let root = root.into();
165 if !root.is_dir() {
166 return Err(Error::Dataset(format!(
167 "dataset path does not exist or is not a directory: {}",
168 root.display()
169 )));
170 }
171 Ok(Self { root })
172 }
173
174 pub fn root(&self) -> &Path {
176 &self.root
177 }
178
179 pub fn manifest_path(&self) -> PathBuf {
181 self.root.join("dataset.json")
182 }
183
184 pub fn load(&self) -> Result<Dataset> {
187 let manifest_path = self.manifest_path();
188 let manifest: DatasetManifest =
189 serde_json::from_reader(BufReader::new(File::open(&manifest_path)?))?;
190 if manifest.format_version != DATASET_FORMAT_VERSION {
191 return Err(Error::Dataset(format!(
192 "unsupported dataset format version {} in {}; expected {}",
193 manifest.format_version,
194 manifest_path.display(),
195 DATASET_FORMAT_VERSION
196 )));
197 }
198 validate_dataset_manifest_paths(&manifest)?;
199 let measurement_manifest_path = self.root.join(&manifest.measurement_manifest);
200 let measurement_spec = MeasurementSpec::load(&measurement_manifest_path)?;
201 validate_measurement_paths(&measurement_spec)?;
202 let measurement_base = measurement_manifest_path
203 .parent()
204 .unwrap_or_else(|| Path::new("."));
205 let measurements =
206 MeasurementStack::from_manifest_definition(measurement_spec, measurement_base)?;
207 let configuration = SimulationConfiguration::load(self.root.join(&manifest.configuration))?;
208 let ground_truth_object = manifest
209 .ground_truth_object
210 .as_ref()
211 .map(|path| load_json_array(self.root.join(path), "ground-truth object"))
212 .transpose()?;
213 let valid_object_mask = manifest
214 .valid_object_mask
215 .as_ref()
216 .map(|path| load_json_array(self.root.join(path), "valid-object mask"))
217 .transpose()?;
218 let mut dataset = Dataset::new(measurements, configuration)?;
219 dataset.source_path = Some(self.root.clone());
220 dataset.with_metadata(
221 ground_truth_object,
222 valid_object_mask,
223 manifest.provenance,
224 manifest.measurement_units,
225 )
226 }
227}
228
229fn validate_dataset_manifest_paths(manifest: &DatasetManifest) -> Result<()> {
230 validate_contained_path("measurement manifest", &manifest.measurement_manifest)?;
231 validate_contained_path("configuration", &manifest.configuration)?;
232 if let Some(path) = &manifest.ground_truth_object {
233 validate_contained_path("ground-truth object", path)?;
234 }
235 if let Some(path) = &manifest.valid_object_mask {
236 validate_contained_path("valid-object mask", path)?;
237 }
238 Ok(())
239}
240
241fn validate_measurement_paths(spec: &MeasurementSpec) -> Result<()> {
242 for frame in &spec.frames {
243 validate_contained_path("measurement frame", &frame.path)?;
244 }
245 if let Some(path) = &spec.dark_frame {
246 validate_contained_path("dark frame", path)?;
247 }
248 if let Some(path) = &spec.flat_field {
249 validate_contained_path("flat field", path)?;
250 }
251 if let Some(images) = &spec.background {
252 validate_image_set_paths("background", images)?;
253 }
254 if let Some(images) = &spec.mask {
255 validate_image_set_paths("mask", images)?;
256 }
257 Ok(())
258}
259
260fn validate_image_set_paths(label: &str, images: &ImageSet) -> Result<()> {
261 match images {
262 ImageSet::Single(path) => validate_contained_path(label, path),
263 ImageSet::PerFrame(paths) => paths
264 .iter()
265 .try_for_each(|path| validate_contained_path(label, path)),
266 }
267}
268
269fn validate_contained_path(label: &str, path: &Path) -> Result<()> {
270 if path.as_os_str().is_empty()
271 || path
272 .components()
273 .any(|component| !matches!(component, std::path::Component::Normal(_)))
274 {
275 return Err(Error::Dataset(format!(
276 "{label} must be a safe relative path contained by its manifest directory: {}",
277 path.display()
278 )));
279 }
280 Ok(())
281}
282
283fn load_json_array<T>(path: PathBuf, label: &str) -> Result<Array2<T>>
284where
285 T: for<'de> Deserialize<'de>,
286{
287 let file = File::open(&path).map_err(|error| {
288 Error::Dataset(format!(
289 "failed to open {label} at {}: {error}",
290 path.display()
291 ))
292 })?;
293 let representation: Array2Data<T> =
294 serde_json::from_reader(BufReader::new(file)).map_err(|error| {
295 Error::Dataset(format!(
296 "failed to parse {label} at {}: {error}",
297 path.display()
298 ))
299 })?;
300 representation
301 .into_array()
302 .map_err(|error| Error::Dataset(format!("invalid {label} at {}: {error}", path.display())))
303}
304
305fn validate_dataset_metadata(
306 configuration: &SimulationConfiguration,
307 ground_truth_object: Option<&Array2<Complex64>>,
308 valid_object_mask: Option<&Array2<u8>>,
309 provenance: &BTreeMap<String, String>,
310 measurement_units: Option<&str>,
311) -> Result<()> {
312 if let Some(ground_truth) = ground_truth_object {
313 if ground_truth.dim() != configuration.reconstruction_shape {
314 return Err(Error::Dataset(format!(
315 "ground-truth shape {:?} differs from reconstruction shape {:?}",
316 ground_truth.dim(),
317 configuration.reconstruction_shape
318 )));
319 }
320 if ground_truth
321 .iter()
322 .any(|value| !value.re.is_finite() || !value.im.is_finite())
323 {
324 return Err(Error::Dataset(
325 "ground-truth object contains non-finite values".into(),
326 ));
327 }
328 }
329 if let Some(mask) = valid_object_mask {
330 if ground_truth_object.is_none() {
331 return Err(Error::Dataset(
332 "a valid-object mask requires a ground-truth object".into(),
333 ));
334 }
335 if mask.dim() != configuration.reconstruction_shape {
336 return Err(Error::Dataset(format!(
337 "valid-object mask shape {:?} differs from reconstruction shape {:?}",
338 mask.dim(),
339 configuration.reconstruction_shape
340 )));
341 }
342 if mask.iter().any(|&value| value > 1) || !mask.iter().any(|&value| value == 1) {
343 return Err(Error::Dataset(
344 "valid-object mask must contain only zero and one and select at least one pixel"
345 .into(),
346 ));
347 }
348 }
349 if provenance
350 .iter()
351 .any(|(key, value)| key.trim().is_empty() || value.trim().is_empty())
352 {
353 return Err(Error::Dataset(
354 "dataset provenance keys and values must not be empty".into(),
355 ));
356 }
357 if measurement_units.is_some_and(|units| units.trim().is_empty()) {
358 return Err(Error::Dataset(
359 "measurement units must not be empty when provided".into(),
360 ));
361 }
362 Ok(())
363}