Skip to main content

fpm_rs/callbacks/
csv_logger.rs

1use std::{fs::File, path::PathBuf};
2
3use crate::{Result, diagnostics::DiagnosticRequest};
4
5use super::{Callback, CallbackAction, CallbackHook, StepContext};
6
7/// Streams universal iteration history to a CSV file.
8pub struct CsvLogger {
9    path: PathBuf,
10    writer: Option<csv::Writer<File>>,
11}
12
13impl CsvLogger {
14    /// Creates a logger that opens or replaces `path` when reconstruction starts.
15    pub fn new(path: impl Into<PathBuf>) -> Self {
16        Self {
17            path: path.into(),
18            writer: None,
19        }
20    }
21}
22
23impl Callback for CsvLogger {
24    fn requires(&self) -> Vec<DiagnosticRequest> {
25        vec![DiagnosticRequest::Objective]
26    }
27
28    fn requires_for(&self, hook: CallbackHook, _iteration: usize) -> Vec<DiagnosticRequest> {
29        if hook == CallbackHook::IterationEnd {
30            self.requires()
31        } else {
32            Vec::new()
33        }
34    }
35
36    fn on_start(&mut self, _context: &StepContext<'_>) -> Result<CallbackAction> {
37        let mut writer = csv::Writer::from_path(&self.path)?;
38        writer.write_record(["iteration", "objective", "elapsed_seconds"])?;
39        self.writer = Some(writer);
40        Ok(CallbackAction::Continue)
41    }
42
43    fn on_iteration_end(&mut self, context: &StepContext<'_>) -> Result<CallbackAction> {
44        if let (Some(writer), Some(record)) = (&mut self.writer, context.trace.iterations.last()) {
45            writer.serialize((record.iteration, record.objective, record.elapsed_seconds))?;
46            writer.flush()?;
47        }
48        Ok(CallbackAction::Continue)
49    }
50}