1use serde::{Deserialize, Serialize};
30
31use crate::error::CoreError;
32
33#[expect(
36 clippy::excessive_precision,
37 reason = "the published digits, which round to the nearest f64"
38)]
39const NODES: [f64; 8] = [
40 0.991_455_371_120_812_639_206_854_697_526_329,
41 0.949_107_912_342_758_524_526_189_684_047_851,
42 0.864_864_423_359_769_072_789_712_788_640_926,
43 0.741_531_185_599_394_439_863_864_773_280_788,
44 0.586_087_235_467_691_130_294_144_845_693_013,
45 0.405_845_151_377_397_166_906_606_412_076_961,
46 0.207_784_955_007_898_467_600_689_403_773_245,
47 0.0,
48];
49
50#[expect(
52 clippy::excessive_precision,
53 reason = "the published digits, which round to the nearest f64"
54)]
55const KRONROD_WEIGHTS: [f64; 8] = [
56 0.022_935_322_010_529_224_963_732_008_058_970,
57 0.063_092_092_629_978_553_290_700_663_189_204,
58 0.104_790_010_322_250_183_839_876_322_541_518,
59 0.140_653_259_715_525_918_745_189_590_510_238,
60 0.169_004_726_639_267_902_826_583_426_598_550,
61 0.190_350_578_064_785_409_913_256_402_421_014,
62 0.204_432_940_075_298_892_414_161_999_234_649,
63 0.209_482_141_084_727_828_012_999_174_891_714,
64];
65
66#[expect(
68 clippy::excessive_precision,
69 reason = "the published digits, which round to the nearest f64"
70)]
71const GAUSS_WEIGHTS: [f64; 4] = [
72 0.129_484_966_168_869_693_270_611_432_679_082,
73 0.279_705_391_489_276_667_901_467_771_423_780,
74 0.381_830_050_505_118_944_950_369_775_488_975,
75 0.417_959_183_673_469_387_755_102_040_816_327,
76];
77
78#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
80pub struct Tolerance {
81 pub relative: f64,
83 pub absolute: f64,
85 pub max_intervals: usize,
87}
88
89impl Default for Tolerance {
90 fn default() -> Self {
93 Self {
94 relative: 1e-12,
95 absolute: 1e-12,
96 max_intervals: 4000,
97 }
98 }
99}
100
101#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
103pub struct Integral<const N: usize> {
104 #[serde(with = "serde_arrays")]
106 pub value: [f64; N],
107 #[serde(with = "serde_arrays")]
109 pub error: [f64; N],
110 pub intervals: usize,
112}
113
114mod serde_arrays {
116 use serde::de::Error as _;
117 use serde::{Deserialize, Deserializer, Serialize, Serializer};
118
119 pub fn serialize<S: Serializer, const N: usize>(
120 values: &[f64; N],
121 serializer: S,
122 ) -> Result<S::Ok, S::Error> {
123 values.as_slice().serialize(serializer)
124 }
125
126 pub fn deserialize<'de, D: Deserializer<'de>, const N: usize>(
127 deserializer: D,
128 ) -> Result<[f64; N], D::Error> {
129 let values = Vec::<f64>::deserialize(deserializer)?;
130 let got = values.len();
131 values
132 .try_into()
133 .map_err(|_| D::Error::custom(format!("expected {N} values, got {got}")))
134 }
135}
136
137#[derive(Debug, Clone, Copy)]
139struct Piece<const N: usize> {
140 a: f64,
141 b: f64,
142 value: [f64; N],
143 error: [f64; N],
144}
145
146fn rule<const N: usize, F>(f: &mut F, a: f64, b: f64) -> Result<Piece<N>, CoreError>
148where
149 F: FnMut(f64) -> [f64; N],
150{
151 let center = 0.5 * (a + b);
152 let half = 0.5 * (b - a);
153 let mut kronrod = [0.0; N];
154 let mut gauss = [0.0; N];
155 let mut add = |x: f64, kw: f64, gw: f64| -> Result<(), CoreError> {
156 let y = f(x);
157 for k in 0..N {
158 if !y[k].is_finite() {
159 return Err(CoreError::QuadratureNotFinite { x });
160 }
161 kronrod[k] += kw * y[k];
162 gauss[k] += gw * y[k];
163 }
164 Ok(())
165 };
166 add(center, KRONROD_WEIGHTS[7], GAUSS_WEIGHTS[3])?;
167 for (i, (&node, &kw)) in NODES.iter().zip(&KRONROD_WEIGHTS).take(7).enumerate() {
168 let gw = if i % 2 == 1 {
169 GAUSS_WEIGHTS[i / 2]
170 } else {
171 0.0
172 };
173 add(center - half * node, kw, gw)?;
174 add(center + half * node, kw, gw)?;
175 }
176 let mut value = [0.0; N];
177 let mut error = [0.0; N];
178 for k in 0..N {
179 value[k] = half * kronrod[k];
180 error[k] = (half * (kronrod[k] - gauss[k])).abs();
181 }
182 Ok(Piece { a, b, value, error })
183}
184
185pub fn integrate<const N: usize, F>(
200 mut f: F,
201 a: f64,
202 b: f64,
203 tolerance: Tolerance,
204) -> Result<Integral<N>, CoreError>
205where
206 F: FnMut(f64) -> [f64; N],
207{
208 for (what, value) in [
209 ("integration lower limit", a),
210 ("integration upper limit", b),
211 ] {
212 if !value.is_finite() {
213 return Err(CoreError::Domain { what, value });
214 }
215 }
216 for (what, value) in [
217 ("relative tolerance", tolerance.relative),
218 ("absolute tolerance", tolerance.absolute),
219 ] {
220 if value.is_nan() || value < 0.0 {
221 return Err(CoreError::Domain { what, value });
222 }
223 }
224 if tolerance.max_intervals == 0 {
225 return Err(CoreError::Domain {
226 what: "maximum number of subintervals",
227 value: 0.0,
228 });
229 }
230 if a == b {
231 return Ok(Integral {
232 value: [0.0; N],
233 error: [0.0; N],
234 intervals: 1,
235 });
236 }
237 let relative = tolerance.relative.max(50.0 * f64::EPSILON);
239 let mut pieces = vec![rule(&mut f, a, b)?];
240 loop {
241 let mut total = [0.0; N];
242 let mut error = [0.0; N];
243 for piece in &pieces {
244 for k in 0..N {
245 total[k] += piece.value[k];
246 error[k] += piece.error[k];
247 }
248 }
249 let bound: [f64; N] =
250 std::array::from_fn(|k| tolerance.absolute.max(relative * total[k].abs()));
251 if (0..N).all(|k| error[k] <= bound[k]) {
252 return Ok(Integral {
253 value: total,
254 error,
255 intervals: pieces.len(),
256 });
257 }
258 let score = |piece: &Piece<N>| {
260 (0..N)
261 .map(|k| piece.error[k] / bound[k].max(f64::MIN_POSITIVE))
262 .fold(0.0, f64::max)
263 };
264 let worst = (0..pieces.len())
265 .max_by(|&i, &j| score(&pieces[i]).total_cmp(&score(&pieces[j])))
266 .unwrap_or(0);
267 let piece = pieces[worst];
268 let middle = 0.5 * (piece.a + piece.b);
269 let tiny = middle == piece.a || middle == piece.b;
270 if pieces.len() >= tolerance.max_intervals || tiny {
271 let worst_component = (0..N)
272 .max_by(|&i, &j| {
273 (error[i] / bound[i].max(f64::MIN_POSITIVE))
274 .total_cmp(&(error[j] / bound[j].max(f64::MIN_POSITIVE)))
275 })
276 .unwrap_or(0);
277 return Err(CoreError::QuadratureDidNotConverge {
278 component: worst_component,
279 value: total.get(worst_component).copied().unwrap_or(0.0),
280 error: error.get(worst_component).copied().unwrap_or(0.0),
281 intervals: pieces.len(),
282 });
283 }
284 pieces[worst] = rule(&mut f, piece.a, middle)?;
285 pieces.push(rule(&mut f, middle, piece.b)?);
286 }
287}
288
289pub fn integrate_scalar<F>(mut f: F, a: f64, b: f64, tolerance: Tolerance) -> Result<f64, CoreError>
295where
296 F: FnMut(f64) -> f64,
297{
298 integrate(|x| [f(x)], a, b, tolerance).map(|integral| integral.value[0])
299}
300
301#[cfg(test)]
302mod tests {
303 use super::*;
304
305 fn legendre(n: usize, x: f64) -> f64 {
307 let (mut p0, mut p1) = (1.0, x);
308 for k in 1..n {
309 let k = k as f64;
310 let p2 = ((2.0 * k + 1.0) * x * p1 - k * p0) / (k + 1.0);
311 p0 = p1;
312 p1 = p2;
313 }
314 if n == 0 { p0 } else { p1 }
315 }
316
317 #[test]
318 fn the_nodes_and_weights_are_the_published_rule() {
319 for j in [1, 3, 5, 7] {
321 assert!(legendre(7, NODES[j]).abs() < 1e-15, "node {j}");
322 }
323 let kronrod_sum: f64 = 2.0 * KRONROD_WEIGHTS[..7].iter().sum::<f64>() + KRONROD_WEIGHTS[7];
324 let gauss_sum: f64 = 2.0 * GAUSS_WEIGHTS[..3].iter().sum::<f64>() + GAUSS_WEIGHTS[3];
325 assert!((kronrod_sum - 2.0).abs() < 1e-15);
326 assert!((gauss_sum - 2.0).abs() < 1e-15);
327 }
328
329 #[test]
330 fn one_panel_is_exact_to_the_rules_degrees() {
331 for degree in 0..=24 {
333 let exact = if degree % 2 == 0 {
334 2.0 / (degree as f64 + 1.0)
335 } else {
336 0.0
337 };
338 let piece = rule(&mut |x: f64| [x.powi(degree)], -1.0, 1.0).unwrap();
339 let kronrod_error = (piece.value[0] - exact).abs();
340 if degree <= 22 {
341 assert!(
342 kronrod_error < 2e-16,
343 "K15 degree {degree}: {kronrod_error}"
344 );
345 } else if degree % 2 == 0 {
346 assert!(kronrod_error > 1e-12, "K15 is not exact at degree {degree}");
348 }
349 if degree <= 13 {
351 assert!(
352 piece.error[0] < 2e-16,
353 "G7 degree {degree}: {}",
354 piece.error[0]
355 );
356 } else if degree % 2 == 0 {
357 assert!(piece.error[0] > 1e-8, "G7 is not exact at degree {degree}");
358 }
359 }
360 let piece = rule(&mut |x: f64| [x.powi(22)], 1.0, 3.0).unwrap();
362 let exact = (3f64.powi(23) - 1.0) / 23.0;
363 assert!((piece.value[0] / exact - 1.0).abs() < 1e-14);
364 }
365
366 #[test]
367 fn smooth_singular_and_kinked_integrands_converge() {
368 type Case<'a> = (&'a dyn Fn(f64) -> f64, f64, f64, f64);
370 let tol = Tolerance::default();
371 let cases: [Case; 6] = [
372 (&|x: f64| x.exp(), 0.0, 1.0, std::f64::consts::E - 1.0),
373 (&|x: f64| x.sin(), 0.0, std::f64::consts::PI, 2.0),
374 (&|x: f64| x.sqrt(), 0.0, 1.0, 2.0 / 3.0),
376 (&|x: f64| 1.0 / x.sqrt(), 0.0, 1.0, 2.0),
377 (&|x: f64| x.powf(0.3), 0.0, 2.0, 2f64.powf(1.3) / 1.3),
378 (&|x: f64| (x - 0.3).abs(), 0.0, 1.0, 0.5 * (0.09 + 0.49)),
380 ];
381 for (i, (f, a, b, exact)) in cases.iter().enumerate() {
382 let value = integrate_scalar(f, *a, *b, tol).unwrap();
383 let relative = ((value - exact) / exact).abs();
384 assert!(
385 relative < 1e-11,
386 "case {i}: {value} vs {exact} ({relative:e})"
387 );
388 }
389 }
390
391 #[test]
392 fn components_converge_together_and_limits_can_be_reversed() {
393 let integral = integrate(
394 |x: f64| [1.0, x, x * x, (x * 10.0).sin()],
395 0.0,
396 2.0,
397 Tolerance::default(),
398 )
399 .unwrap();
400 let exact = [2.0, 2.0, 8.0 / 3.0, (1.0 - 20f64.cos()) / 10.0];
401 for (k, want) in exact.iter().enumerate() {
402 assert!((integral.value[k] - want).abs() < 1e-12, "component {k}");
403 assert!(integral.error[k] <= 1e-12f64.max(1e-12 * want.abs()));
404 }
405 let backwards = integrate_scalar(|x| x * x, 2.0, 0.0, Tolerance::default()).unwrap();
406 assert!((backwards + 8.0 / 3.0).abs() < 1e-14);
407 assert_eq!(
408 integrate(|_| [1.0, 2.0], 1.5, 1.5, Tolerance::default())
409 .unwrap()
410 .value,
411 [0.0, 0.0]
412 );
413 }
414
415 #[test]
416 fn bad_inputs_and_hard_integrands_are_errors() {
417 let tol = Tolerance::default();
418 assert!(matches!(
419 integrate_scalar(|x| x, f64::NAN, 1.0, tol),
420 Err(CoreError::Domain { .. })
421 ));
422 assert!(matches!(
423 integrate_scalar(
424 |x| x,
425 0.0,
426 1.0,
427 Tolerance {
428 relative: -1.0,
429 ..tol
430 }
431 ),
432 Err(CoreError::Domain { .. })
433 ));
434 assert!(matches!(
435 integrate_scalar(|x| if x > 0.5 { f64::NAN } else { x }, 0.0, 1.0, tol),
436 Err(CoreError::QuadratureNotFinite { .. })
437 ));
438 let limited = Tolerance {
440 max_intervals: 50,
441 ..tol
442 };
443 assert!(matches!(
444 integrate_scalar(|x| 1.0 / x, 0.0, 1.0, limited),
445 Err(CoreError::QuadratureDidNotConverge { intervals: 50, .. })
446 ));
447 }
448
449 #[test]
450 fn a_result_serializes_as_plain_arrays() {
451 let integral = integrate(|x| [x, 1.0], 0.0, 1.0, Tolerance::default()).unwrap();
452 let json = serde_json::to_string(&integral).unwrap();
453 let back: Integral<2> = serde_json::from_str(&json).unwrap();
454 assert_eq!(back, integral);
455 assert!(serde_json::from_str::<Integral<3>>(&json).is_err());
456 }
457}