Skip to main content

fpm_rs/datasets/
loader.rs

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
20/// Current version of the language-neutral dataset bundle format.
21pub const DATASET_FORMAT_VERSION: u32 = 1;
22
23/// The `dataset.json` entry point defined by `dataset_spec.md`.
24#[derive(Clone, Debug, Serialize, Deserialize)]
25#[serde(deny_unknown_fields)]
26pub struct DatasetManifest {
27    /// Dataset-bundle schema version; must equal [`DATASET_FORMAT_VERSION`].
28    pub format_version: u32,
29    /// Relative path to a [`crate::measurements::MeasurementSpec`] JSON file.
30    pub measurement_manifest: PathBuf,
31    /// Relative path to a [`crate::configuration::SimulationConfiguration`] JSON file.
32    pub configuration: PathBuf,
33    /// Optional JSON-serialized `Array2<Complex64>` in reconstruction space.
34    #[serde(default)]
35    pub ground_truth_object: Option<PathBuf>,
36    /// Optional JSON-serialized binary `Array2<u8>` in reconstruction space.
37    #[serde(default)]
38    pub valid_object_mask: Option<PathBuf>,
39    #[serde(default)]
40    /// User-visible provenance key/value metadata.
41    pub provenance: BTreeMap<String, String>,
42    #[serde(default)]
43    /// Optional semantic units for stored measurements, such as camera counts.
44    pub measurement_units: Option<String>,
45}
46
47/// A validated dataset loaded from a bundle conforming to `dataset_spec.md`.
48#[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    /// Constructs a dataset from already-loaded values.
61    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    /// Borrows the resident measured-intensity stack.
83    pub fn measurements(&self) -> &MeasurementStack {
84        &self.measurements
85    }
86
87    /// Root directory of the loaded bundle, or `None` for programmatic data.
88    pub fn source_path(&self) -> Option<&Path> {
89        self.source_path.as_deref()
90    }
91
92    /// Borrows the validated true/assumed experiment configuration.
93    pub fn configuration(&self) -> &SimulationConfiguration {
94        &self.configuration
95    }
96
97    /// Borrows optional high-resolution complex ground truth shaped `(height, width)`.
98    pub fn ground_truth_object(&self) -> Option<&Array2<Complex64>> {
99        self.ground_truth_object.as_ref()
100    }
101
102    /// Borrows an optional same-shaped binary validity mask for object metrics.
103    pub fn valid_object_mask(&self) -> Option<&Array2<u8>> {
104        self.valid_object_mask.as_ref()
105    }
106
107    /// Borrows provenance metadata from the dataset manifest.
108    pub fn provenance(&self) -> &BTreeMap<String, String> {
109        &self.provenance
110    }
111
112    /// Borrows optional semantic measurement units.
113    pub fn measurement_units(&self) -> Option<&str> {
114        self.measurement_units.as_deref()
115    }
116
117    /// Clones measurements and the assumed model into a validated reconstruction problem.
118    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    /// Starts a deterministic frame-selection and spatial-cropping builder.
129    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/// Loads an already-converted dataset bundle from a local directory.
156#[derive(Clone, Debug)]
157pub struct DatasetLoader {
158    root: PathBuf,
159}
160
161impl DatasetLoader {
162    /// Opens an existing local bundle root without performing network access.
163    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    /// Returns the local dataset-bundle root.
175    pub fn root(&self) -> &Path {
176        &self.root
177    }
178
179    /// Returns `root/dataset.json` without checking whether the file exists.
180    pub fn manifest_path(&self) -> PathBuf {
181        self.root.join("dataset.json")
182    }
183
184    /// Loads and validates the manifest, measurements, configuration, optional ground truth,
185    /// mask, provenance, and measurement units entirely from local files.
186    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}