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        let seed_fact = ctx.get(ContextKey::Seeds).first();
125
126        // 1. Compute features using Polars
127        let features = match self.compute_features() {
128            Ok(f) => f,
129            Err(e) => {
130                let diagnostic = PRISM_PROVENANCE.proposed_fact(
131                    ContextKey::Diagnostic,
132                    "feature-agent-error",
133                    TextPayload::new(e.to_string()),
134                );
135                return AgentEffect::with_proposal(match seed_fact {
136                    Some(seed) => diagnostic.with_subject_from(seed),
137                    None => diagnostic,
138                });
139            }
140        };
141
142        // 2. Propose the features
143        let proposal =
144            PRISM_PROVENANCE.proposed_fact(ContextKey::Proposals, "features-001", features);
145        let proposal = match seed_fact {
146            Some(seed) => proposal.with_subject_from(seed),
147            None => proposal,
148        };
149
150        // Note: In a real agent, we might emit a Fact directly if trusted, or a ProposedFact.
151        // converge_core usually requires TryFrom implementation or specific flow.
152        // For simplicity, we assume we can emit effects.
153        // Wait, AgentEffect::with_proposal?
154        // Let's check AgentEffect definition.
155
156        // Use the constructor for single proposal
157        AgentEffect::with_proposal(proposal)
158    }
159}
160
161fn compute_features_from_df(
162    df: &DataFrame,
163    columns: Option<&FeatureColumns>,
164) -> Result<FeatureVector> {
165    let (left, right) = if let Some(columns) = columns {
166        let left = df
167            .column(&columns.left)
168            .map_err(|_| anyhow!("missing column {}", columns.left))?;
169        let right = df
170            .column(&columns.right)
171            .map_err(|_| anyhow!("missing column {}", columns.right))?;
172        (left.clone(), right.clone())
173    } else {
174        let mut numeric = df
175            .get_columns()
176            .iter()
177            .filter(|col| is_numeric_dtype(col.dtype()))
178            .cloned()
179            .collect::<Vec<_>>();
180        if numeric.len() < 2 {
181            return Err(anyhow!("need at least two numeric columns"));
182        }
183        (numeric.remove(0), numeric.remove(0))
184    };
185
186    if left.is_empty() || right.is_empty() {
187        return Err(anyhow!("input data is empty"));
188    }
189
190    let left = left.cast(&DataType::Float32)?;
191    let right = right.cast(&DataType::Float32)?;
192
193    let left_val = left
194        .f32()?
195        .get(0)
196        .ok_or_else(|| anyhow!("missing left value"))?;
197    let right_val = right
198        .f32()?
199        .get(0)
200        .ok_or_else(|| anyhow!("missing right value"))?;
201
202    let interaction = left_val * right_val;
203    Ok(FeatureVector::row(vec![left_val, right_val, interaction]))
204}
205
206fn load_dataframe(path: &Path) -> Result<DataFrame> {
207    let extension = path
208        .extension()
209        .and_then(|ext| ext.to_str())
210        .unwrap_or("")
211        .to_ascii_lowercase();
212
213    let path_str = path
214        .to_str()
215        .ok_or_else(|| anyhow!("path is not valid utf-8: {}", path.display()))?;
216
217    match extension.as_str() {
218        "parquet" => {
219            let pl_path = PlPath::new(path_str);
220            Ok(LazyFrame::scan_parquet(pl_path, Default::default())?.collect()?)
221        }
222        "csv" => Ok(CsvReadOptions::default()
223            .with_has_header(true)
224            .try_into_reader_with_file_path(Some(path.to_path_buf()))?
225            .finish()?),
226        _ => Err(anyhow!(
227            "unsupported data format for path {} (expected .csv or .parquet)",
228            path.display()
229        )),
230    }
231}
232
233fn is_numeric_dtype(dtype: &DataType) -> bool {
234    matches!(
235        dtype,
236        DataType::Int8
237            | DataType::Int16
238            | DataType::Int32
239            | DataType::Int64
240            | DataType::UInt8
241            | DataType::UInt16
242            | DataType::UInt32
243            | DataType::UInt64
244            | DataType::Float32
245            | DataType::Float64
246    )
247}
248
249#[cfg(test)]
250mod tests {
251    use super::*;
252    use proptest::prelude::*;
253    use std::collections::HashMap;
254    use std::fs;
255    use std::hint::black_box;
256    use std::time::Instant;
257    use std::time::{SystemTime, UNIX_EPOCH};
258
259    #[test]
260    fn feature_vector_validates_shape() {
261        let ok = FeatureVector::new(vec![1.0, 2.0], [1, 2]).unwrap();
262        assert_eq!(ok.rows(), 1);
263        assert_eq!(ok.cols(), 2);
264        assert!(FeatureVector::new(vec![1.0], [1, 2]).is_err());
265    }
266
267    #[test]
268    fn feature_vector_new_multi_row() {
269        let fv = FeatureVector::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], [2, 3]).unwrap();
270        assert_eq!(fv.rows(), 2);
271        assert_eq!(fv.cols(), 3);
272        assert_eq!(fv.data.len(), 6);
273    }
274
275    #[test]
276    fn feature_vector_new_rejects_mismatched_length() {
277        assert!(FeatureVector::new(vec![1.0, 2.0, 3.0], [2, 2]).is_err());
278        assert!(FeatureVector::new(vec![], [1, 1]).is_err());
279        assert!(FeatureVector::new(vec![1.0], [0, 1]).is_err());
280    }
281
282    #[test]
283    fn feature_vector_new_empty() {
284        let fv = FeatureVector::new(vec![], [0, 0]).unwrap();
285        assert_eq!(fv.rows(), 0);
286        assert_eq!(fv.cols(), 0);
287        assert!(fv.data.is_empty());
288    }
289
290    #[test]
291    fn feature_vector_new_zero_cols() {
292        let fv = FeatureVector::new(vec![], [5, 0]).unwrap();
293        assert_eq!(fv.rows(), 5);
294        assert_eq!(fv.cols(), 0);
295    }
296
297    #[test]
298    fn feature_vector_row_creates_single_row() {
299        let fv = FeatureVector::row(vec![10.0, 20.0, 30.0]);
300        assert_eq!(fv.rows(), 1);
301        assert_eq!(fv.cols(), 3);
302        assert_eq!(fv.data, vec![10.0, 20.0, 30.0]);
303    }
304
305    #[test]
306    fn feature_vector_row_empty() {
307        let fv = FeatureVector::row(vec![]);
308        assert_eq!(fv.rows(), 1);
309        assert_eq!(fv.cols(), 0);
310        assert!(fv.data.is_empty());
311    }
312
313    #[test]
314    fn feature_vector_row_single_element() {
315        let fv = FeatureVector::row(vec![42.0]);
316        assert_eq!(fv.rows(), 1);
317        assert_eq!(fv.cols(), 1);
318        assert_eq!(fv.data, vec![42.0]);
319    }
320
321    #[test]
322    fn feature_columns_construction() {
323        let fc = FeatureColumns {
324            left: "price".to_string(),
325            right: "quantity".to_string(),
326        };
327        assert_eq!(fc.left, "price");
328        assert_eq!(fc.right, "quantity");
329    }
330
331    #[test]
332    fn feature_columns_roundtrip_serde() {
333        let fc = FeatureColumns {
334            left: "a".to_string(),
335            right: "b".to_string(),
336        };
337        let json = serde_json::to_string(&fc).unwrap();
338        let deserialized: FeatureColumns = serde_json::from_str(&json).unwrap();
339        assert_eq!(deserialized.left, "a");
340        assert_eq!(deserialized.right, "b");
341    }
342
343    #[test]
344    fn feature_vector_roundtrip_serde() {
345        let fv = FeatureVector::new(vec![1.0, 2.0, 3.0, 4.0], [2, 2]).unwrap();
346        let json = serde_json::to_string(&fv).unwrap();
347        let deserialized: FeatureVector = serde_json::from_str(&json).unwrap();
348        assert_eq!(fv, deserialized);
349    }
350
351    #[test]
352    fn feature_agent_new_without_columns() {
353        let agent = FeatureAgent::new(None);
354        assert!(agent.source_path.is_none());
355        assert!(agent.columns.is_none());
356    }
357
358    #[test]
359    fn feature_agent_with_columns() {
360        let agent = FeatureAgent::new(None).with_columns("x", "y");
361        let cols = agent.columns.unwrap();
362        assert_eq!(cols.left, "x");
363        assert_eq!(cols.right, "y");
364    }
365
366    #[test]
367    fn feature_agent_with_source_path() {
368        let agent = FeatureAgent::new(Some(PathBuf::from("/tmp/data.csv")));
369        assert_eq!(agent.source_path.unwrap(), PathBuf::from("/tmp/data.csv"));
370    }
371
372    #[test]
373    fn is_numeric_dtype_covers_all_numeric_types() {
374        let numeric = [
375            DataType::Int8,
376            DataType::Int16,
377            DataType::Int32,
378            DataType::Int64,
379            DataType::UInt8,
380            DataType::UInt16,
381            DataType::UInt32,
382            DataType::UInt64,
383            DataType::Float32,
384            DataType::Float64,
385        ];
386        for dt in &numeric {
387            assert!(is_numeric_dtype(dt), "{dt:?} should be numeric");
388        }
389    }
390
391    #[test]
392    fn is_numeric_dtype_rejects_non_numeric() {
393        assert!(!is_numeric_dtype(&DataType::String));
394        assert!(!is_numeric_dtype(&DataType::Boolean));
395        assert!(!is_numeric_dtype(&DataType::Date));
396    }
397
398    #[test]
399    fn compute_features_rejects_empty_dataframe() {
400        let df = df![
401            "a" => Vec::<f32>::new(),
402            "b" => Vec::<f32>::new(),
403        ]
404        .unwrap();
405        let cols = FeatureColumns {
406            left: "a".into(),
407            right: "b".into(),
408        };
409        assert!(compute_features_from_df(&df, Some(&cols)).is_err());
410    }
411
412    #[test]
413    fn compute_features_rejects_missing_column() {
414        let df = df!["a" => [1.0f32]].unwrap();
415        let cols = FeatureColumns {
416            left: "a".into(),
417            right: "missing".into(),
418        };
419        assert!(compute_features_from_df(&df, Some(&cols)).is_err());
420    }
421
422    #[test]
423    fn compute_features_rejects_insufficient_numeric_columns() {
424        let df = df!["text" => ["a", "b"]].unwrap();
425        assert!(compute_features_from_df(&df, None).is_err());
426    }
427
428    proptest! {
429        #[test]
430        fn feature_vector_shape_invariant(
431            rows in 0usize..50,
432            cols in 0usize..50,
433        ) {
434            let len = rows.saturating_mul(cols);
435            let data = vec![0.0f32; len];
436            let fv = FeatureVector::new(data, [rows, cols]).unwrap();
437            prop_assert_eq!(fv.rows() * fv.cols(), fv.data.len());
438        }
439    }
440
441    #[test]
442    fn compute_features_from_df_uses_named_columns() {
443        let df = df![
444            "a" => [2.0f32, 3.0],
445            "b" => [4.0f32, 5.0],
446        ]
447        .unwrap();
448
449        let columns = FeatureColumns {
450            left: "a".into(),
451            right: "b".into(),
452        };
453        let features = compute_features_from_df(&df, Some(&columns)).unwrap();
454        assert_eq!(features.data, vec![2.0, 4.0, 8.0]);
455        assert_eq!(features.shape, [1, 3]);
456    }
457
458    #[test]
459    fn compute_features_from_df_falls_back_to_first_numeric_columns() {
460        let df = df![
461            "text" => ["x", "y"],
462            "a" => [1.5f32, 2.5],
463            "b" => [3.0f32, 4.0],
464        ]
465        .unwrap();
466
467        let features = compute_features_from_df(&df, None).unwrap();
468        assert_eq!(features.data, vec![1.5, 3.0, 4.5]);
469    }
470
471    #[test]
472    fn compute_features_handles_large_dataset() {
473        let rows = 10_000;
474        let left: Vec<f32> = (0..rows).map(|i| i as f32).collect();
475        let right: Vec<f32> = (0..rows).map(|i| (i as f32) + 1.0).collect();
476        let df = df![
477            "left" => left,
478            "right" => right,
479        ]
480        .unwrap();
481
482        let columns = FeatureColumns {
483            left: "left".into(),
484            right: "right".into(),
485        };
486        let features = compute_features_from_df(&df, Some(&columns)).unwrap();
487        assert_eq!(features.data, vec![0.0, 1.0, 0.0]);
488    }
489
490    #[test]
491    // RP-DETERMINISM exemption (QF-2026-06-08-01): SystemTime::now() is
492    // used only to mint a unique tempdir suffix. The value is never
493    // asserted on; the test result does not depend on the clock.
494    #[allow(clippy::disallowed_methods)]
495    fn load_dataframe_reads_csv() {
496        let mut path = std::env::temp_dir();
497        let nanos = SystemTime::now()
498            .duration_since(UNIX_EPOCH)
499            .unwrap()
500            .as_nanos();
501        path.push(format!("prism_{nanos}.csv"));
502
503        let contents = "left,right\n2.0,4.0\n3.0,5.0\n";
504        fs::write(&path, contents).unwrap();
505
506        let df = load_dataframe(&path).unwrap();
507        assert_eq!(df.height(), 2);
508        assert_eq!(df.width(), 2);
509    }
510
511    proptest! {
512        #[test]
513        fn compute_features_matches_first_row(
514            left in proptest::collection::vec(prop::num::f32::NORMAL, 1..50),
515            right in proptest::collection::vec(prop::num::f32::NORMAL, 1..50),
516        ) {
517            let len = left.len().min(right.len());
518            let df = df![
519                "left" => left[..len].to_vec(),
520                "right" => right[..len].to_vec(),
521            ]
522            .unwrap();
523
524            let columns = FeatureColumns {
525                left: "left".into(),
526                right: "right".into(),
527            };
528            let features = compute_features_from_df(&df, Some(&columns)).unwrap();
529            let expected_left = left[0];
530            let expected_right = right[0];
531            prop_assert_eq!(features.data, vec![expected_left, expected_right, expected_left * expected_right]);
532        }
533    }
534
535    #[test]
536    fn polars_vectorized_dot_product_matches_naive() {
537        let rows = 50_000;
538        let left: Vec<f32> = (0..rows).map(|i| (i % 100) as f32).collect();
539        let right: Vec<f32> = (0..rows).map(|i| ((i + 3) % 100) as f32).collect();
540        let df = df![
541            "left" => left.clone(),
542            "right" => right.clone(),
543        ]
544        .unwrap();
545
546        let product = (df.column("left").unwrap() * df.column("right").unwrap()).unwrap();
547        let polars_sum = product
548            .as_materialized_series()
549            .cast(&DataType::Float64)
550            .unwrap()
551            .f64()
552            .unwrap()
553            .sum()
554            .unwrap_or(0.0);
555
556        let mut naive_sum = 0.0f64;
557        for (l, r) in left.iter().zip(right.iter()) {
558            naive_sum += (*l as f64) * (*r as f64);
559        }
560
561        assert!((polars_sum - naive_sum).abs() < 1e-6);
562    }
563
564    #[test]
565    fn polars_groupby_sum_matches_naive() {
566        let rows = 10_000;
567        let keys: Vec<&str> = (0..rows)
568            .map(|i| {
569                if i % 3 == 0 {
570                    "alpha"
571                } else if i % 3 == 1 {
572                    "beta"
573                } else {
574                    "gamma"
575                }
576            })
577            .collect();
578        let values: Vec<f32> = (0..rows).map(|i| (i % 7) as f32).collect();
579        let df = df![
580            "key" => keys.clone(),
581            "value" => values.clone(),
582        ]
583        .unwrap();
584
585        let grouped = df
586            .lazy()
587            .group_by([col("key")])
588            .agg([col("value").sum().alias("value_sum")])
589            .collect()
590            .unwrap();
591        let keys_series = grouped.column("key").unwrap().str().unwrap();
592        let sums_series = grouped.column("value_sum").unwrap().f32().unwrap();
593
594        let mut naive = HashMap::<&str, f32>::new();
595        for (key, value) in keys.iter().zip(values.iter()) {
596            *naive.entry(*key).or_insert(0.0) += value;
597        }
598
599        for idx in 0..grouped.height() {
600            if let Some(key) = keys_series.get(idx) {
601                let polars_value = sums_series.get(idx).unwrap_or(0.0);
602                let naive_value = naive.get(key).copied().unwrap_or(0.0);
603                assert!((polars_value - naive_value).abs() < 1e-3);
604            }
605        }
606    }
607
608    #[test]
609    #[ignore]
610    // RP-DETERMINISM exemption (QF-2026-06-08-01): Instant::now() is
611    // used for benchmark timing instrumentation in an #[ignore]'d
612    // benchmark. The single assert! is a coarse 20x catastrophic-
613    // regression ceiling, not a precision check; minor clock noise
614    // doesn't change the verdict. This benchmark runs manually, not
615    // in default `cargo test`.
616    #[allow(clippy::disallowed_methods)]
617    fn polars_vectorized_dot_product_is_fast() {
618        let rows = 300_000;
619        let left: Vec<f32> = (0..rows).map(|i| (i % 100) as f32).collect();
620        let right: Vec<f32> = (0..rows).map(|i| ((i + 5) % 100) as f32).collect();
621
622        let df = df![
623            "left" => left.clone(),
624            "right" => right.clone(),
625        ]
626        .unwrap();
627
628        let polars_start = Instant::now();
629        let product = (df.column("left").unwrap() * df.column("right").unwrap()).unwrap();
630        let polars_sum = product
631            .as_materialized_series()
632            .f32()
633            .unwrap()
634            .sum()
635            .unwrap_or(0.0);
636        let polars_elapsed = polars_start.elapsed();
637        black_box(polars_sum);
638
639        let naive_start = Instant::now();
640        let mut naive_sum = 0.0f32;
641        for (l, r) in left.iter().zip(right.iter()) {
642            naive_sum += l * r;
643        }
644        let naive_elapsed = naive_start.elapsed();
645        black_box(naive_sum);
646
647        println!(
648            "polars dot product: {:?}, naive loop: {:?}",
649            polars_elapsed, naive_elapsed
650        );
651
652        assert!(polars_elapsed <= naive_elapsed * 20);
653    }
654
655    #[test]
656    #[ignore]
657    // RP-DETERMINISM exemption (QF-2026-06-08-01): Instant::now() is
658    // used for benchmark timing instrumentation in an #[ignore]'d
659    // benchmark. The single assert! is a coarse 20x catastrophic-
660    // regression ceiling, not a precision check; minor clock noise
661    // doesn't change the verdict. This benchmark runs manually, not
662    // in default `cargo test`.
663    #[allow(clippy::disallowed_methods)]
664    fn polars_groupby_is_fast() {
665        let rows = 200_000;
666        let keys: Vec<&str> = (0..rows)
667            .map(|i| {
668                if i % 4 == 0 {
669                    "alpha"
670                } else if i % 4 == 1 {
671                    "beta"
672                } else if i % 4 == 2 {
673                    "gamma"
674                } else {
675                    "delta"
676                }
677            })
678            .collect();
679        let values: Vec<f32> = (0..rows).map(|i| (i % 9) as f32).collect();
680        let df = df![
681            "key" => keys.clone(),
682            "value" => values.clone(),
683        ]
684        .unwrap();
685
686        let polars_start = Instant::now();
687        let grouped = df
688            .lazy()
689            .group_by([col("key")])
690            .agg([col("value").sum().alias("value_sum")])
691            .collect()
692            .unwrap();
693        let polars_elapsed = polars_start.elapsed();
694        black_box(grouped.height());
695
696        let naive_start = Instant::now();
697        let mut naive = HashMap::<&str, f32>::new();
698        for (key, value) in keys.iter().zip(values.iter()) {
699            *naive.entry(*key).or_insert(0.0) += value;
700        }
701        let naive_elapsed = naive_start.elapsed();
702        black_box(naive.len());
703
704        println!(
705            "polars groupby: {:?}, naive hashmap: {:?}",
706            polars_elapsed, naive_elapsed
707        );
708
709        assert!(polars_elapsed <= naive_elapsed * 20);
710    }
711}