1use std::sync::Arc;
6
7use antecedent_core::{
8 CausalSchemaBuilder, MeasurementSpec, RoleHint, SmallRoleSet, ValueType, VariableId,
9};
10
11use crate::column::{ColumnView, Float64Column, OwnedColumn, ValidityBitmap};
12use crate::dataset::TabularData;
13use crate::error::DataError;
14use crate::storage::OwnedColumnarStorage;
15use crate::table::TableView;
16
17impl TabularData {
18 pub fn complete_case_mask(&self, ids: &[VariableId]) -> Result<Vec<bool>, DataError> {
24 let mut keep = Vec::new();
25 self.complete_case_mask_into(ids, &mut keep)?;
26 Ok(keep)
27 }
28
29 pub fn complete_case_mask_into(
36 &self,
37 ids: &[VariableId],
38 out: &mut Vec<bool>,
39 ) -> Result<(), DataError> {
40 let n = self.row_count();
41 out.clear();
42 out.resize(n, true);
43 let keep = out;
44 if let Some(mask) = self.storage().analysis_mask() {
45 for (i, slot) in keep.iter_mut().enumerate() {
46 *slot = mask.is_valid(i);
47 }
48 }
49 for &id in ids {
50 let validity = self.column(id)?.validity();
51 for (i, slot) in keep.iter_mut().enumerate() {
52 if *slot && !validity.is_valid(i) {
53 *slot = false;
54 }
55 }
56 }
57 if !keep.iter().any(|k| *k) {
58 return Err(DataError::EmptySelection {
59 context: "complete-case mask after validity/analysis filtering",
60 });
61 }
62 Ok(())
63 }
64
65 pub fn float64_masked(&self, id: VariableId, keep: &[bool]) -> Result<Vec<f64>, DataError> {
71 if keep.len() != self.row_count() {
72 return Err(DataError::LengthMismatch {
73 expected: self.row_count(),
74 actual: keep.len(),
75 context: "complete-case keep mask",
76 });
77 }
78 let ColumnView::Float64(c) = self.column(id)? else {
79 return Err(DataError::TypeMismatch { id, expected: "float64" });
80 };
81 let mut out = Vec::with_capacity(keep.iter().filter(|k| **k).count());
82 for (i, &k) in keep.iter().enumerate() {
83 if k {
84 out.push(c.values[i]);
85 }
86 }
87 Ok(out)
88 }
89
90 pub fn with_replaced_float(
98 &self,
99 id: VariableId,
100 values: Arc<[f64]>,
101 ) -> Result<Self, DataError> {
102 self.with_replaced_floats(&[(id, values)])
103 }
104
105 pub fn with_replaced_floats(
118 &self,
119 replacements: &[(VariableId, Arc<[f64]>)],
120 ) -> Result<Self, DataError> {
121 let n = self.row_count();
122 let storage = self.storage();
123 let mut cols: Vec<OwnedColumn> = storage.columns().to_vec();
124 for (id, values) in replacements {
125 let id = *id;
126 if values.len() != n {
127 return Err(DataError::LengthMismatch {
128 expected: n,
129 actual: values.len(),
130 context: "replacement float column",
131 });
132 }
133 let idx = id.as_usize();
134 if idx >= cols.len() {
135 return Err(DataError::UnknownVariable { id });
136 }
137 if !matches!(cols[idx], OwnedColumn::Float64(_)) {
138 return Err(DataError::TypeMismatch { id, expected: "float64" });
139 }
140 cols[idx] = OwnedColumn::Float64(Float64Column::new(
141 id,
142 Arc::clone(values),
143 ValidityBitmap::all_valid(n),
144 )?);
145 }
146 let storage = OwnedColumnarStorage::try_new(
147 storage.schema().clone(),
148 cols,
149 storage.analysis_mask().cloned(),
150 storage.weights().map(Arc::from),
151 )?;
152 Ok(Self::new(storage))
153 }
154
155 pub fn with_analysis_mask(&self, mask: ValidityBitmap) -> Result<Self, DataError> {
162 let storage = self.storage();
163 let n = storage.row_count();
164 if mask.len() != n {
165 return Err(DataError::LengthMismatch {
166 expected: n,
167 actual: mask.len(),
168 context: "analysis mask",
169 });
170 }
171 let combined = match storage.analysis_mask() {
172 Some(existing) => {
173 let mut bytes = vec![0u8; n.div_ceil(8)];
174 for i in 0..n {
175 if existing.is_valid(i) && mask.is_valid(i) {
176 bytes[i / 8] |= 1 << (i % 8);
177 }
178 }
179 ValidityBitmap::from_bytes(bytes, n)?
180 }
181 None => mask,
182 };
183 let new_storage = OwnedColumnarStorage::try_new(
184 storage.schema().clone(),
185 storage.columns().to_vec(),
186 Some(combined),
187 storage.weights().map(Arc::from),
188 )?;
189 Ok(Self::new(new_storage))
190 }
191
192 pub fn with_appended_float(
198 &self,
199 name: &str,
200 values: Arc<[f64]>,
201 ) -> Result<(Self, VariableId), DataError> {
202 let n = self.row_count();
203 if values.len() != n {
204 return Err(DataError::LengthMismatch {
205 expected: n,
206 actual: values.len(),
207 context: "appended float column",
208 });
209 }
210 let storage = self.storage();
211 let mut builder = CausalSchemaBuilder::new();
212 for v in storage.schema().variables() {
213 builder
214 .add_variable(
215 Arc::clone(&v.name),
216 v.value_type.clone(),
217 v.role_hints,
218 v.unit.clone(),
219 v.category_domain,
220 v.measurement.clone(),
221 )
222 .map_err(|e| DataError::Schema(e.to_string()))?;
223 }
224 builder
225 .add_variable(
226 name,
227 ValueType::Continuous,
228 SmallRoleSet::from_hint(RoleHint::Context),
229 None,
230 None,
231 MeasurementSpec::default(),
232 )
233 .map_err(|e| DataError::Schema(e.to_string()))?;
234 let schema = builder.build().map_err(|e| DataError::Schema(e.to_string()))?;
235 let new_id = VariableId::from_raw(u32::try_from(schema.len() - 1).map_err(|_| {
236 DataError::InvalidArgument { message: "schema exceeds VariableId range".into() }
237 })?);
238 let mut cols: Vec<OwnedColumn> = storage.columns().to_vec();
239 cols.push(OwnedColumn::Float64(Float64Column::new(
240 new_id,
241 values,
242 ValidityBitmap::all_valid(n),
243 )?));
244 let storage = OwnedColumnarStorage::try_new(
245 schema,
246 cols,
247 storage.analysis_mask().cloned(),
248 storage.weights().map(Arc::from),
249 )?;
250 Ok((Self::new(storage), new_id))
251 }
252}