1use serde::{Deserialize, Serialize};
21use thiserror::Error;
22
23pub const EVENT_TIME_RESOLUTION_S: f64 = 1e-12;
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
30#[serde(rename_all = "snake_case")]
31pub enum Direction {
32 Rising,
34 Falling,
36 Either,
38}
39
40impl Direction {
41 #[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#[derive(Debug, Clone, Copy, PartialEq, Error)]
59#[non_exhaustive]
60pub enum RootError {
61 #[error("the function is not finite at {x}")]
63 NotFinite {
64 x: f64,
66 },
67 #[error("the ends don't bracket a root")]
69 NotBracketed,
70 #[error("the tolerance must not be negative or NaN, not {tolerance}")]
72 Tolerance {
73 tolerance: f64,
75 },
76 #[error("no convergence within the iteration limit")]
78 NotConverged,
79}
80
81pub 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 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 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 p = 2.0 * m * s;
166 q = 1.0 - s;
167 } else {
168 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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}