antecedent_data/
project.rs1use std::collections::{HashMap, HashSet};
6
7use antecedent_core::{CausalSchemaBuilder, VariableId};
8
9use crate::dataset::TabularData;
10use crate::error::DataError;
11use crate::storage::OwnedColumnarStorage;
12use crate::table::TableView;
13
14#[derive(Clone, Debug)]
18pub struct IdRemap {
19 old_to_new: HashMap<VariableId, VariableId>,
20}
21
22impl IdRemap {
23 pub fn map(&self, old: VariableId) -> Result<VariableId, DataError> {
29 self.old_to_new.get(&old).copied().ok_or(DataError::UnknownVariable { id: old })
30 }
31
32 #[must_use]
34 pub fn len(&self) -> usize {
35 self.old_to_new.len()
36 }
37
38 #[must_use]
40 pub fn is_empty(&self) -> bool {
41 self.old_to_new.is_empty()
42 }
43}
44
45impl TabularData {
46 pub fn project(&self, ids: &[VariableId]) -> Result<(Self, IdRemap), DataError> {
55 let mut seen = HashSet::new();
56 let mut ordered: Vec<VariableId> = Vec::with_capacity(ids.len());
57 for &id in ids {
58 if seen.insert(id) {
59 ordered.push(id);
60 }
61 }
62 if ordered.is_empty() {
63 return Err(DataError::EmptySelection {
64 context: "column projection: no variables requested",
65 });
66 }
67
68 let storage = self.storage();
69 let schema = storage.schema();
70 let mut builder = CausalSchemaBuilder::new();
71 let mut cols = Vec::with_capacity(ordered.len());
72 let mut old_to_new = HashMap::with_capacity(ordered.len());
73
74 for (new_idx, &old_id) in ordered.iter().enumerate() {
75 let meta = schema.get(old_id).map_err(|_| DataError::UnknownVariable { id: old_id })?;
76 builder
77 .add_variable(
78 std::sync::Arc::clone(&meta.name),
79 meta.value_type.clone(),
80 meta.role_hints,
81 meta.unit.clone(),
82 meta.category_domain,
83 meta.measurement.clone(),
84 )
85 .map_err(|e| DataError::Schema(e.to_string()))?;
86 let new_id = VariableId::from_raw(u32::try_from(new_idx).map_err(|_| {
87 DataError::InvalidArgument {
88 message: "projected schema exceeds VariableId range".into(),
89 }
90 })?);
91 old_to_new.insert(old_id, new_id);
92 let col = storage
93 .columns()
94 .get(old_id.as_usize())
95 .ok_or(DataError::UnknownVariable { id: old_id })?;
96 cols.push(col.with_id(new_id));
97 }
98
99 let new_schema = builder.build().map_err(|e| DataError::Schema(e.to_string()))?;
100 let new_storage = OwnedColumnarStorage::try_new(
101 new_schema,
102 cols,
103 storage.analysis_mask().cloned(),
104 storage.weights().map(std::sync::Arc::from),
105 )?;
106 Ok((Self::new(new_storage), IdRemap { old_to_new }))
107 }
108}
109
110#[must_use]
112pub fn dedupe_variable_ids(ids: impl IntoIterator<Item = VariableId>) -> Vec<VariableId> {
113 let mut seen = HashSet::new();
114 let mut out = Vec::new();
115 for id in ids {
116 if seen.insert(id) {
117 out.push(id);
118 }
119 }
120 out
121}
122
123#[cfg(test)]
124mod tests {
125 #![allow(clippy::cast_precision_loss)]
126
127 use super::*;
128 use crate::column::{Float64Column, OwnedColumn, ValidityBitmap};
129 use antecedent_core::{MeasurementSpec, RoleHint, SmallRoleSet, ValueType};
130 use std::sync::Arc;
131
132 fn float_table(names: &[&str], rows: usize) -> TabularData {
133 let mut b = CausalSchemaBuilder::new();
134 for name in names {
135 b.add_variable(
136 *name,
137 ValueType::Continuous,
138 SmallRoleSet::from_hint(RoleHint::Context),
139 None,
140 None,
141 MeasurementSpec::default(),
142 )
143 .unwrap();
144 }
145 let schema = b.build().unwrap();
146 let cols: Vec<OwnedColumn> = names
147 .iter()
148 .enumerate()
149 .map(|(i, _)| {
150 let id = VariableId::from_raw(u32::try_from(i).unwrap());
151 let values: Arc<[f64]> =
152 (0..rows).map(|r| (r + i * 100) as f64).collect::<Vec<_>>().into();
153 OwnedColumn::Float64(
154 Float64Column::new(id, values, ValidityBitmap::all_valid(rows)).unwrap(),
155 )
156 })
157 .collect();
158 TabularData::new(OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap())
159 }
160
161 #[test]
162 fn project_preserves_values_and_remaps_ids() {
163 let data = float_table(&["noise", "t", "y", "z", "extra"], 4);
164 let t = VariableId::from_raw(1);
165 let y = VariableId::from_raw(2);
166 let z = VariableId::from_raw(3);
167 let (proj, remap) = data.project(&[t, y, z]).unwrap();
168 assert_eq!(proj.schema().len(), 3);
169 assert_eq!(proj.schema().get(VariableId::from_raw(0)).unwrap().name.as_ref(), "t");
170 assert_eq!(proj.schema().get(VariableId::from_raw(1)).unwrap().name.as_ref(), "y");
171 assert_eq!(proj.schema().get(VariableId::from_raw(2)).unwrap().name.as_ref(), "z");
172 assert_eq!(remap.map(t).unwrap(), VariableId::from_raw(0));
173 assert_eq!(remap.map(y).unwrap(), VariableId::from_raw(1));
174 assert_eq!(remap.map(z).unwrap(), VariableId::from_raw(2));
175 assert!(remap.map(VariableId::from_raw(0)).is_err());
176
177 let view = proj.column(VariableId::from_raw(0)).unwrap();
178 let crate::column::ColumnView::Float64(c) = view else {
179 panic!("expected float64");
180 };
181 assert_eq!(c.values.as_slice(), &[100.0, 101.0, 102.0, 103.0]);
182 }
183
184 #[test]
185 fn project_shares_float_buffers() {
186 let data = float_table(&["t", "y", "noise"], 8);
187 let t = VariableId::from_raw(0);
188 let y = VariableId::from_raw(1);
189 let before = match data.storage().columns()[0] {
190 OwnedColumn::Float64(ref c) => c.values.as_slice().as_ptr(),
191 _ => panic!("float"),
192 };
193 let (proj, _) = data.project(&[t, y]).unwrap();
194 let after = match proj.storage().columns()[0] {
195 OwnedColumn::Float64(ref c) => c.values.as_slice().as_ptr(),
196 _ => panic!("float"),
197 };
198 assert_eq!(before, after);
199 }
200}