1use std::collections::BTreeSet;
2
3use serde::{Deserialize, Serialize};
4
5use crate::error::{DataError, Result};
6use crate::handle::CoordinatorFeatureBlock;
7use crate::ids::{ObservationId, RepresentationId, SampleId};
8
9#[derive(Clone, Copy, Debug, Default, Eq, PartialEq, Serialize, Deserialize)]
10#[serde(rename_all = "snake_case")]
11pub enum CollationPadding {
12 #[default]
13 None,
14 Right,
15 Left,
16 Center,
17}
18
19#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
20pub struct CollationPolicy {
21 #[serde(default)]
22 pub padding: CollationPadding,
23 #[serde(default)]
24 pub truncate: bool,
25 #[serde(default)]
26 pub batch_container: Option<String>,
27 #[serde(default = "default_true")]
28 pub emit_mask: bool,
29 #[serde(default)]
30 pub max_length: Option<usize>,
31 #[serde(default)]
32 pub pad_value: f64,
33}
34
35impl Default for CollationPolicy {
36 fn default() -> Self {
37 Self {
38 padding: CollationPadding::None,
39 truncate: false,
40 batch_container: None,
41 emit_mask: true,
42 max_length: None,
43 pad_value: 0.0,
44 }
45 }
46}
47
48fn default_true() -> bool {
49 true
50}
51
52#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
53pub struct NumericCollationInputBlock {
54 pub block_id: String,
55 pub representation_id: RepresentationId,
56 pub observation_ids: Vec<ObservationId>,
57 pub sample_ids: Vec<SampleId>,
58 pub rows: Vec<Vec<Option<f64>>>,
59 #[serde(default)]
60 pub feature_names: Option<Vec<String>>,
61}
62
63#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
64pub struct NumericTensorBlock {
65 pub block_id: String,
66 pub representation_id: RepresentationId,
67 pub batch_container: String,
68 pub observation_ids: Vec<ObservationId>,
69 pub sample_ids: Vec<SampleId>,
70 pub shape: Vec<usize>,
71 pub values: Vec<f64>,
72 #[serde(default)]
73 pub presence_mask: Option<Vec<bool>>,
74 #[serde(default)]
75 pub validity_mask: Option<Vec<bool>>,
76 #[serde(default)]
77 pub feature_names: Option<Vec<String>>,
78}
79
80pub fn numeric_input_from_feature_block(
81 block: &CoordinatorFeatureBlock,
82) -> Result<NumericCollationInputBlock> {
83 validate_feature_block_shape(block)?;
84 let rows = block
85 .values
86 .iter()
87 .enumerate()
88 .map(|(row_idx, row)| {
89 row.iter()
90 .enumerate()
91 .map(|(feature_idx, value)| match value {
92 serde_json::Value::Null => Ok(None),
93 serde_json::Value::Number(number) => number.as_f64().map(Some).ok_or_else(|| {
94 DataError::Validation(format!(
95 "feature block `{}` row `{}` feature `{}` contains a non-f64 numeric value",
96 block.feature_set_id,
97 block.observation_ids[row_idx],
98 block.feature_names[feature_idx]
99 ))
100 }),
101 _ => Err(DataError::Validation(format!(
102 "feature block `{}` row `{}` feature `{}` must be numeric or null for collation",
103 block.feature_set_id,
104 block.observation_ids[row_idx],
105 block.feature_names[feature_idx]
106 ))),
107 })
108 .collect::<Result<Vec<_>>>()
109 })
110 .collect::<Result<Vec<_>>>()?;
111 Ok(NumericCollationInputBlock {
112 block_id: block.feature_set_id.clone(),
113 representation_id: block.representation_id.clone(),
114 observation_ids: block.observation_ids.clone(),
115 sample_ids: block.sample_ids.clone(),
116 rows,
117 feature_names: Some(block.feature_names.clone()),
118 })
119}
120
121pub fn collate_feature_block(
122 block: &CoordinatorFeatureBlock,
123 policy: &CollationPolicy,
124) -> Result<NumericTensorBlock> {
125 let input = numeric_input_from_feature_block(block)?;
126 collate_numeric_block(&input, policy)
127}
128
129pub fn collate_numeric_block(
130 block: &NumericCollationInputBlock,
131 policy: &CollationPolicy,
132) -> Result<NumericTensorBlock> {
133 validate_numeric_input_block(block)?;
134 validate_collation_policy(policy)?;
135 let target_len = target_length(block, policy)?;
136 validate_feature_names_for_collation(block, policy, target_len)?;
137
138 let batch = block.rows.len();
139 let cells = batch.checked_mul(target_len).ok_or_else(|| {
140 DataError::Validation("collation dimensions exceed addressable capacity".into())
141 })?;
142 let mut values = reserve_collation::<f64>(cells)?;
143 let mut presence = reserve_collation::<bool>(cells)?;
144 let mut validity = reserve_collation::<bool>(cells)?;
145 let mut has_invalid = false;
146 for row in &block.rows {
147 let projected = project_row(row, target_len, policy)?;
148 values.extend(
149 projected
150 .values
151 .iter()
152 .map(|value| value.unwrap_or(policy.pad_value)),
153 );
154 presence.extend(projected.presence.iter().copied());
155 for (value, present) in projected.values.iter().zip(projected.presence.iter()) {
156 let valid = *present && value.is_some();
157 has_invalid |= !valid;
158 validity.push(valid);
159 }
160 }
161
162 Ok(NumericTensorBlock {
163 block_id: block.block_id.clone(),
164 representation_id: block.representation_id.clone(),
165 batch_container: policy
166 .batch_container
167 .clone()
168 .unwrap_or_else(|| "ndarray".to_string()),
169 observation_ids: block.observation_ids.clone(),
170 sample_ids: block.sample_ids.clone(),
171 shape: vec![batch, target_len],
172 values,
173 presence_mask: policy.emit_mask.then_some(presence),
174 validity_mask: has_invalid.then_some(validity),
175 feature_names: projected_feature_names(block.feature_names.as_deref(), target_len, policy)?,
176 })
177}
178
179struct ProjectedRow {
180 values: Vec<Option<f64>>,
181 presence: Vec<bool>,
182}
183
184fn reserve_collation<T>(len: usize) -> Result<Vec<T>> {
185 let mut values = Vec::new();
186 values
187 .try_reserve_exact(len)
188 .map_err(|error| DataError::Validation(format!("collation allocation failed: {error}")))?;
189 Ok(values)
190}
191
192fn validate_collation_policy(policy: &CollationPolicy) -> Result<()> {
193 if policy.max_length == Some(0) {
194 return Err(DataError::Validation(
195 "collation max_length must be greater than zero".to_string(),
196 ));
197 }
198 if let Some(container) = &policy.batch_container {
199 if container.trim().is_empty() {
200 return Err(DataError::Validation(
201 "collation batch_container must not be empty".to_string(),
202 ));
203 }
204 }
205 if !policy.pad_value.is_finite() {
206 return Err(DataError::Validation(
207 "collation pad_value must be finite".to_string(),
208 ));
209 }
210 Ok(())
211}
212
213fn validate_numeric_input_block(block: &NumericCollationInputBlock) -> Result<()> {
214 if block.block_id.trim().is_empty() {
215 return Err(DataError::Validation(
216 "collation input block_id is empty".to_string(),
217 ));
218 }
219 if block.observation_ids.is_empty() {
220 return Err(DataError::Validation(format!(
221 "collation input block `{}` contains no rows",
222 block.block_id
223 )));
224 }
225 if block.observation_ids.len() != block.sample_ids.len()
226 || block.sample_ids.len() != block.rows.len()
227 {
228 return Err(DataError::Validation(format!(
229 "collation input block `{}` row identity/value lengths differ",
230 block.block_id
231 )));
232 }
233 let mut observations = BTreeSet::new();
234 for (idx, row) in block.rows.iter().enumerate() {
235 if !observations.insert(&block.observation_ids[idx]) {
236 return Err(DataError::Validation(format!(
237 "collation input block `{}` contains duplicate observation `{}`",
238 block.block_id, block.observation_ids[idx]
239 )));
240 }
241 if row.is_empty() {
242 return Err(DataError::Validation(format!(
243 "collation input block `{}` row `{}` is empty",
244 block.block_id, block.observation_ids[idx]
245 )));
246 }
247 for value in row.iter().flatten() {
248 if !value.is_finite() {
249 return Err(DataError::Validation(format!(
250 "collation input block `{}` row `{}` contains a non-finite value",
251 block.block_id, block.observation_ids[idx]
252 )));
253 }
254 }
255 }
256 if let Some(feature_names) = &block.feature_names {
257 if feature_names.is_empty() {
258 return Err(DataError::Validation(format!(
259 "collation input block `{}` has empty feature_names",
260 block.block_id
261 )));
262 }
263 if feature_names.iter().any(|name| name.trim().is_empty()) {
264 return Err(DataError::Validation(format!(
265 "collation input block `{}` has an empty feature name",
266 block.block_id
267 )));
268 }
269 let mut names = BTreeSet::new();
270 for name in feature_names {
271 if !names.insert(name) {
272 return Err(DataError::Validation(format!(
273 "collation input block `{}` has duplicate feature `{name}`",
274 block.block_id
275 )));
276 }
277 }
278 }
279 Ok(())
280}
281
282fn validate_feature_block_shape(block: &CoordinatorFeatureBlock) -> Result<()> {
283 if block.feature_set_id.trim().is_empty() {
284 return Err(DataError::Validation(
285 "feature block feature_set_id is empty".to_string(),
286 ));
287 }
288 if block.feature_names.is_empty() {
289 return Err(DataError::Validation(format!(
290 "feature block `{}` contains no features",
291 block.feature_set_id
292 )));
293 }
294 if block.observation_ids.len() != block.sample_ids.len()
295 || block.sample_ids.len() != block.values.len()
296 {
297 return Err(DataError::Validation(format!(
298 "feature block `{}` row identity/value lengths differ",
299 block.feature_set_id
300 )));
301 }
302 for (idx, row) in block.values.iter().enumerate() {
303 if row.len() != block.feature_names.len() {
304 return Err(DataError::Validation(format!(
305 "feature block `{}` row `{}` has {} values for {} features",
306 block.feature_set_id,
307 block.observation_ids[idx],
308 row.len(),
309 block.feature_names.len()
310 )));
311 }
312 }
313 Ok(())
314}
315
316fn validate_feature_names_for_collation(
317 block: &NumericCollationInputBlock,
318 policy: &CollationPolicy,
319 target_len: usize,
320) -> Result<()> {
321 let Some(feature_names) = &block.feature_names else {
322 return Ok(());
323 };
324 for row in &block.rows {
325 if row.len() != feature_names.len() {
326 return Err(DataError::Validation(format!(
327 "collation input block `{}` named features require rectangular rows",
328 block.block_id
329 )));
330 }
331 }
332 if target_len > feature_names.len() {
333 return Err(DataError::Validation(format!(
334 "collation input block `{}` cannot pad named feature rows",
335 block.block_id
336 )));
337 }
338 if target_len < feature_names.len() && !policy.truncate {
339 return Err(DataError::Validation(format!(
340 "collation input block `{}` feature names require truncation for max_length",
341 block.block_id
342 )));
343 }
344 Ok(())
345}
346
347fn target_length(block: &NumericCollationInputBlock, policy: &CollationPolicy) -> Result<usize> {
348 let max_observed = block.rows.iter().map(Vec::len).max().unwrap_or(0);
349 let target = policy.max_length.unwrap_or(max_observed);
350 if target == 0 {
351 return Err(DataError::Validation(format!(
352 "collation input block `{}` produced an empty target length",
353 block.block_id
354 )));
355 }
356 if !policy.truncate && block.rows.iter().any(|row| row.len() > target) {
357 return Err(DataError::Validation(format!(
358 "collation input block `{}` has rows longer than max_length without truncate",
359 block.block_id
360 )));
361 }
362 if policy.padding == CollationPadding::None && block.rows.iter().any(|row| row.len() < target) {
363 return Err(DataError::Validation(format!(
364 "collation input block `{}` has ragged rows but padding is none",
365 block.block_id
366 )));
367 }
368 Ok(target)
369}
370
371fn project_row(
372 row: &[Option<f64>],
373 target_len: usize,
374 policy: &CollationPolicy,
375) -> Result<ProjectedRow> {
376 let truncated = if row.len() > target_len {
377 if !policy.truncate {
378 return Err(DataError::Validation(
379 "collation row is longer than target length without truncate".to_string(),
380 ));
381 }
382 let start = truncate_start(row.len(), target_len, policy.padding);
383 row[start..start + target_len].to_vec()
384 } else {
385 row.to_vec()
386 };
387
388 if truncated.len() == target_len {
389 return Ok(ProjectedRow {
390 presence: vec![true; target_len],
391 values: truncated,
392 });
393 }
394 if policy.padding == CollationPadding::None {
395 return Err(DataError::Validation(
396 "collation row is shorter than target length but padding is none".to_string(),
397 ));
398 }
399
400 let missing = target_len - truncated.len();
401 let (left_pad, right_pad) = match policy.padding {
402 CollationPadding::None => unreachable!("padding none handled above"),
403 CollationPadding::Right => (0, missing),
404 CollationPadding::Left => (missing, 0),
405 CollationPadding::Center => (missing / 2, missing - (missing / 2)),
406 };
407 let mut values = reserve_collation::<Option<f64>>(target_len)?;
408 let mut presence = reserve_collation::<bool>(target_len)?;
409 values.extend(std::iter::repeat_n(None, left_pad));
410 presence.extend(std::iter::repeat_n(false, left_pad));
411 values.extend(truncated);
412 presence.extend(std::iter::repeat_n(true, target_len - left_pad - right_pad));
413 values.extend(std::iter::repeat_n(None, right_pad));
414 presence.extend(std::iter::repeat_n(false, right_pad));
415 Ok(ProjectedRow { values, presence })
416}
417
418fn truncate_start(row_len: usize, target_len: usize, padding: CollationPadding) -> usize {
419 match padding {
420 CollationPadding::Left => row_len - target_len,
421 CollationPadding::Center => (row_len - target_len) / 2,
422 CollationPadding::None | CollationPadding::Right => 0,
423 }
424}
425
426fn projected_feature_names(
427 feature_names: Option<&[String]>,
428 target_len: usize,
429 policy: &CollationPolicy,
430) -> Result<Option<Vec<String>>> {
431 let Some(feature_names) = feature_names else {
432 return Ok(None);
433 };
434 if target_len == feature_names.len() {
435 return Ok(Some(feature_names.to_vec()));
436 }
437 if target_len > feature_names.len() {
438 return Err(DataError::Validation(
439 "named feature collation cannot add padded feature names".to_string(),
440 ));
441 }
442 let start = truncate_start(feature_names.len(), target_len, policy.padding);
443 Ok(Some(feature_names[start..start + target_len].to_vec()))
444}
445
446#[cfg(test)]
447mod tests {
448 use super::*;
449 use crate::ids::RepresentationId;
450 use serde_json::json;
451
452 fn obs(value: &str) -> ObservationId {
453 ObservationId::new(value).unwrap()
454 }
455
456 fn sample(value: &str) -> SampleId {
457 SampleId::new(value).unwrap()
458 }
459
460 fn feature_block() -> CoordinatorFeatureBlock {
461 CoordinatorFeatureBlock {
462 feature_set_id: "x".to_string(),
463 representation_id: RepresentationId::new("tabular_numeric").unwrap(),
464 feature_names: vec!["f0".to_string(), "f1".to_string()],
465 observation_ids: vec![obs("obs.S001"), obs("obs.S002")],
466 sample_ids: vec![sample("S001"), sample("S002")],
467 values: vec![vec![json!(1.0), json!(2.0)], vec![json!(3.0), json!(4.0)]],
468 }
469 }
470
471 #[test]
472 fn collates_rectangular_feature_block_to_row_major_tensor() {
473 let tensor = collate_feature_block(
474 &feature_block(),
475 &CollationPolicy {
476 emit_mask: false,
477 ..Default::default()
478 },
479 )
480 .unwrap();
481
482 assert_eq!(tensor.shape, vec![2, 2]);
483 assert_eq!(tensor.values, vec![1.0, 2.0, 3.0, 4.0]);
484 assert_eq!(tensor.presence_mask, None);
485 assert_eq!(tensor.validity_mask, None);
486 assert_eq!(
487 tensor.feature_names,
488 Some(vec!["f0".to_string(), "f1".to_string()])
489 );
490 }
491
492 #[test]
493 fn right_padding_emits_presence_and_validity_masks() {
494 let block = NumericCollationInputBlock {
495 block_id: "seq".to_string(),
496 representation_id: RepresentationId::new("sequence_tensor").unwrap(),
497 observation_ids: vec![obs("obs.S001"), obs("obs.S002")],
498 sample_ids: vec![sample("S001"), sample("S002")],
499 rows: vec![vec![Some(1.0), Some(2.0)], vec![Some(3.0), None]],
500 feature_names: None,
501 };
502 let tensor = collate_numeric_block(
503 &block,
504 &CollationPolicy {
505 padding: CollationPadding::Right,
506 max_length: Some(3),
507 pad_value: -1.0,
508 ..Default::default()
509 },
510 )
511 .unwrap();
512
513 assert_eq!(tensor.shape, vec![2, 3]);
514 assert_eq!(tensor.values, vec![1.0, 2.0, -1.0, 3.0, -1.0, -1.0]);
515 assert_eq!(
516 tensor.presence_mask,
517 Some(vec![true, true, false, true, true, false])
518 );
519 assert_eq!(
520 tensor.validity_mask,
521 Some(vec![true, true, false, true, false, false])
522 );
523 }
524
525 #[test]
526 fn no_padding_refuses_ragged_rows() {
527 let block = NumericCollationInputBlock {
528 block_id: "seq".to_string(),
529 representation_id: RepresentationId::new("sequence_tensor").unwrap(),
530 observation_ids: vec![obs("obs.S001"), obs("obs.S002")],
531 sample_ids: vec![sample("S001"), sample("S002")],
532 rows: vec![vec![Some(1.0), Some(2.0)], vec![Some(3.0)]],
533 feature_names: None,
534 };
535
536 let err = collate_numeric_block(&block, &CollationPolicy::default()).unwrap_err();
537
538 assert!(err.to_string().contains("ragged rows"));
539 }
540
541 #[test]
542 fn left_truncation_keeps_suffix_and_projects_feature_names() {
543 let tensor = collate_feature_block(
544 &feature_block(),
545 &CollationPolicy {
546 padding: CollationPadding::Left,
547 truncate: true,
548 max_length: Some(1),
549 ..Default::default()
550 },
551 )
552 .unwrap();
553
554 assert_eq!(tensor.shape, vec![2, 1]);
555 assert_eq!(tensor.values, vec![2.0, 4.0]);
556 assert_eq!(tensor.feature_names, Some(vec!["f1".to_string()]));
557 }
558
559 #[test]
560 fn collation_refuses_non_numeric_feature_values() {
561 let mut block = feature_block();
562 block.values[0][0] = json!("bad");
563
564 let err = collate_feature_block(&block, &CollationPolicy::default()).unwrap_err();
565
566 assert!(err.to_string().contains("must be numeric or null"));
567 }
568}