Skip to main content

hpr_sim/
events.rs

1//! Event directions and their location: zero crossings of `g(t, y)` found on an integrator's dense
2//! output and polished with Brent's method.
3//!
4//! An event is a sign change of a scalar function `g(t, y)` across an accepted step, in the
5//! direction the event asks for. A system declares its events through
6//! [`crate::integrator::OdeSystem::event_count`] and its sibling methods.
7//! [`crate::integrator::Integrator::advance`] evaluates every event function at the end of each
8//! accepted step and locates the earliest crossing on the step's dense output with [`find_root`]
9//! to [`EVENT_TIME_RESOLUTION_S`]. The integration stops at the end of the final bracket on the far
10//! side of the crossing, with the dense output's state there. Every event past its zero at that
11//! state is reported together, so coincident events are never lost, and none of them is reported
12//! again when the integration resumes.
13//!
14//! Crossings are seen only as sign changes between step ends: a function that crosses zero and
15//! comes back inside one step goes unseen. Bound the step (`Adaptive::max_step_s`) when that
16//! matters.
17//!
18//! Method: `docs/physics/integration.md`.
19
20use serde::{Deserialize, Serialize};
21use thiserror::Error;
22
23/// The time resolution of event location, in seconds. Roots are polished to within this (plus
24/// four units of rounding in `t`) of the dense output's zero; the accuracy of the event itself is
25/// the integration's.
26pub const EVENT_TIME_RESOLUTION_S: f64 = 1e-12;
27
28/// Which sign changes of `g` count as an event.
29#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
30#[serde(rename_all = "snake_case")]
31pub enum Direction {
32    /// `g` goes from negative to zero or positive.
33    Rising,
34    /// `g` goes from positive to zero or negative.
35    Falling,
36    /// Either of the two.
37    Either,
38}
39
40impl Direction {
41    /// Whether `g` moving from `g0` to `g1` is a crossing in this direction.
42    ///
43    /// The start must be strictly on one side: a `g0` of exactly zero is an event that has already
44    /// happened, so a step that starts on the root doesn't report it again.
45    #[must_use]
46    pub fn crosses(self, g0: f64, g1: f64) -> bool {
47        let rising = g0 < 0.0 && g1 >= 0.0;
48        let falling = g0 > 0.0 && g1 <= 0.0;
49        match self {
50            Self::Rising => rising,
51            Self::Falling => falling,
52            Self::Either => rising || falling,
53        }
54    }
55}
56
57/// Why [`find_root`] found no root.
58#[derive(Debug, Clone, Copy, PartialEq, Error)]
59#[non_exhaustive]
60pub enum RootError {
61    /// The function or an end is not finite at `x`.
62    #[error("the function is not finite at {x}")]
63    NotFinite {
64        /// Where.
65        x: f64,
66    },
67    /// The ends have the same strict sign.
68    #[error("the ends don't bracket a root")]
69    NotBracketed,
70    /// The tolerance is negative or not a number.
71    #[error("the tolerance must not be negative or NaN, not {tolerance}")]
72    Tolerance {
73        /// The tolerance.
74        tolerance: f64,
75    },
76    /// The bracket didn't shrink to the tolerance within the iteration limit.
77    #[error("no convergence within the iteration limit")]
78    NotConverged,
79}
80
81/// A zero of `f` in `[a, b]` by Brent's method, given `f(a)` and `f(b)` of opposite signs (or one
82/// of them zero).
83///
84/// Brent's algorithm (R. P. Brent, *Algorithms for Minimization without Derivatives*,
85/// Prentice-Hall, 1973, ch. 4) combines bisection, the secant rule and inverse quadratic
86/// interpolation. It keeps a bracket `[b, c]` with `|f(b)| ≤ |f(c)|` and stops when
87/// `|c − b|/2 ≤ 2ε|b| + tolerance/2`. It converges superlinearly on smooth functions and falls
88/// back on bisection otherwise.
89///
90/// The result is the end of the final bracket on `b`'s side of the root: an evaluated point whose
91/// sign matches `f(b)`, or a point where `f` is zero, within about twice the tolerance of the zero.
92/// Callers that stop at an event use it to stand past the crossing.
93///
94/// # Errors
95///
96/// [`RootError`]: a non-finite value, ends that don't bracket, a bad tolerance, or no convergence
97/// in 500 iterations.
98pub fn find_root(
99    mut f: impl FnMut(f64) -> f64,
100    a: f64,
101    b: f64,
102    fa: f64,
103    fb: f64,
104    tolerance: f64,
105) -> Result<f64, RootError> {
106    if tolerance.is_nan() || tolerance < 0.0 {
107        return Err(RootError::Tolerance { tolerance });
108    }
109    for (x, fx) in [(a, fa), (b, fb)] {
110        if !x.is_finite() || !fx.is_finite() {
111            return Err(RootError::NotFinite { x });
112        }
113    }
114    if fa == 0.0 {
115        return Ok(a);
116    }
117    if fb == 0.0 {
118        return Ok(b);
119    }
120    if (fa > 0.0) == (fb > 0.0) {
121        return Err(RootError::NotBracketed);
122    }
123    let (mut a, mut b, mut fa, mut fb) = (a, b, fa, fb);
124    let far_side_positive = fb > 0.0;
125    // The bracket end on the far side: `b` or `c`, whichever has `f(b)`'s original sign. The
126    // bracket always holds one end of each sign, so this is never on the near side.
127    let far = |b: f64, fb: f64, c: f64| {
128        if fb == 0.0 || (fb > 0.0) == far_side_positive {
129            b
130        } else {
131            c
132        }
133    };
134    let (mut c, mut fc) = (a, fa);
135    let mut d = b - a;
136    let mut e = d;
137    // Brent's safeguards bisect whenever interpolation fails to shrink the bracket enough, so the
138    // iterations are bounded by a small multiple of bisection's (at most about 2100 halvings for
139    // any f64 interval, about 50 for an integration step).
140    for _ in 0..500 {
141        if (fb > 0.0) == (fc > 0.0) {
142            c = a;
143            fc = fa;
144            d = b - a;
145            e = d;
146        }
147        if fc.abs() < fb.abs() {
148            a = b;
149            b = c;
150            c = a;
151            fa = fb;
152            fb = fc;
153            fc = fa;
154        }
155        let tol = 2.0 * f64::EPSILON * b.abs() + 0.5 * tolerance;
156        let m = 0.5 * (c - b);
157        if m.abs() <= tol || fb == 0.0 {
158            return Ok(far(b, fb, c));
159        }
160        if e.abs() >= tol && fa.abs() > fb.abs() {
161            let s = fb / fa;
162            let (mut p, mut q);
163            if a == c {
164                // Secant step.
165                p = 2.0 * m * s;
166                q = 1.0 - s;
167            } else {
168                // Inverse quadratic interpolation.
169                let qa = fa / fc;
170                let r = fb / fc;
171                p = s * (2.0 * m * qa * (qa - r) - (b - a) * (r - 1.0));
172                q = (qa - 1.0) * (r - 1.0) * (s - 1.0);
173            }
174            if p > 0.0 {
175                q = -q;
176            } else {
177                p = -p;
178            }
179            if 2.0 * p < (3.0 * m * q - (tol * q).abs()).min((e * q).abs()) {
180                e = d;
181                d = p / q;
182            } else {
183                d = m;
184                e = m;
185            }
186        } else {
187            d = m;
188            e = m;
189        }
190        a = b;
191        fa = fb;
192        b += if d.abs() > tol { d } else { tol.copysign(m) };
193        fb = f(b);
194        if !fb.is_finite() {
195            return Err(RootError::NotFinite { x: b });
196        }
197    }
198    Err(RootError::NotConverged)
199}
200
201#[cfg(test)]
202mod tests {
203    use super::*;
204    use crate::integrator::{Advance, Integrator, Method, OdeSystem};
205    use crate::testing::{Oscillator, QuadraticDragFall, WithEvents, closed_form_quadratic_drag};
206
207    #[test]
208    fn directions_need_a_strict_start_side() {
209        assert!(Direction::Rising.crosses(-1.0, 0.0));
210        assert!(Direction::Rising.crosses(-1.0, 2.0));
211        assert!(!Direction::Rising.crosses(0.0, 1.0));
212        assert!(!Direction::Rising.crosses(1.0, -1.0));
213        assert!(Direction::Falling.crosses(1.0, 0.0));
214        assert!(!Direction::Falling.crosses(0.0, -1.0));
215        assert!(Direction::Either.crosses(1.0, -1.0));
216        assert!(Direction::Either.crosses(-1.0, 1.0));
217        assert!(!Direction::Either.crosses(1.0, 1.0));
218        assert!(!Direction::Either.crosses(f64::NAN, 1.0));
219    }
220
221    #[test]
222    fn brent_finds_smooth_flat_and_discontinuous_roots() {
223        // Smooth: cos x = x at 0.739085133215160641655312087673873404...
224        let x = find_root(|x| x.cos() - x, 0.0, 1.0, 1.0, 1.0_f64.cos() - 1.0, 1e-15).unwrap();
225        assert!((x - 0.739_085_133_215_160_6).abs() < 1e-14, "{x}");
226
227        // A triple root, where interpolation stalls and bisection has to carry it.
228        let cube = |x: f64| (x - 0.3).powi(3);
229        let x = find_root(cube, 0.0, 1.0, cube(0.0), cube(1.0), 1e-12).unwrap();
230        assert!((x - 0.3).abs() < 1e-11, "{x}");
231
232        // A step: the "root" is the jump, found to the tolerance, on the far side.
233        let step = |x: f64| if x < 0.123_456 { -1.0 } else { 1.0 };
234        let mut calls = 0;
235        let x = find_root(
236            |x| {
237                calls += 1;
238                step(x)
239            },
240            0.0,
241            1.0,
242            -1.0,
243            1.0,
244            1e-12,
245        )
246        .unwrap();
247        assert!((x - 0.123_456).abs() < 1e-12, "{x}");
248        assert_eq!(step(x), 1.0, "the far side of the jump");
249        assert!(calls < 60, "{calls} evaluations");
250        // Reversed signs: the far side is now the negative one.
251        let x = find_root(|x| -step(x), 0.0, 1.0, 1.0, -1.0, 1e-12).unwrap();
252        assert_eq!(-step(x), -1.0);
253
254        // Ends that are roots, ends that don't bracket, a function that turns NaN, and bad
255        // tolerances.
256        assert_eq!(find_root(|x| x, 0.0, 1.0, 0.0, 1.0, 1e-12), Ok(0.0));
257        assert_eq!(find_root(|x| x - 1.0, 0.0, 1.0, -1.0, 0.0, 1e-12), Ok(1.0));
258        assert_eq!(
259            find_root(|x| x + 1.0, 0.0, 1.0, 1.0, 2.0, 1e-12),
260            Err(RootError::NotBracketed)
261        );
262        assert!(matches!(
263            find_root(|_| f64::NAN, -1.0, 1.0, -1.0, 1.0, 1e-12),
264            Err(RootError::NotFinite { .. })
265        ));
266        assert!(matches!(
267            find_root(|x| x, -1.0, 1.0, -1.0, 1.0, f64::NAN),
268            Err(RootError::Tolerance { .. })
269        ));
270    }
271
272    /// Runs to `t_stop`, collecting every fired event with its time and state.
273    fn collect_events<S: OdeSystem<N>, const N: usize>(
274        integrator: &mut Integrator<N>,
275        system: &mut S,
276        t_stop: f64,
277    ) -> Vec<(usize, f64, [f64; N])>
278    where
279        S::Error: std::fmt::Debug,
280    {
281        let mut found = Vec::new();
282        loop {
283            match integrator.advance(system, t_stop).unwrap() {
284                Advance::Reached => return found,
285                Advance::Events => {
286                    for index in integrator.fired_events() {
287                        found.push((*index, integrator.time_s(), *integrator.state()));
288                    }
289                }
290                other => panic!("{other:?}"),
291            }
292        }
293    }
294
295    #[test]
296    fn oscillator_crossings_are_located_within_1e_6_s_by_both_methods() {
297        // x = cos t crosses zero at π/2 + kπ: falling at even k. The events are falling zeros of
298        // x and every extremum (zeros of x').
299        for method in [Method::default(), Method::Rk4 { step_s: 0.01 }] {
300            let mut system = WithEvents::new(
301                Oscillator,
302                vec![Direction::Falling, Direction::Either],
303                |i: usize, _t: f64, y: &[f64; 2]| y[i],
304            );
305            let mut integrator = Integrator::new(method, 0.0, [1.0, 0.0]).unwrap();
306            let found = collect_events(&mut integrator, &mut system, 20.0);
307            let pi = std::f64::consts::PI;
308            let falling: Vec<f64> = (0..3).map(|k| pi / 2.0 + 2.0 * pi * f64::from(k)).collect();
309            let extrema: Vec<f64> = (1..7).map(|k| pi * f64::from(k)).collect();
310            let got = |index: usize| -> Vec<f64> {
311                found
312                    .iter()
313                    .filter(|(i, _, _)| *i == index)
314                    .map(|(_, t, _)| *t)
315                    .collect()
316            };
317            for (expected, actual) in [(falling, got(0)), (extrema, got(1))] {
318                assert_eq!(expected.len(), actual.len(), "{method:?}: {found:?}");
319                for (e, a) in expected.iter().zip(&actual) {
320                    assert!((e - a).abs() <= 1e-6, "{method:?}: {a} vs {e}");
321                }
322            }
323        }
324    }
325
326    #[test]
327    fn apogee_landing_and_altitude_deploy_located_within_1e_6_s() {
328        // Loft lesson L22: events were never root-found, so apogee was quantised to the step and
329        // an altitude deploy overshot by v·dt. A vertical flight with quadratic drag has closed
330        // forms for all three events (`testing::closed_form_quadratic_drag`), and its drag term
331        // `v|v|` is not smooth at apogee, which a polynomial test would hide.
332        let flight = QuadraticDragFall::example();
333        let truth = closed_form_quadratic_drag(&flight, 150.0);
334        let deploy_m = 300.0;
335        let g = move |i: usize, _t: f64, y: &[f64; 2]| match i {
336            0 => y[1],
337            1 => y[0] - deploy_m,
338            _ => y[0],
339        };
340        for method in [Method::default(), Method::Rk4 { step_s: 0.01 }] {
341            let mut system = WithEvents::new(flight.clone(), vec![Direction::Falling; 3], g);
342            let mut integrator = Integrator::new(method, 0.0, [0.0, 150.0]).unwrap();
343            let landing_s = truth.time_at_descending_height_s(0.0);
344            let found = collect_events(&mut integrator, &mut system, landing_s + 1.0);
345            // Landing is the last event; the flight continues underground to the stop time.
346            let [apogee, deploy, landing] = found.as_slice() else {
347                panic!("{method:?}: {found:?}");
348            };
349            let expected = [
350                (0, truth.apogee_s),
351                (1, truth.time_at_descending_height_s(deploy_m)),
352                (2, landing_s),
353            ];
354            for ((index, t, y), (want_index, want_t)) in
355                [apogee, deploy, landing].into_iter().zip(expected)
356            {
357                assert_eq!(*index, want_index, "{method:?}");
358                assert!(
359                    (t - want_t).abs() <= 1e-6,
360                    "{method:?} event {index}: {t} vs {want_t}"
361                );
362                // The state is the dense output's at the stop: on the root to the integration's
363                // accuracy.
364                let value = g(*index, *t, y);
365                assert!(value.abs() < 1e-5, "{method:?} event {index}: g = {value}");
366            }
367            assert!((apogee.2[0] - truth.apogee_m).abs() < 1e-6, "{method:?}");
368            let landing_speed = truth.state(landing.1)[1];
369            assert!((landing.2[1] - landing_speed).abs() < 1e-6, "{method:?}");
370        }
371    }
372
373    #[test]
374    fn coincident_events_are_all_reported_once() {
375        // Two identical apogee events and a third that crosses at the same instant but is written
376        // differently: all three fire together, once per crossing.
377        let g = |i: usize, _t: f64, y: &[f64; 2]| match i {
378            0 | 1 => y[1],
379            _ => 2.0 * y[1],
380        };
381        for method in [Method::default(), Method::Rk4 { step_s: 0.01 }] {
382            let mut system = WithEvents::new(Oscillator, vec![Direction::Falling; 3], g);
383            let mut integrator = Integrator::new(method, 0.0, [0.0, 1.0]).unwrap();
384            let found = collect_events(&mut integrator, &mut system, 10.0);
385            let indices: Vec<usize> = found.iter().map(|(i, _, _)| *i).collect();
386            assert_eq!(indices, [0, 1, 2, 0, 1, 2], "{method:?}: {found:?}");
387            let pi = std::f64::consts::PI;
388            for (k, (_, t, _)) in found.iter().enumerate() {
389                let want = pi / 2.0 + 2.0 * pi * (k / 3) as f64;
390                assert!((t - want).abs() < 1e-6, "{method:?}: {found:?}");
391            }
392        }
393    }
394
395    #[test]
396    fn event_times_are_resolved_to_the_root_finder_tolerance() {
397        // On y' = 1 both dense outputs are exact, so the only error left is the root finder's.
398        // The event functions are nonlinear in t, where a secant step is not exact.
399        struct Clock;
400        impl OdeSystem<1> for Clock {
401            type Error = std::convert::Infallible;
402            fn derivative(&mut self, _t: f64, _y: &[f64; 1]) -> Result<[f64; 1], Self::Error> {
403                Ok([1.0])
404            }
405        }
406        let g = |i: usize, _t: f64, y: &[f64; 1]| match i {
407            0 => y[0].powi(3) - 0.3,
408            _ => y[0].sin() - 0.5,
409        };
410        // Each g crosses once on [0, 2], so no step can hide a return.
411        let roots = [0.3_f64.cbrt(), std::f64::consts::FRAC_PI_6];
412        for method in [Method::default(), Method::Rk4 { step_s: 0.9 }] {
413            for (index, root) in roots.into_iter().enumerate() {
414                let mut system = WithEvents::new(
415                    Clock,
416                    vec![Direction::Rising],
417                    move |_: usize, t: f64, y: &[f64; 1]| g(index, t, y),
418                );
419                let mut integrator = Integrator::new(method, 0.0, [0.0]).unwrap();
420                assert_eq!(integrator.advance(&mut system, 2.0), Ok(Advance::Events));
421                let t = integrator.time_s();
422                let error = t - root;
423                assert!(
424                    (-1e-15..=2.5e-12).contains(&error),
425                    "{method:?} event {index}: {error:e} past the root"
426                );
427                assert!(g(index, t, integrator.state()) >= 0.0, "on the far side");
428            }
429        }
430    }
431
432    #[test]
433    fn a_restart_on_an_event_does_not_report_it_again() {
434        let mut system = WithEvents::new(
435            Oscillator,
436            vec![Direction::Either],
437            |_: usize, _t: f64, y: &[f64; 2]| y[0],
438        );
439        let mut integrator = Integrator::new(Method::default(), 0.0, [1.0, 0.0]).unwrap();
440        let mut times = Vec::new();
441        for _ in 0..3 {
442            let outcome = integrator.advance(&mut system, 100.0).unwrap();
443            assert_eq!(outcome, Advance::Events);
444            assert_eq!(integrator.fired_events(), [0]);
445            times.push(integrator.time_s());
446        }
447        let pi = std::f64::consts::PI;
448        for (k, t) in times.iter().enumerate() {
449            let want = pi / 2.0 + pi * k as f64;
450            assert!((t - want).abs() < 1e-6, "{times:?}");
451        }
452    }
453
454    #[test]
455    fn the_earliest_of_several_crossings_in_one_step_wins() {
456        // One long RK4 step holds crossings of x = t − 0.7, x = t − 0.2 and x = t − 0.5; the
457        // integration stops at 0.2 and then finds 0.5 and 0.7 in turn.
458        struct Clock;
459        impl OdeSystem<1> for Clock {
460            type Error = std::convert::Infallible;
461            fn derivative(&mut self, _t: f64, _y: &[f64; 1]) -> Result<[f64; 1], Self::Error> {
462                Ok([1.0])
463            }
464        }
465        let offsets = [0.7, 0.2, 0.5];
466        let mut system = WithEvents::new(
467            Clock,
468            vec![Direction::Rising; 3],
469            |i: usize, _t: f64, y: &[f64; 1]| y[0] - offsets[i],
470        );
471        let mut integrator = Integrator::new(Method::Rk4 { step_s: 10.0 }, 0.0, [0.0]).unwrap();
472        let order = collect_events(&mut integrator, &mut system, 5.0);
473        assert_eq!(order.len(), 3, "{order:?}");
474        for ((index, t, _), (want_index, want_t)) in
475            order.iter().zip([(1, 0.2), (2, 0.5), (0, 0.7)])
476        {
477            assert_eq!(*index, want_index);
478            assert!((t - want_t).abs() < 1e-12, "{order:?}");
479        }
480        assert_eq!(integrator.time_s(), 5.0);
481    }
482}