Skip to main content

prism/
engine.rs

1// Copyright 2024-2026 Reflective Labs
2
3use anyhow::{Result, anyhow};
4use converge_pack::{
5    AgentEffect, Context, ContextKey, FactPayload, Provenance, ProvenanceSource, Suggestor,
6    TextPayload,
7};
8use polars::prelude::*;
9use serde::{Deserialize, Serialize};
10use std::path::{Path, PathBuf};
11
12use crate::provenance::PRISM_PROVENANCE;
13
14/// Typed payload representing computed features.
15#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
16#[serde(deny_unknown_fields)]
17pub struct FeatureVector {
18    pub data: Vec<f32>,
19    pub shape: [usize; 2],
20}
21
22impl FactPayload for FeatureVector {
23    const FAMILY: &'static str = "prism.feature-vector";
24    const VERSION: u16 = 1;
25}
26
27impl FeatureVector {
28    pub fn new(data: Vec<f32>, shape: [usize; 2]) -> Result<Self> {
29        let expected = shape
30            .first()
31            .and_then(|rows| shape.get(1).map(|cols| rows.saturating_mul(*cols)))
32            .unwrap_or(0);
33        if data.len() != expected {
34            return Err(anyhow!(
35                "feature data length {} does not match shape {:?}",
36                data.len(),
37                shape
38            ));
39        }
40        Ok(Self { data, shape })
41    }
42
43    pub fn row(data: Vec<f32>) -> Self {
44        let cols = data.len();
45        Self {
46            data,
47            shape: [1, cols],
48        }
49    }
50
51    pub fn rows(&self) -> usize {
52        self.shape[0]
53    }
54
55    pub fn cols(&self) -> usize {
56        self.shape[1]
57    }
58}
59
60#[derive(Clone, Debug, Serialize, Deserialize)]
61pub struct FeatureColumns {
62    pub left: String,
63    pub right: String,
64}
65
66#[derive(Clone, Debug)]
67pub struct FeatureAgent {
68    source_path: Option<PathBuf>,
69    columns: Option<FeatureColumns>,
70}
71
72impl FeatureAgent {
73    pub fn new(source_path: Option<PathBuf>) -> Self {
74        Self {
75            source_path,
76            columns: None,
77        }
78    }
79
80    pub fn with_columns(mut self, left: impl Into<String>, right: impl Into<String>) -> Self {
81        self.columns = Some(FeatureColumns {
82            left: left.into(),
83            right: right.into(),
84        });
85        self
86    }
87
88    /// Internal Polars logic to compute features
89    fn compute_features(&self) -> Result<FeatureVector> {
90        let df = if let Some(path) = &self.source_path {
91            load_dataframe(path)?
92        } else {
93            df! [
94                "x1" => [1.0, 2.0, 3.0],
95                "x2" => [4.0, 5.0, 6.0],
96                "x3" => [7.0, 8.0, 9.0],
97            ]?
98        };
99        compute_features_from_df(&df, self.columns.as_ref())
100    }
101}
102
103#[async_trait::async_trait]
104impl Suggestor for FeatureAgent {
105    fn name(&self) -> &'static str {
106        "FeatureAgent (Polars)"
107    }
108
109    fn dependencies(&self) -> &[ContextKey] {
110        // Depends on Seeds to know WHAT to process
111        &[ContextKey::Seeds]
112    }
113
114    fn accepts(&self, ctx: &dyn Context) -> bool {
115        // Run if we have Seeds but haven't produced Proposals yet
116        ctx.has(ContextKey::Seeds) && !ctx.has(ContextKey::Proposals)
117    }
118
119    fn provenance(&self) -> Provenance {
120        PRISM_PROVENANCE.provenance()
121    }
122
123    async fn execute(&self, _ctx: &dyn Context) -> AgentEffect {
124        // 1. Compute features using Polars
125        let features = match self.compute_features() {
126            Ok(f) => f,
127            Err(e) => {
128                return AgentEffect::with_proposal(PRISM_PROVENANCE.proposed_fact(
129                    ContextKey::Diagnostic,
130                    "feature-agent-error",
131                    TextPayload::new(e.to_string()),
132                ));
133            }
134        };
135
136        // 2. Propose the features
137        let proposal =
138            PRISM_PROVENANCE.proposed_fact(ContextKey::Proposals, "features-001", features);
139
140        // Note: In a real agent, we might emit a Fact directly if trusted, or a ProposedFact.
141        // converge_core usually requires TryFrom implementation or specific flow.
142        // For simplicity, we assume we can emit effects.
143        // Wait, AgentEffect::with_proposal?
144        // Let's check AgentEffect definition.
145
146        // Use the constructor for single proposal
147        AgentEffect::with_proposal(proposal)
148    }
149}
150
151fn compute_features_from_df(
152    df: &DataFrame,
153    columns: Option<&FeatureColumns>,
154) -> Result<FeatureVector> {
155    let (left, right) = if let Some(columns) = columns {
156        let left = df
157            .column(&columns.left)
158            .map_err(|_| anyhow!("missing column {}", columns.left))?;
159        let right = df
160            .column(&columns.right)
161            .map_err(|_| anyhow!("missing column {}", columns.right))?;
162        (left.clone(), right.clone())
163    } else {
164        let mut numeric = df
165            .get_columns()
166            .iter()
167            .filter(|col| is_numeric_dtype(col.dtype()))
168            .cloned()
169            .collect::<Vec<_>>();
170        if numeric.len() < 2 {
171            return Err(anyhow!("need at least two numeric columns"));
172        }
173        (numeric.remove(0), numeric.remove(0))
174    };
175
176    if left.is_empty() || right.is_empty() {
177        return Err(anyhow!("input data is empty"));
178    }
179
180    let left = left.cast(&DataType::Float32)?;
181    let right = right.cast(&DataType::Float32)?;
182
183    let left_val = left
184        .f32()?
185        .get(0)
186        .ok_or_else(|| anyhow!("missing left value"))?;
187    let right_val = right
188        .f32()?
189        .get(0)
190        .ok_or_else(|| anyhow!("missing right value"))?;
191
192    let interaction = left_val * right_val;
193    Ok(FeatureVector::row(vec![left_val, right_val, interaction]))
194}
195
196fn load_dataframe(path: &Path) -> Result<DataFrame> {
197    let extension = path
198        .extension()
199        .and_then(|ext| ext.to_str())
200        .unwrap_or("")
201        .to_ascii_lowercase();
202
203    let path_str = path
204        .to_str()
205        .ok_or_else(|| anyhow!("path is not valid utf-8: {}", path.display()))?;
206
207    match extension.as_str() {
208        "parquet" => {
209            let pl_path = PlPath::new(path_str);
210            Ok(LazyFrame::scan_parquet(pl_path, Default::default())?.collect()?)
211        }
212        "csv" => Ok(CsvReadOptions::default()
213            .with_has_header(true)
214            .try_into_reader_with_file_path(Some(path.to_path_buf()))?
215            .finish()?),
216        _ => Err(anyhow!(
217            "unsupported data format for path {} (expected .csv or .parquet)",
218            path.display()
219        )),
220    }
221}
222
223fn is_numeric_dtype(dtype: &DataType) -> bool {
224    matches!(
225        dtype,
226        DataType::Int8
227            | DataType::Int16
228            | DataType::Int32
229            | DataType::Int64
230            | DataType::UInt8
231            | DataType::UInt16
232            | DataType::UInt32
233            | DataType::UInt64
234            | DataType::Float32
235            | DataType::Float64
236    )
237}
238
239#[cfg(test)]
240mod tests {
241    use super::*;
242    use proptest::prelude::*;
243    use std::collections::HashMap;
244    use std::fs;
245    use std::hint::black_box;
246    use std::time::Instant;
247    use std::time::{SystemTime, UNIX_EPOCH};
248
249    #[test]
250    fn feature_vector_validates_shape() {
251        let ok = FeatureVector::new(vec![1.0, 2.0], [1, 2]).unwrap();
252        assert_eq!(ok.rows(), 1);
253        assert_eq!(ok.cols(), 2);
254        assert!(FeatureVector::new(vec![1.0], [1, 2]).is_err());
255    }
256
257    #[test]
258    fn feature_vector_new_multi_row() {
259        let fv = FeatureVector::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], [2, 3]).unwrap();
260        assert_eq!(fv.rows(), 2);
261        assert_eq!(fv.cols(), 3);
262        assert_eq!(fv.data.len(), 6);
263    }
264
265    #[test]
266    fn feature_vector_new_rejects_mismatched_length() {
267        assert!(FeatureVector::new(vec![1.0, 2.0, 3.0], [2, 2]).is_err());
268        assert!(FeatureVector::new(vec![], [1, 1]).is_err());
269        assert!(FeatureVector::new(vec![1.0], [0, 1]).is_err());
270    }
271
272    #[test]
273    fn feature_vector_new_empty() {
274        let fv = FeatureVector::new(vec![], [0, 0]).unwrap();
275        assert_eq!(fv.rows(), 0);
276        assert_eq!(fv.cols(), 0);
277        assert!(fv.data.is_empty());
278    }
279
280    #[test]
281    fn feature_vector_new_zero_cols() {
282        let fv = FeatureVector::new(vec![], [5, 0]).unwrap();
283        assert_eq!(fv.rows(), 5);
284        assert_eq!(fv.cols(), 0);
285    }
286
287    #[test]
288    fn feature_vector_row_creates_single_row() {
289        let fv = FeatureVector::row(vec![10.0, 20.0, 30.0]);
290        assert_eq!(fv.rows(), 1);
291        assert_eq!(fv.cols(), 3);
292        assert_eq!(fv.data, vec![10.0, 20.0, 30.0]);
293    }
294
295    #[test]
296    fn feature_vector_row_empty() {
297        let fv = FeatureVector::row(vec![]);
298        assert_eq!(fv.rows(), 1);
299        assert_eq!(fv.cols(), 0);
300        assert!(fv.data.is_empty());
301    }
302
303    #[test]
304    fn feature_vector_row_single_element() {
305        let fv = FeatureVector::row(vec![42.0]);
306        assert_eq!(fv.rows(), 1);
307        assert_eq!(fv.cols(), 1);
308        assert_eq!(fv.data, vec![42.0]);
309    }
310
311    #[test]
312    fn feature_columns_construction() {
313        let fc = FeatureColumns {
314            left: "price".to_string(),
315            right: "quantity".to_string(),
316        };
317        assert_eq!(fc.left, "price");
318        assert_eq!(fc.right, "quantity");
319    }
320
321    #[test]
322    fn feature_columns_roundtrip_serde() {
323        let fc = FeatureColumns {
324            left: "a".to_string(),
325            right: "b".to_string(),
326        };
327        let json = serde_json::to_string(&fc).unwrap();
328        let deserialized: FeatureColumns = serde_json::from_str(&json).unwrap();
329        assert_eq!(deserialized.left, "a");
330        assert_eq!(deserialized.right, "b");
331    }
332
333    #[test]
334    fn feature_vector_roundtrip_serde() {
335        let fv = FeatureVector::new(vec![1.0, 2.0, 3.0, 4.0], [2, 2]).unwrap();
336        let json = serde_json::to_string(&fv).unwrap();
337        let deserialized: FeatureVector = serde_json::from_str(&json).unwrap();
338        assert_eq!(fv, deserialized);
339    }
340
341    #[test]
342    fn feature_agent_new_without_columns() {
343        let agent = FeatureAgent::new(None);
344        assert!(agent.source_path.is_none());
345        assert!(agent.columns.is_none());
346    }
347
348    #[test]
349    fn feature_agent_with_columns() {
350        let agent = FeatureAgent::new(None).with_columns("x", "y");
351        let cols = agent.columns.unwrap();
352        assert_eq!(cols.left, "x");
353        assert_eq!(cols.right, "y");
354    }
355
356    #[test]
357    fn feature_agent_with_source_path() {
358        let agent = FeatureAgent::new(Some(PathBuf::from("/tmp/data.csv")));
359        assert_eq!(agent.source_path.unwrap(), PathBuf::from("/tmp/data.csv"));
360    }
361
362    #[test]
363    fn is_numeric_dtype_covers_all_numeric_types() {
364        let numeric = [
365            DataType::Int8,
366            DataType::Int16,
367            DataType::Int32,
368            DataType::Int64,
369            DataType::UInt8,
370            DataType::UInt16,
371            DataType::UInt32,
372            DataType::UInt64,
373            DataType::Float32,
374            DataType::Float64,
375        ];
376        for dt in &numeric {
377            assert!(is_numeric_dtype(dt), "{dt:?} should be numeric");
378        }
379    }
380
381    #[test]
382    fn is_numeric_dtype_rejects_non_numeric() {
383        assert!(!is_numeric_dtype(&DataType::String));
384        assert!(!is_numeric_dtype(&DataType::Boolean));
385        assert!(!is_numeric_dtype(&DataType::Date));
386    }
387
388    #[test]
389    fn compute_features_rejects_empty_dataframe() {
390        let df = df![
391            "a" => Vec::<f32>::new(),
392            "b" => Vec::<f32>::new(),
393        ]
394        .unwrap();
395        let cols = FeatureColumns {
396            left: "a".into(),
397            right: "b".into(),
398        };
399        assert!(compute_features_from_df(&df, Some(&cols)).is_err());
400    }
401
402    #[test]
403    fn compute_features_rejects_missing_column() {
404        let df = df!["a" => [1.0f32]].unwrap();
405        let cols = FeatureColumns {
406            left: "a".into(),
407            right: "missing".into(),
408        };
409        assert!(compute_features_from_df(&df, Some(&cols)).is_err());
410    }
411
412    #[test]
413    fn compute_features_rejects_insufficient_numeric_columns() {
414        let df = df!["text" => ["a", "b"]].unwrap();
415        assert!(compute_features_from_df(&df, None).is_err());
416    }
417
418    proptest! {
419        #[test]
420        fn feature_vector_shape_invariant(
421            rows in 0usize..50,
422            cols in 0usize..50,
423        ) {
424            let len = rows.saturating_mul(cols);
425            let data = vec![0.0f32; len];
426            let fv = FeatureVector::new(data, [rows, cols]).unwrap();
427            prop_assert_eq!(fv.rows() * fv.cols(), fv.data.len());
428        }
429    }
430
431    #[test]
432    fn compute_features_from_df_uses_named_columns() {
433        let df = df![
434            "a" => [2.0f32, 3.0],
435            "b" => [4.0f32, 5.0],
436        ]
437        .unwrap();
438
439        let columns = FeatureColumns {
440            left: "a".into(),
441            right: "b".into(),
442        };
443        let features = compute_features_from_df(&df, Some(&columns)).unwrap();
444        assert_eq!(features.data, vec![2.0, 4.0, 8.0]);
445        assert_eq!(features.shape, [1, 3]);
446    }
447
448    #[test]
449    fn compute_features_from_df_falls_back_to_first_numeric_columns() {
450        let df = df![
451            "text" => ["x", "y"],
452            "a" => [1.5f32, 2.5],
453            "b" => [3.0f32, 4.0],
454        ]
455        .unwrap();
456
457        let features = compute_features_from_df(&df, None).unwrap();
458        assert_eq!(features.data, vec![1.5, 3.0, 4.5]);
459    }
460
461    #[test]
462    fn compute_features_handles_large_dataset() {
463        let rows = 10_000;
464        let left: Vec<f32> = (0..rows).map(|i| i as f32).collect();
465        let right: Vec<f32> = (0..rows).map(|i| (i as f32) + 1.0).collect();
466        let df = df![
467            "left" => left,
468            "right" => right,
469        ]
470        .unwrap();
471
472        let columns = FeatureColumns {
473            left: "left".into(),
474            right: "right".into(),
475        };
476        let features = compute_features_from_df(&df, Some(&columns)).unwrap();
477        assert_eq!(features.data, vec![0.0, 1.0, 0.0]);
478    }
479
480    #[test]
481    fn load_dataframe_reads_csv() {
482        let mut path = std::env::temp_dir();
483        let nanos = SystemTime::now()
484            .duration_since(UNIX_EPOCH)
485            .unwrap()
486            .as_nanos();
487        path.push(format!("prism_{nanos}.csv"));
488
489        let contents = "left,right\n2.0,4.0\n3.0,5.0\n";
490        fs::write(&path, contents).unwrap();
491
492        let df = load_dataframe(&path).unwrap();
493        assert_eq!(df.height(), 2);
494        assert_eq!(df.width(), 2);
495    }
496
497    proptest! {
498        #[test]
499        fn compute_features_matches_first_row(
500            left in proptest::collection::vec(prop::num::f32::NORMAL, 1..50),
501            right in proptest::collection::vec(prop::num::f32::NORMAL, 1..50),
502        ) {
503            let len = left.len().min(right.len());
504            let df = df![
505                "left" => left[..len].to_vec(),
506                "right" => right[..len].to_vec(),
507            ]
508            .unwrap();
509
510            let columns = FeatureColumns {
511                left: "left".into(),
512                right: "right".into(),
513            };
514            let features = compute_features_from_df(&df, Some(&columns)).unwrap();
515            let expected_left = left[0];
516            let expected_right = right[0];
517            prop_assert_eq!(features.data, vec![expected_left, expected_right, expected_left * expected_right]);
518        }
519    }
520
521    #[test]
522    fn polars_vectorized_dot_product_matches_naive() {
523        let rows = 50_000;
524        let left: Vec<f32> = (0..rows).map(|i| (i % 100) as f32).collect();
525        let right: Vec<f32> = (0..rows).map(|i| ((i + 3) % 100) as f32).collect();
526        let df = df![
527            "left" => left.clone(),
528            "right" => right.clone(),
529        ]
530        .unwrap();
531
532        let product = (df.column("left").unwrap() * df.column("right").unwrap()).unwrap();
533        let polars_sum = product
534            .as_materialized_series()
535            .cast(&DataType::Float64)
536            .unwrap()
537            .f64()
538            .unwrap()
539            .sum()
540            .unwrap_or(0.0);
541
542        let mut naive_sum = 0.0f64;
543        for (l, r) in left.iter().zip(right.iter()) {
544            naive_sum += (*l as f64) * (*r as f64);
545        }
546
547        assert!((polars_sum - naive_sum).abs() < 1e-6);
548    }
549
550    #[test]
551    fn polars_groupby_sum_matches_naive() {
552        let rows = 10_000;
553        let keys: Vec<&str> = (0..rows)
554            .map(|i| {
555                if i % 3 == 0 {
556                    "alpha"
557                } else if i % 3 == 1 {
558                    "beta"
559                } else {
560                    "gamma"
561                }
562            })
563            .collect();
564        let values: Vec<f32> = (0..rows).map(|i| (i % 7) as f32).collect();
565        let df = df![
566            "key" => keys.clone(),
567            "value" => values.clone(),
568        ]
569        .unwrap();
570
571        let grouped = df
572            .lazy()
573            .group_by([col("key")])
574            .agg([col("value").sum().alias("value_sum")])
575            .collect()
576            .unwrap();
577        let keys_series = grouped.column("key").unwrap().str().unwrap();
578        let sums_series = grouped.column("value_sum").unwrap().f32().unwrap();
579
580        let mut naive = HashMap::<&str, f32>::new();
581        for (key, value) in keys.iter().zip(values.iter()) {
582            *naive.entry(*key).or_insert(0.0) += value;
583        }
584
585        for idx in 0..grouped.height() {
586            if let Some(key) = keys_series.get(idx) {
587                let polars_value = sums_series.get(idx).unwrap_or(0.0);
588                let naive_value = naive.get(key).copied().unwrap_or(0.0);
589                assert!((polars_value - naive_value).abs() < 1e-3);
590            }
591        }
592    }
593
594    #[test]
595    #[ignore]
596    fn polars_vectorized_dot_product_is_fast() {
597        let rows = 300_000;
598        let left: Vec<f32> = (0..rows).map(|i| (i % 100) as f32).collect();
599        let right: Vec<f32> = (0..rows).map(|i| ((i + 5) % 100) as f32).collect();
600
601        let df = df![
602            "left" => left.clone(),
603            "right" => right.clone(),
604        ]
605        .unwrap();
606
607        let polars_start = Instant::now();
608        let product = (df.column("left").unwrap() * df.column("right").unwrap()).unwrap();
609        let polars_sum = product
610            .as_materialized_series()
611            .f32()
612            .unwrap()
613            .sum()
614            .unwrap_or(0.0);
615        let polars_elapsed = polars_start.elapsed();
616        black_box(polars_sum);
617
618        let naive_start = Instant::now();
619        let mut naive_sum = 0.0f32;
620        for (l, r) in left.iter().zip(right.iter()) {
621            naive_sum += l * r;
622        }
623        let naive_elapsed = naive_start.elapsed();
624        black_box(naive_sum);
625
626        println!(
627            "polars dot product: {:?}, naive loop: {:?}",
628            polars_elapsed, naive_elapsed
629        );
630
631        assert!(polars_elapsed <= naive_elapsed * 20);
632    }
633
634    #[test]
635    #[ignore]
636    fn polars_groupby_is_fast() {
637        let rows = 200_000;
638        let keys: Vec<&str> = (0..rows)
639            .map(|i| {
640                if i % 4 == 0 {
641                    "alpha"
642                } else if i % 4 == 1 {
643                    "beta"
644                } else if i % 4 == 2 {
645                    "gamma"
646                } else {
647                    "delta"
648                }
649            })
650            .collect();
651        let values: Vec<f32> = (0..rows).map(|i| (i % 9) as f32).collect();
652        let df = df![
653            "key" => keys.clone(),
654            "value" => values.clone(),
655        ]
656        .unwrap();
657
658        let polars_start = Instant::now();
659        let grouped = df
660            .lazy()
661            .group_by([col("key")])
662            .agg([col("value").sum().alias("value_sum")])
663            .collect()
664            .unwrap();
665        let polars_elapsed = polars_start.elapsed();
666        black_box(grouped.height());
667
668        let naive_start = Instant::now();
669        let mut naive = HashMap::<&str, f32>::new();
670        for (key, value) in keys.iter().zip(values.iter()) {
671            *naive.entry(*key).or_insert(0.0) += value;
672        }
673        let naive_elapsed = naive_start.elapsed();
674        black_box(naive.len());
675
676        println!(
677            "polars groupby: {:?}, naive hashmap: {:?}",
678            polars_elapsed, naive_elapsed
679        );
680
681        assert!(polars_elapsed <= naive_elapsed * 20);
682    }
683}