1#![allow(clippy::type_complexity)]
2
3use crate::engine::EngineSignal;
4use crate::state::{StateView, UserState};
5
6use std::sync::{Arc, Mutex};
7
8#[cfg(feature = "writing")]
9mod csv_file;
10
11mod failure;
12mod metrics;
13
14#[cfg(feature = "plotting")]
15mod plot;
16
17mod sampler;
18mod tracing;
19
20#[cfg(feature = "writing")]
21pub use csv_file::CsvProgressWriter;
22
23#[cfg(feature = "plotting")]
24pub use plot::PlotObserver;
25
26pub use tracing::Tracer;
27
28pub trait Observe<S: UserState>: Send + Sync {
38 fn observe(&self, name: &'static str, state: StateView<'_, S>, event: &EngineSignal<S::Float>);
39}
40
41impl<S, T> Observe<S> for Arc<T>
45where
46 S: UserState,
47 T: Observe<S> + ?Sized,
48{
49 fn observe(&self, name: &'static str, state: StateView<'_, S>, event: &EngineSignal<S::Float>) {
50 (**self).observe(name, state, event)
51 }
52}
53
54#[derive(Copy, Clone, Debug, Eq, PartialEq)]
62pub enum Frequency {
63 Always,
64 Every(usize),
65 OnExit,
66 Never,
67}
68
69impl Frequency {
70 fn should_run<F>(&self, event: &EngineSignal<F>, iteration: usize) -> bool {
71 match self {
72 Self::Never => false,
73 Self::Always => true,
74 Self::OnExit => matches!(event, EngineSignal::Termination(_)),
75 Self::Every(n) => *n != 0 && iteration.is_multiple_of(*n),
76 }
77 }
78}
79
80pub struct Observers<S> {
85 inner: Vec<(Arc<Mutex<dyn Observe<S>>>, Frequency)>,
86}
87
88impl<S: UserState> Observers<S> {
89 pub fn new() -> Self {
91 Self { inner: Vec::new() }
92 }
93
94 pub fn attach(&mut self, observer: Arc<Mutex<dyn Observe<S>>>, frequency: Frequency) {
96 self.inner.push((observer, frequency));
97 }
98
99 pub fn dispatch(
103 &self,
104 ident: &'static str,
105 state: StateView<'_, S>,
106 event: &EngineSignal<S::Float>,
107 ) {
108 let iter = state.iteration();
109
110 for (obs, freq) in &self.inner {
111 if !freq.should_run(event, iter) {
112 continue;
113 }
114
115 let obs = obs.lock().unwrap();
116 obs.observe(ident, state, event);
117 }
118 }
119}
120
121#[cfg(test)]
122mod tests {
123 use super::*;
124
125 use crate::{
126 CancellationGuard, GenerateBuilder, MaxIterationPolicy, Procedure, Progress,
127 ProgressDiagnostics, Snapshotable, StateRestorer, UserState,
128 };
129
130 #[derive(Clone, Debug)]
131 pub struct DummyProblem {
132 pub target: f64,
133 }
134
135 #[derive(Clone, Debug, PartialEq)]
136 pub struct DummySnapshot {
137 pub value: f64,
138 pub steps: usize,
139 }
140
141 #[derive(Clone, Debug)]
142 pub struct DummyState {
143 pub value: f64,
144 pub steps: usize,
145 }
146
147 impl Default for DummyState {
148 fn default() -> Self {
149 Self {
150 value: 10.0,
151 steps: 0,
152 }
153 }
154 }
155
156 impl UserState for DummyState {
157 type Float = f64;
158
159 fn progress(&self) -> Progress<Self::Float> {
160 Progress::Report {
161 measure: self.value,
162 diagnostics: ProgressDiagnostics {
163 absolute_error: Some(self.value.abs()),
164 relative_error: Some(self.value.abs() / 10.0),
165 ..Default::default()
166 },
167 }
168 }
169 }
170
171 impl Snapshotable for DummyState {
172 type Snapshot = DummySnapshot;
173
174 fn snapshot(&self) -> Self::Snapshot {
175 DummySnapshot {
176 value: self.value,
177 steps: self.steps,
178 }
179 }
180 }
181
182 impl StateRestorer<DummyState> for DummyState {
183 fn restore(snapshot: DummySnapshot) -> Self {
184 Self {
185 value: snapshot.value,
186 steps: snapshot.steps,
187 }
188 }
189 }
190
191 pub struct DummyProcedure;
192
193 impl Procedure<DummyProblem> for DummyProcedure {
194 const NAME: &'static str = "Dummy Procedure";
195
196 type State = DummyState;
197 type Output = DummyState;
198
199 fn initialise(&self, _problem: &mut DummyProblem, _state: &mut Self::State) {}
200
201 fn step(
202 &self,
203 problem: &mut DummyProblem,
204 state: &mut Self::State,
205 _guard: CancellationGuard<'_>,
206 ) {
207 state.steps += 1;
208
209 let delta = state.value - problem.target;
210
211 if delta.abs() > 1e-12 {
212 state.value -= 0.5 * delta;
213 }
214 }
215
216 fn finalise(&self, _problem: &mut DummyProblem, state: &Self::State) -> Self::Output {
217 state.clone()
218 }
219 }
220
221 #[derive(Default)]
222 struct Spy {
223 events: std::sync::Mutex<Vec<&'static str>>,
224 }
225
226 impl<S> Observe<S> for Spy
227 where
228 S: UserState,
229 {
230 fn observe(
231 &self,
232 _name: &'static str,
233 _state: StateView<'_, S>,
234 event: &EngineSignal<S::Float>,
235 ) {
236 self.events.lock().unwrap().push(event.as_tag());
237 }
238 }
239
240 #[test]
241 fn observer_receives_lifecycle_events() {
242 let spy = std::sync::Arc::new(Spy::default());
243 let target = 1.0;
244
245 let _ = DummyProcedure
246 .build_for(DummyProblem { target })
247 .with_initial_state(DummyState::default())
248 .attach_observer(spy.clone(), Frequency::Always)
249 .and_policy(MaxIterationPolicy::new(3))
250 .finalise()
251 .run();
252
253 let events = spy.events.lock().unwrap();
254
255 assert!(events.contains(&"initialised"));
256 assert!(events.contains(&"progress"));
257 assert!(events.contains(&"termination"));
258 }
259
260 #[test]
261 fn observer_frequency_every_filters_progress_events() {
262 let spy = std::sync::Arc::new(Spy::default());
263
264 let target = 1.0;
265 let _ = DummyProcedure
266 .build_for(DummyProblem { target })
267 .with_initial_state(DummyState::default())
268 .attach_observer(spy.clone(), Frequency::Every(2))
269 .and_policy(MaxIterationPolicy::new(5))
270 .finalise()
271 .run();
272
273 let events = spy.events.lock().unwrap();
274
275 let progress_count = events.iter().filter(|e| **e == "progress").count();
276
277 assert_eq!(progress_count, 2);
278 }
279}