Skip to main content

trellis_runner/watchers/
mod.rs

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
28/// Core observer trait for the engine event system.
29///
30/// Observers receive a stream of signal events during execution
31/// along with a read-only view over the iteration state
32///
33/// ### Design
34/// - Uses `&self` to support shared observers (`Arc`)
35/// - State mutation must use interior mutability if required
36/// - Must be object-safe to allow dynamic dispatch
37pub trait Observe<S: UserState>: Send + Sync {
38    fn observe(&self, name: &'static str, state: StateView<'_, S>, event: &EngineSignal<S::Float>);
39}
40
41/// Blanket implementation for shared observers.
42///
43/// This allows `Arc<T>` to be used directly as an observer.
44impl<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/// Frequency control for observer execution.
55///
56/// Allows observers to be:
57/// - always active
58/// - sampled at intervals
59/// - triggered only on termination
60/// - disabled entirely
61#[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
80/// Container for all registered observers.
81///
82/// Observers are stored as `(observer, frequency)` pairs and dispatched
83/// during engine execution.
84pub struct Observers<S> {
85    inner: Vec<(Arc<Mutex<dyn Observe<S>>>, Frequency)>,
86}
87
88impl<S: UserState> Observers<S> {
89    /// Create an empty observer set.
90    pub fn new() -> Self {
91        Self { inner: Vec::new() }
92    }
93
94    /// Attach a new observer with a frequency policy.
95    pub fn attach(&mut self, observer: Arc<Mutex<dyn Observe<S>>>, frequency: Frequency) {
96        self.inner.push((observer, frequency));
97    }
98
99    /// Dispatch an event to all eligible observers.
100    ///
101    /// Observers are filtered using their [`Frequency`] policy.
102    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}