use std::collections::BTreeMap;
use sim_lib_pitch_core::PitchClass;
use sim_lib_pitch_serial::PitchClassAlphabet;
use sim_lib_serial_core::{AggregateRule, Series, SeriesError};
use thiserror::Error;
use crate::rotate_sequence_left;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SerialArrayRow {
pub label: String,
pub order: Series<PitchClassAlphabet>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ColumnPartition {
pub id: String,
pub columns: Vec<usize>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct VerticalAggregateRequirement {
pub id: String,
pub partitions: Vec<ColumnPartition>,
pub rule: AggregateRule,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PartitionCoverageReport {
pub duplicate_columns: Vec<usize>,
pub omitted_columns: Vec<usize>,
pub complete: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct AggregatePartitionReport {
pub id: String,
pub columns: Vec<usize>,
pub values: Vec<PitchClass>,
pub duplicates: Vec<PitchClass>,
pub omissions: Vec<PitchClass>,
pub satisfied: bool,
pub error: Option<SeriesError>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct AggregateArrayReport {
pub id: String,
pub coverage: PartitionCoverageReport,
pub partitions: Vec<AggregatePartitionReport>,
pub satisfied: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SerialArray {
rows: Vec<SerialArrayRow>,
vertical_requirements: Vec<VerticalAggregateRequirement>,
column_count: usize,
}
impl SerialArray {
pub fn try_new(
rows: Vec<SerialArrayRow>,
vertical_requirements: Vec<VerticalAggregateRequirement>,
) -> Result<Self, SerialArrayError> {
let Some(first_row) = rows.first() else {
return Err(SerialArrayError::EmptyArray);
};
let column_count = first_row.order.order().len();
if column_count == 0 {
return Err(SerialArrayError::EmptyRowOrder {
row_label: first_row.label.clone(),
});
}
for row in &rows {
if row.label.trim().is_empty() {
return Err(SerialArrayError::EmptyRowLabel);
}
if row.order.order().len() != column_count {
return Err(SerialArrayError::RowLengthMismatch {
row_label: row.label.clone(),
expected: column_count,
found: row.order.order().len(),
});
}
}
for requirement in &vertical_requirements {
if requirement.id.trim().is_empty() {
return Err(SerialArrayError::EmptyRequirementId);
}
for partition in &requirement.partitions {
if partition.id.trim().is_empty() {
return Err(SerialArrayError::EmptyPartitionId {
requirement_id: requirement.id.clone(),
});
}
if partition.columns.is_empty() {
return Err(SerialArrayError::EmptyPartition {
requirement_id: requirement.id.clone(),
partition_id: partition.id.clone(),
});
}
for &column in &partition.columns {
if column >= column_count {
return Err(SerialArrayError::ColumnOutOfRange {
requirement_id: requirement.id.clone(),
partition_id: partition.id.clone(),
column,
column_count,
});
}
}
}
}
Ok(Self {
rows,
vertical_requirements,
column_count,
})
}
pub fn row_count(&self) -> usize {
self.rows.len()
}
pub fn column_count(&self) -> usize {
self.column_count
}
pub fn rows(&self) -> &[SerialArrayRow] {
&self.rows
}
pub fn vertical_requirements(&self) -> &[VerticalAggregateRequirement] {
&self.vertical_requirements
}
pub fn rotate_columns(&self, steps: usize) -> Self {
let rows = self
.rows
.iter()
.map(|row| SerialArrayRow {
label: row.label.clone(),
order: Series::try_new(
row.order.alphabet().clone(),
row.order.rule().clone(),
rotate_sequence_left(row.order.order(), steps),
)
.expect("rotating a validated row preserves its aggregate contract"),
})
.collect();
Self {
rows,
vertical_requirements: self.vertical_requirements.clone(),
column_count: self.column_count,
}
}
pub fn aggregate_reports(&self) -> Vec<AggregateArrayReport> {
self.vertical_requirements
.iter()
.map(|requirement| self.aggregate_report(requirement))
.collect()
}
pub fn all_partition_report(
&self,
block_width: usize,
) -> Result<AggregateArrayReport, SerialArrayError> {
if block_width == 0 {
return Err(SerialArrayError::ZeroBlockWidth);
}
if !self.column_count.is_multiple_of(block_width) {
return Err(SerialArrayError::BlockWidthMismatch {
block_width,
column_count: self.column_count,
});
}
let partitions = (0..self.column_count / block_width)
.map(|index| ColumnPartition {
id: format!("partition/{}", index + 1),
columns: ((index * block_width)..((index + 1) * block_width)).collect(),
})
.collect::<Vec<_>>();
Ok(self.aggregate_report(&VerticalAggregateRequirement {
id: format!("all-partition/{block_width}"),
partitions,
rule: AggregateRule::exhaustive_exactly_once(),
}))
}
fn aggregate_report(&self, requirement: &VerticalAggregateRequirement) -> AggregateArrayReport {
let coverage = coverage_report(self.column_count, &requirement.partitions);
let alphabet = PitchClassAlphabet::try_new().expect("canonical pitch-class alphabet");
let partitions = requirement
.partitions
.iter()
.map(|partition| {
let values = self.flatten_partition(&partition.columns);
let duplicates = duplicate_pitch_classes(&values);
let omissions = omitted_pitch_classes(&values);
match Series::try_new(alphabet.clone(), requirement.rule.clone(), values.clone()) {
Ok(_) => AggregatePartitionReport {
id: partition.id.clone(),
columns: partition.columns.clone(),
values,
duplicates,
omissions,
satisfied: true,
error: None,
},
Err(error) => AggregatePartitionReport {
id: partition.id.clone(),
columns: partition.columns.clone(),
values,
duplicates,
omissions,
satisfied: false,
error: Some(error),
},
}
})
.collect::<Vec<_>>();
let satisfied = coverage.complete && partitions.iter().all(|partition| partition.satisfied);
AggregateArrayReport {
id: requirement.id.clone(),
coverage,
partitions,
satisfied,
}
}
fn flatten_partition(&self, columns: &[usize]) -> Vec<PitchClass> {
let mut values = Vec::with_capacity(self.rows.len() * columns.len());
for row in &self.rows {
for &column in columns {
values.push(row.order.order()[column]);
}
}
values
}
}
#[derive(Clone, Debug, PartialEq, Eq, Error)]
pub enum SerialArrayError {
#[error("serial array must contain at least one row")]
EmptyArray,
#[error("serial array rows must carry a non-empty label")]
EmptyRowLabel,
#[error("serial array row {row_label:?} must contain at least one value")]
EmptyRowOrder {
row_label: String,
},
#[error("serial array row {row_label:?} has length {found}; expected {expected}")]
RowLengthMismatch {
row_label: String,
expected: usize,
found: usize,
},
#[error("serial array vertical requirements must carry a non-empty id")]
EmptyRequirementId,
#[error("serial array requirement {requirement_id:?} contains an empty partition id")]
EmptyPartitionId {
requirement_id: String,
},
#[error(
"serial array requirement {requirement_id:?} partition {partition_id:?} must name at least one column"
)]
EmptyPartition {
requirement_id: String,
partition_id: String,
},
#[error(
"serial array requirement {requirement_id:?} partition {partition_id:?} names column {column} outside 0..{column_count}"
)]
ColumnOutOfRange {
requirement_id: String,
partition_id: String,
column: usize,
column_count: usize,
},
#[error("all-partition block width must be at least 1")]
ZeroBlockWidth,
#[error("all-partition block width {block_width} does not divide column count {column_count}")]
BlockWidthMismatch {
block_width: usize,
column_count: usize,
},
}
fn coverage_report(column_count: usize, partitions: &[ColumnPartition]) -> PartitionCoverageReport {
let mut counts = vec![0usize; column_count];
for partition in partitions {
for &column in &partition.columns {
counts[column] += 1;
}
}
let duplicate_columns = counts
.iter()
.enumerate()
.filter_map(|(column, count)| (*count > 1).then_some(column))
.collect::<Vec<_>>();
let omitted_columns = counts
.iter()
.enumerate()
.filter_map(|(column, count)| (*count == 0).then_some(column))
.collect::<Vec<_>>();
PartitionCoverageReport {
complete: duplicate_columns.is_empty() && omitted_columns.is_empty(),
duplicate_columns,
omitted_columns,
}
}
fn duplicate_pitch_classes(values: &[PitchClass]) -> Vec<PitchClass> {
let mut counts = BTreeMap::new();
for &value in values {
*counts.entry(value).or_insert(0usize) += 1;
}
counts
.into_iter()
.filter_map(|(pitch_class, count)| (count > 1).then_some(pitch_class))
.collect()
}
fn omitted_pitch_classes(values: &[PitchClass]) -> Vec<PitchClass> {
PitchClassAlphabet::try_new()
.expect("canonical pitch-class alphabet")
.classes()
.iter()
.copied()
.filter(|pitch_class| !values.contains(pitch_class))
.collect()
}