1#[cfg(feature = "polars")]
4use polars::prelude::*;
5
6#[cfg(feature = "serde")]
7use serde::{Deserialize, Serialize};
8
9use crate::{
10 stats::{Evals, Steps, Timer},
11 status::Status,
12 traits::{Real, State},
13};
14
15#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
26#[derive(Debug, Clone)]
27pub struct Solution<T, Y>
28where
29 T: Real,
30 Y: State<T>,
31{
32 pub t: Vec<T>,
34
35 pub y: Vec<Y>,
37
38 pub status: Status<T, Y>,
40
41 pub evals: Evals,
43
44 pub steps: Steps,
46
47 #[cfg(not(target_arch = "wasm32"))]
49 pub timer: Timer<T>,
50}
51
52impl<T, Y> Default for Solution<T, Y>
54where
55 T: Real,
56 Y: State<T>,
57{
58 fn default() -> Self {
59 Self::new()
60 }
61}
62
63impl<T, Y> Solution<T, Y>
64where
65 T: Real,
66 Y: State<T>,
67{
68 pub fn new() -> Self {
70 Solution {
71 t: Vec::new(),
72 y: Vec::new(),
73 status: Status::Uninitialized,
74 evals: Evals::new(),
75 steps: Steps::new(),
76 #[cfg(not(target_arch = "wasm32"))]
77 timer: Timer::Off,
78 }
79 }
80
81 pub fn new_with_capacity(capacity: usize) -> Self {
86 Solution {
87 t: Vec::with_capacity(capacity),
88 y: Vec::with_capacity(capacity),
89 status: Status::Uninitialized,
90 evals: Evals::new(),
91 steps: Steps::new(),
92 #[cfg(not(target_arch = "wasm32"))]
93 timer: Timer::Off,
94 }
95 }
96}
97
98impl<T, Y> Solution<T, Y>
100where
101 T: Real,
102 Y: State<T>,
103{
104 pub fn push(&mut self, t: T, y: Y) {
111 self.t.push(t);
112 self.y.push(y);
113 }
114
115 pub fn pop(&mut self) -> Option<(T, Y)> {
121 if self.t.is_empty() || self.y.is_empty() {
122 return None;
123 }
124 let t = self.t.pop().unwrap();
125 let y = self.y.pop().unwrap();
126 Some((t, y))
127 }
128
129 pub fn truncate(&mut self, index: usize) {
135 self.t.truncate(index);
136 self.y.truncate(index);
137 }
138}
139
140impl<T, Y> Solution<T, Y>
142where
143 T: Real,
144 Y: State<T>,
145{
146 pub fn into_tuple(self) -> (Vec<T>, Vec<Y>) {
154 (self.t, self.y)
155 }
156
157 pub fn last(&self) -> Result<(&T, &Y), Box<dyn std::error::Error>> {
163 let t = self.t.last().ok_or("No t steps available")?;
164 let y = self.y.last().ok_or("No y vectors available")?;
165 Ok((t, y))
166 }
167
168 pub fn iter(&self) -> std::iter::Zip<std::slice::Iter<'_, T>, std::slice::Iter<'_, Y>> {
175 self.t.iter().zip(self.y.iter())
176 }
177
178 #[cfg(not(feature = "polars"))]
189 pub fn to_csv(&self, filename: &str) -> Result<(), Box<dyn std::error::Error>> {
190 use std::io::{BufWriter, Write};
191
192 let path = std::path::Path::new(filename);
194 if let Some(parent) = path.parent()
195 && !parent.exists()
196 {
197 std::fs::create_dir_all(parent)?;
198 }
199 let file = std::fs::File::create(filename)?;
200 let mut writer = BufWriter::new(file);
201
202 let n = self.y[0].len();
204
205 let mut header = String::from("t");
207 for i in 0..n {
208 header.push_str(&format!(",y{}", i));
209 }
210 writeln!(writer, "{}", header)?;
211
212 for (t, y) in self.iter() {
214 let mut row = format!("{:?}", t);
215 for i in 0..n {
216 row.push_str(&format!(",{:?}", y.get_component(i)));
217 }
218 writeln!(writer, "{}", row)?;
219 }
220
221 writer.flush()?;
222
223 Ok(())
224 }
225
226 #[cfg(feature = "polars")]
237 pub fn to_csv(&self, filename: &str) -> Result<(), Box<dyn std::error::Error>> {
238 let path = std::path::Path::new(filename);
240 if let Some(parent) = path.parent()
241 && !parent.exists()
242 {
243 std::fs::create_dir_all(parent)?;
244 }
245 let mut file = std::fs::File::create(filename)?;
246
247 let t = self
248 .t
249 .iter()
250 .map(simba::scalar::SupersetOf::<f64>::to_subset_unchecked)
251 .collect::<Vec<f64>>();
252 let mut columns = vec![Column::new("t".into(), t)];
253 let n = self.y[0].len();
254 for i in 0..n {
255 let header = format!("y{}", i);
256 columns.push(Column::new(
257 header.into(),
258 self.y
259 .iter()
260 .map(|y| {
261 simba::scalar::SupersetOf::<f64>::to_subset_unchecked(&y.get_component(i))
262 })
263 .collect::<Vec<f64>>(),
264 ));
265 }
266 let mut df = DataFrame::new(self.t.len(), columns)?;
267
268 CsvWriter::new(&mut file).finish(&mut df)?;
270
271 Ok(())
272 }
273
274 #[cfg(feature = "polars")]
284 pub fn to_polars(&self) -> Result<DataFrame, PolarsError> {
285 let t = self
286 .t
287 .iter()
288 .map(simba::scalar::SupersetOf::<f64>::to_subset_unchecked)
289 .collect::<Vec<f64>>();
290 let mut columns = vec![Column::new("t".into(), t)];
291 let n = self.y[0].len();
292 for i in 0..n {
293 let header = format!("y{}", i);
294 columns.push(Column::new(
295 header.into(),
296 self.y
297 .iter()
298 .map(|y| {
299 simba::scalar::SupersetOf::<f64>::to_subset_unchecked(&y.get_component(i))
300 })
301 .collect::<Vec<f64>>(),
302 ));
303 }
304
305 DataFrame::new(self.t.len(), columns)
306 }
307
308 #[cfg(feature = "polars")]
320 pub fn to_named_polars(
321 &self,
322 t_name: &str,
323 y_names: Vec<&str>,
324 ) -> Result<DataFrame, PolarsError> {
325 let t = self
326 .t
327 .iter()
328 .map(simba::scalar::SupersetOf::<f64>::to_subset_unchecked)
329 .collect::<Vec<f64>>();
330 let mut columns = vec![Column::new(t_name.into(), t)];
331
332 let n = self.y[0].len();
333
334 if y_names.len() != n {
336 return Err(PolarsError::ComputeError(
337 format!(
338 "Expected {} column names for state variables, but got {}",
339 n,
340 y_names.len()
341 )
342 .into(),
343 ));
344 }
345
346 for (i, name) in y_names.iter().enumerate() {
347 columns.push(Column::new(
348 (*name).into(),
349 self.y
350 .iter()
351 .map(|y| {
352 simba::scalar::SupersetOf::<f64>::to_subset_unchecked(&y.get_component(i))
353 })
354 .collect::<Vec<f64>>(),
355 ));
356 }
357
358 DataFrame::new(self.t.len(), columns)
359 }
360}
361
362#[cfg(test)]
363mod tests {
364 use super::*;
365
366 #[test]
367 fn test_into_tuple() {
368 let mut sol: Solution<f64, f64> = Solution::new();
369 sol.push(0.0, 10.0);
370 sol.push(1.0, 20.0);
371
372 let (t, y) = sol.into_tuple();
373 assert_eq!(t, vec![0.0, 1.0]);
374 assert_eq!(y, vec![10.0, 20.0]);
375 }
376
377 #[test]
378 fn test_solution_lifecycle() {
379 let sol_new: Solution<f64, f64> = Solution::new();
381 assert!(sol_new.t.is_empty());
382 assert!(sol_new.y.is_empty());
383
384 let sol_cap: Solution<f64, f64> = Solution::new_with_capacity(10);
385 assert!(sol_cap.t.is_empty());
386 assert!(sol_cap.y.is_empty());
387 assert!(sol_cap.t.capacity() >= 10);
388 assert!(sol_cap.y.capacity() >= 10);
389
390 let mut sol = sol_new;
392 sol.push(2.0, 30.0);
393 assert_eq!(sol.t.len(), 1);
394 assert_eq!(sol.y.len(), 1);
395 assert_eq!(sol.t[0], 2.0);
396 assert_eq!(sol.y[0], 30.0);
397
398 let last = sol.last().unwrap();
400 assert_eq!(*last.0, 2.0);
401 assert_eq!(*last.1, 30.0);
402
403 let popped = sol.pop();
405 assert_eq!(popped, Some((2.0, 30.0)));
406 assert!(sol.t.is_empty());
407 assert!(sol.y.is_empty());
408
409 assert!(sol.last().is_err());
411
412 assert_eq!(sol.pop(), None);
414
415 sol.push(0.0, 10.0);
417 sol.push(1.0, 20.0);
418 sol.push(2.0, 30.0);
419
420 let expected = vec![(0.0, 10.0), (1.0, 20.0), (2.0, 30.0)];
421 let actual: Vec<(f64, f64)> = sol.iter().map(|(&t, &y)| (t, y)).collect();
422 assert_eq!(actual, expected);
423
424 sol.truncate(1);
425 assert_eq!(sol.t.len(), 1);
426 assert_eq!(sol.y.len(), 1);
427 assert_eq!(sol.t[0], 0.0);
428 assert_eq!(sol.y[0], 10.0);
429 }
430}