1#[derive(Debug, Clone, Copy, PartialEq)]
4pub enum OdeError {
5 InvalidDomain,
6 DimensionMismatch,
7 WorkspaceTooSmall { required: usize, available: usize },
8 NonFiniteDerivative,
9 StepUnderflow,
10 StepLimitExceeded,
11}
12
13#[repr(C)]
14#[derive(Debug, Clone, Copy, PartialEq)]
15pub struct AdaptiveOdeConfig {
16 pub absolute_tolerance: f64,
17 pub relative_tolerance: f64,
18 pub initial_step: f64,
19 pub minimum_step: f64,
20 pub maximum_step: f64,
21 pub max_steps: u32,
22}
23
24impl Default for AdaptiveOdeConfig {
25 fn default() -> Self {
26 Self {
27 absolute_tolerance: 1e-9,
28 relative_tolerance: 1e-7,
29 initial_step: 1e-3,
30 minimum_step: 1e-14,
31 maximum_step: 1.0,
32 max_steps: 100_000,
33 }
34 }
35}
36
37#[repr(C)]
38#[derive(Debug, Clone, Copy, PartialEq)]
39pub struct AdaptiveOdeResult {
40 pub final_time: f64,
41 pub accepted_steps: u32,
42 pub rejected_steps: u32,
43 pub derivative_evaluations: u32,
44 pub last_error_norm: f64,
45 pub last_step: f64,
46}
47
48pub const fn dopri5_workspace_len(dimension: usize) -> Option<usize> {
49 dimension.checked_mul(8)
50}
51
52pub fn integrate_dopri5<F>(
53 derivative: F,
54 state: &mut [f64],
55 t0: f64,
56 t_final: f64,
57 config: AdaptiveOdeConfig,
58 workspace: &mut [f64],
59) -> Result<AdaptiveOdeResult, OdeError>
60where
61 F: Fn(f64, &[f64], &mut [f64]) -> Result<(), OdeError>,
62{
63 let dimension = state.len();
64 let required = dopri5_workspace_len(dimension).ok_or(OdeError::InvalidDomain)?;
65 if dimension == 0
66 || !t0.is_finite()
67 || !t_final.is_finite()
68 || t_final < t0
69 || state.iter().any(|value| !value.is_finite())
70 || !valid_config(config)
71 {
72 return Err(OdeError::InvalidDomain);
73 }
74 if workspace.len() < required {
75 return Err(OdeError::WorkspaceTooSmall {
76 required,
77 available: workspace.len(),
78 });
79 }
80 if t_final == t0 {
81 return Ok(AdaptiveOdeResult {
82 final_time: t0,
83 accepted_steps: 0,
84 rejected_steps: 0,
85 derivative_evaluations: 0,
86 last_error_norm: 0.0,
87 last_step: 0.0,
88 });
89 }
90
91 let (k1, rest) = workspace.split_at_mut(dimension);
92 let (k2, rest) = rest.split_at_mut(dimension);
93 let (k3, rest) = rest.split_at_mut(dimension);
94 let (k4, rest) = rest.split_at_mut(dimension);
95 let (k5, rest) = rest.split_at_mut(dimension);
96 let (k6, rest) = rest.split_at_mut(dimension);
97 let (k7, rest) = rest.split_at_mut(dimension);
98 let temp = &mut rest[..dimension];
99
100 let mut time = t0;
101 let mut step = config.initial_step.min(config.maximum_step);
102 let mut accepted = 0;
103 let mut rejected = 0;
104 let mut evaluations = 0;
105 let mut last_error = f64::INFINITY;
106 let mut last_step = 0.0;
107
108 for _ in 0..config.max_steps {
109 if time >= t_final {
110 return Ok(AdaptiveOdeResult {
111 final_time: time,
112 accepted_steps: accepted,
113 rejected_steps: rejected,
114 derivative_evaluations: evaluations,
115 last_error_norm: last_error,
116 last_step,
117 });
118 }
119 step = step.min(t_final - time).min(config.maximum_step);
120 if step < config.minimum_step || time + step == time {
121 return Err(OdeError::StepUnderflow);
122 }
123
124 eval(&derivative, time, state, k1)?;
125 stage(state, temp, step, &[(1.0 / 5.0, k1)]);
126 eval(&derivative, time + step / 5.0, temp, k2)?;
127
128 stage(state, temp, step, &[(3.0 / 40.0, k1), (9.0 / 40.0, k2)]);
129 eval(&derivative, time + 3.0 * step / 10.0, temp, k3)?;
130
131 stage(
132 state,
133 temp,
134 step,
135 &[(44.0 / 45.0, k1), (-56.0 / 15.0, k2), (32.0 / 9.0, k3)],
136 );
137 eval(&derivative, time + 4.0 * step / 5.0, temp, k4)?;
138
139 stage(
140 state,
141 temp,
142 step,
143 &[
144 (19372.0 / 6561.0, k1),
145 (-25360.0 / 2187.0, k2),
146 (64448.0 / 6561.0, k3),
147 (-212.0 / 729.0, k4),
148 ],
149 );
150 eval(&derivative, time + 8.0 * step / 9.0, temp, k5)?;
151
152 stage(
153 state,
154 temp,
155 step,
156 &[
157 (9017.0 / 3168.0, k1),
158 (-355.0 / 33.0, k2),
159 (46732.0 / 5247.0, k3),
160 (49.0 / 176.0, k4),
161 (-5103.0 / 18656.0, k5),
162 ],
163 );
164 eval(&derivative, time + step, temp, k6)?;
165
166 stage(
167 state,
168 temp,
169 step,
170 &[
171 (35.0 / 384.0, k1),
172 (500.0 / 1113.0, k3),
173 (125.0 / 192.0, k4),
174 (-2187.0 / 6784.0, k5),
175 (11.0 / 84.0, k6),
176 ],
177 );
178 eval(&derivative, time + step, temp, k7)?;
179 evaluations += 7;
180
181 let mut error_norm = 0.0_f64;
182 for index in 0..dimension {
183 let fifth = state[index]
184 + step
185 * (35.0 / 384.0 * k1[index]
186 + 500.0 / 1113.0 * k3[index]
187 + 125.0 / 192.0 * k4[index]
188 - 2187.0 / 6784.0 * k5[index]
189 + 11.0 / 84.0 * k6[index]);
190 let fourth = state[index]
191 + step
192 * (5179.0 / 57600.0 * k1[index]
193 + 7571.0 / 16695.0 * k3[index]
194 + 393.0 / 640.0 * k4[index]
195 - 92097.0 / 339200.0 * k5[index]
196 + 187.0 / 2100.0 * k6[index]
197 + 1.0 / 40.0 * k7[index]);
198 let scale = config.absolute_tolerance
199 + config.relative_tolerance * state[index].abs().max(fifth.abs());
200 error_norm = error_norm.max((fifth - fourth).abs() / scale);
201 temp[index] = fifth;
202 }
203 if !error_norm.is_finite() || temp.iter().any(|value| !value.is_finite()) {
204 return Err(OdeError::NonFiniteDerivative);
205 }
206
207 last_error = error_norm;
208 let factor = if error_norm == 0.0 {
209 5.0
210 } else {
211 (0.9 * error_norm.powf(-0.2)).clamp(0.2, 5.0)
212 };
213 if error_norm <= 1.0 {
214 state.copy_from_slice(temp);
215 time += step;
216 last_step = step;
217 accepted += 1;
218 } else {
219 rejected += 1;
220 }
221 step = (step * factor).clamp(config.minimum_step, config.maximum_step);
222 }
223
224 Err(OdeError::StepLimitExceeded)
225}
226
227fn valid_config(config: AdaptiveOdeConfig) -> bool {
228 config.absolute_tolerance.is_finite()
229 && config.absolute_tolerance > 0.0
230 && config.relative_tolerance.is_finite()
231 && config.relative_tolerance >= 0.0
232 && config.initial_step.is_finite()
233 && config.initial_step > 0.0
234 && config.minimum_step.is_finite()
235 && config.minimum_step > 0.0
236 && config.maximum_step.is_finite()
237 && config.maximum_step >= config.minimum_step
238 && config.max_steps > 0
239}
240
241fn eval<F>(derivative: &F, time: f64, state: &[f64], output: &mut [f64]) -> Result<(), OdeError>
242where
243 F: Fn(f64, &[f64], &mut [f64]) -> Result<(), OdeError>,
244{
245 derivative(time, state, output)?;
246 if output.len() != state.len() {
247 return Err(OdeError::DimensionMismatch);
248 }
249 if output.iter().any(|value| !value.is_finite()) {
250 return Err(OdeError::NonFiniteDerivative);
251 }
252 Ok(())
253}
254
255fn stage(state: &[f64], output: &mut [f64], step: f64, terms: &[(f64, &[f64])]) {
256 for index in 0..state.len() {
257 let mut value = state[index];
258 for (coefficient, derivative) in terms {
259 value += step * coefficient * derivative[index];
260 }
261 output[index] = value;
262 }
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268
269 fn decay(_time: f64, state: &[f64], output: &mut [f64]) -> Result<(), OdeError> {
270 if state.len() != output.len() {
271 return Err(OdeError::DimensionMismatch);
272 }
273 for (out, value) in output.iter_mut().zip(state) {
274 *out = -*value;
275 }
276 Ok(())
277 }
278
279 #[test]
280 fn dopri5_lands_exactly_at_final_time_with_tolerance_control() {
281 let mut state = [1.0, 2.0, 3.0];
282 let mut workspace = [0.0; 24];
283 let report = integrate_dopri5(
284 decay,
285 &mut state,
286 0.0,
287 2.0,
288 AdaptiveOdeConfig::default(),
289 &mut workspace,
290 )
291 .unwrap();
292 assert_eq!(report.final_time, 2.0);
293 for (index, value) in state.iter().enumerate() {
294 let expected = (index + 1) as f64 * (-2.0_f64).exp();
295 assert!((value - expected).abs() < 3e-8);
296 }
297 assert!(report.accepted_steps > 0);
298 }
299
300 #[test]
301 fn dopri5_rejects_workspace_and_non_finite_derivatives() {
302 let mut state = [1.0, 2.0];
303 let mut too_small = [0.0; 15];
304 assert_eq!(
305 integrate_dopri5(
306 decay,
307 &mut state,
308 0.0,
309 1.0,
310 AdaptiveOdeConfig::default(),
311 &mut too_small,
312 ),
313 Err(OdeError::WorkspaceTooSmall {
314 required: 16,
315 available: 15
316 })
317 );
318
319 let mut workspace = [0.0; 16];
320 assert_eq!(
321 integrate_dopri5(
322 |_t, _y, out| {
323 out.fill(f64::NAN);
324 Ok(())
325 },
326 &mut state,
327 0.0,
328 1.0,
329 AdaptiveOdeConfig::default(),
330 &mut workspace,
331 ),
332 Err(OdeError::NonFiniteDerivative)
333 );
334 }
335}