Skip to main content

sde_sim_rs/
filtration.rs

1use crate::proc::ProcessUniverse;
2use ordered_float::OrderedFloat;
3use polars::prelude::*;
4use std::collections::BTreeMap;
5use std::collections::HashMap;
6
7pub struct ScenarioFiltrationCache {
8    pub time: OrderedFloat<f64>,
9    pub values: BTreeMap<String, f64>,
10}
11
12pub struct ScenarioFiltration {
13    pub scenario: i32,
14    pub times: Vec<OrderedFloat<f64>>,
15    pub process_universe: ProcessUniverse,
16    raw_values: Vec<f64>,
17    time_registry: HashMap<OrderedFloat<f64>, usize>,
18    pub cache: ScenarioFiltrationCache,
19}
20
21impl ScenarioFiltration {
22    pub fn new(
23        scenario: i32,
24        process_universe: ProcessUniverse,
25        times: Vec<OrderedFloat<f64>>,
26        initial_values: HashMap<String, f64>,
27    ) -> Self {
28        let raw_values = vec![0.0; times.len() * process_universe.processes.len()];
29        let time_registry = times.iter().enumerate().map(|(i, t)| (*t, i)).collect();
30        let value_cache = ScenarioFiltrationCache {
31            time: times[0],
32            values: BTreeMap::new(),
33        };
34        let mut scenario_filtration = ScenarioFiltration {
35            scenario,
36            process_universe,
37            times,
38            raw_values,
39            time_registry,
40            cache: value_cache,
41        };
42        for (process_name, val) in initial_values.into_iter() {
43            if let Some(process_idx) = scenario_filtration
44                .process_universe
45                .process_registry
46                .get(&process_name)
47            {
48                scenario_filtration.set(0, *process_idx, val);
49            }
50        }
51        scenario_filtration.refresh_cache(scenario_filtration.times[0]);
52        scenario_filtration
53    }
54
55    #[inline]
56    pub fn get(&self, time_idx: usize, process_idx: usize) -> f64 {
57        self.raw_values[time_idx * self.process_universe.processes.len() + process_idx]
58    }
59
60    #[inline]
61    pub fn set(&mut self, time_idx: usize, process_idx: usize, val: f64) {
62        let idx = time_idx * self.process_universe.processes.len() + process_idx;
63        self.raw_values[idx] = val;
64    }
65
66    pub fn get_time_idx(&self, time: OrderedFloat<f64>) -> Option<&usize> {
67        self.time_registry.get(&time)
68    }
69
70    pub fn refresh_cache(&mut self, time: OrderedFloat<f64>) {
71        self.cache.time = time;
72        self.cache.values.insert("t".to_string(), time.into_inner());
73        let t_idx = self.get_time_idx(time).copied().unwrap_or(0);
74        for (p_name, p_idx) in self.process_universe.process_registry.iter() {
75            self.cache
76                .values
77                .insert(p_name.clone(), self.get(t_idx, *p_idx));
78        }
79    }
80
81    pub fn to_lazyframe(&self) -> LazyFrame {
82        let num_procs = self.process_universe.processes.len();
83        let num_times = self.times.len();
84
85        // 1. Fixed PlSmallStr by adding .into()
86        // and using StringChunked::from_iter for cleaner collection
87        let process_names: Series = StringChunked::from_iter(
88            self.times
89                .iter()
90                .flat_map(|_| self.process_universe.processes.iter().map(|p| p.name())),
91        )
92        .with_name("process_name".into())
93        .into_series();
94
95        // 2. Fixed Float64Chunked collection
96        // We use Float64Chunked::from_iter and .into() for the name
97        let times: Series = Float64Chunked::from_iter(
98            self.times
99                .iter()
100                .flat_map(|t| std::iter::repeat_n(Some(t.0), num_procs)),
101        )
102        .with_name("time".into())
103        .into_series();
104
105        // 3. Build the DataFrame
106        // Note: The df! macro in 0.51 also expects PlSmallStr for column names
107        // but the macro usually handles string literals via internal conversion.
108        df![
109            "scenario" => [self.scenario].repeat(num_procs * num_times),
110            "time" => times,
111            "process_name" => process_names,
112            "value" => &self.raw_values
113        ]
114        .expect("Failed to create DataFrame")
115        .lazy()
116    }
117}