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 seed_fact = ctx.get(ContextKey::Seeds).first();
125
126 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 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 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 #[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 #[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 #[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}