fpm_rs/callbacks/
csv_logger.rs1use std::{fs::File, path::PathBuf};
2
3use crate::{Result, diagnostics::DiagnosticRequest};
4
5use super::{Callback, CallbackAction, CallbackHook, StepContext};
6
7pub struct CsvLogger {
9 path: PathBuf,
10 writer: Option<csv::Writer<File>>,
11}
12
13impl CsvLogger {
14 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}