1use crate::domain::{EvaluationDomain, PowfExtension};
10use crate::error::Result;
11use crate::evaluator::ExpressionEvaluator;
12
13pub struct StreamingEvaluator<'a, T: EvaluationDomain> {
31 evaluator: &'a ExpressionEvaluator<T>,
32 params: Vec<T>,
33 stack: Vec<T>,
34 results: Vec<T>,
35}
36
37impl<T: EvaluationDomain + PowfExtension> StreamingEvaluator<'_, T> {
38 pub fn new(evaluator: &ExpressionEvaluator<T>) -> StreamingEvaluator<'_, T> {
40 StreamingEvaluator {
41 evaluator,
42 params: Vec::with_capacity(evaluator.param_count()),
43 stack: Vec::with_capacity(evaluator.stack_size()),
44 results: Vec::with_capacity(evaluator.result_count()),
45 }
46 }
47
48 pub fn for_each<I, S, F>(&mut self, rows: I, mut sink: F) -> Result<usize>
66 where
67 I: IntoIterator<Item = S>,
68 S: AsRef<[T]>,
69 F: FnMut(&[T]),
70 {
71 let mut count = 0usize;
72 for row in rows {
73 let row = row.as_ref();
74 self.params.clear();
75 self.params.extend(row.iter().cloned());
76 self.evaluator
77 .evaluate_with_stack(&self.params, &mut self.stack, &mut self.results)?;
78 sink(&self.results);
79 count += 1;
80 }
81 Ok(count)
82 }
83
84 pub fn evaluate_chunk<S: AsRef<[T]>>(&mut self, rows: &[S]) -> Result<Vec<Vec<T>>> {
95 let mut out = Vec::with_capacity(rows.len());
96 self.for_each(rows, |results| out.push(results.to_vec()))?;
97 Ok(out)
98 }
99}
100
101#[cfg(test)]
102mod tests {
103 use super::*;
104 use ocas_atom::AtomArena;
105 use ocas_core::arena::Arena;
106
107 fn build_eval() -> (Arena, ExpressionEvaluator<f64>) {
108 let arena = Arena::new();
109 let ctx = AtomArena::new(&arena);
110 let sum = ctx.add(&[ctx.var("x"), ctx.var("y")]);
111 let prod = ctx.mul(&[ctx.var("x"), ctx.var("y")]);
112 let eval = ExpressionEvaluator::compile_multi(&[sum, prod]).unwrap();
113 (arena, eval)
114 }
115
116 #[test]
117 fn streaming_multi_output() {
118 let (_arena, eval) = build_eval();
119 let mut stream = StreamingEvaluator::new(&eval);
120 let rows: Vec<[f64; 2]> = (0..100).map(|i| [i as f64, 2.0]).collect();
121 let mut seen = Vec::new();
122 let n = stream
123 .for_each(&rows, |results| seen.push((results[0], results[1])))
124 .unwrap();
125 assert_eq!(n, 100);
126 for (i, &(sum, prod)) in seen.iter().enumerate() {
127 assert!((sum - (i as f64 + 2.0)).abs() < 1e-10);
128 assert!((prod - (i as f64 * 2.0)).abs() < 1e-10);
129 }
130 }
131
132 #[test]
133 fn streaming_constant_memory_million_rows() {
134 let (_arena, eval) = build_eval();
135 let mut stream = StreamingEvaluator::new(&eval);
136
137 let warm: Vec<[f64; 2]> = vec![[1.0, 2.0]; 10];
139 stream.for_each(&warm, |_| {}).unwrap();
140 let stack_cap = stream.stack.capacity();
141 let results_cap = stream.results.capacity();
142 let params_cap = stream.params.capacity();
143
144 let rows = (0..1_000_000u64).map(|i| [i as f64 % 100.0, 3.0]);
146 let mut count = 0usize;
147 let mut checksum = 0.0f64;
148 let n = stream
149 .for_each(rows, |results| {
150 count += 1;
151 checksum += results[0];
152 })
153 .unwrap();
154 assert_eq!(n, 1_000_000);
155 assert_eq!(count, 1_000_000);
156 assert!(checksum > 0.0);
157
158 assert_eq!(stream.stack.capacity(), stack_cap);
160 assert_eq!(stream.results.capacity(), results_cap);
161 assert_eq!(stream.params.capacity(), params_cap);
162 }
163
164 #[test]
165 fn streaming_wrong_arity_errors() {
166 let (_arena, eval) = build_eval();
167 let mut stream = StreamingEvaluator::new(&eval);
168 let rows: Vec<Vec<f64>> = vec![vec![1.0], vec![2.0, 3.0]];
169 assert!(stream.for_each(&rows, |_| {}).is_err());
170 }
171
172 #[test]
173 fn streaming_evaluate_chunk() {
174 let (_arena, eval) = build_eval();
175 let mut stream = StreamingEvaluator::new(&eval);
176 let rows: Vec<[f64; 2]> = vec![[1.0, 2.0], [3.0, 4.0]];
177 let out = stream.evaluate_chunk(&rows).unwrap();
178 assert_eq!(out.len(), 2);
179 assert!((out[0][0] - 3.0).abs() < 1e-10);
180 assert!((out[0][1] - 2.0).abs() < 1e-10);
181 assert!((out[1][0] - 7.0).abs() < 1e-10);
182 assert!((out[1][1] - 12.0).abs() < 1e-10);
183 }
184}