1use 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#[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 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 &[ContextKey::Seeds]
112 }
113
114 fn accepts(&self, ctx: &dyn Context) -> bool {
115 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 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 let proposal =
138 PRISM_PROVENANCE.proposed_fact(ContextKey::Proposals, "features-001", features);
139
140 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}