1use crate::aggregate::ProjectedClassEvidence;
4use crate::alphabet::validate_alphabet;
5use crate::{AggregateLedger, AggregateRule, SerialAlphabet, SeriesError};
6use sim_lib_discrete_rank::PermutationSpace;
7use sim_lib_rank::Nat;
8use std::collections::BTreeMap;
9
10#[derive(Clone, Debug, PartialEq, Eq)]
12pub struct Series<A: SerialAlphabet> {
13 alphabet: A,
14 rule: AggregateRule,
15 order: Vec<A::Symbol>,
16 ledger: AggregateLedger<A::Symbol>,
17}
18
19impl<A: SerialAlphabet> Series<A> {
20 pub fn try_new(
22 alphabet: A,
23 rule: AggregateRule,
24 order: Vec<A::Symbol>,
25 ) -> Result<Self, SeriesError> {
26 let positions = validate_alphabet(&alphabet)?;
27 let mut observed = alphabet
28 .symbols()
29 .iter()
30 .cloned()
31 .map(|symbol| (symbol, 0usize))
32 .collect::<BTreeMap<_, _>>();
33 let mut first_occurrence = vec![None; alphabet.symbols().len()];
34 let mut order_positions = Vec::with_capacity(order.len());
35 for (series_position, symbol) in order.iter().enumerate() {
36 let Some(&alphabet_position) = positions.get(symbol) else {
37 return Err(SeriesError::ForeignSymbol {
38 position: series_position,
39 alphabet_id: alphabet.id().clone(),
40 });
41 };
42 order_positions.push(alphabet_position);
43 if first_occurrence[alphabet_position].is_none() {
44 first_occurrence[alphabet_position] = Some(series_position);
45 }
46 if let Some(count) = observed.get_mut(symbol) {
47 *count += 1;
48 }
49 }
50
51 let expected = validate_rule(
52 &alphabet,
53 &rule,
54 &observed,
55 &first_occurrence,
56 &order_positions,
57 order.len(),
58 )?;
59 let omitted = alphabet
60 .symbols()
61 .iter()
62 .filter(|symbol| observed.get(*symbol) == Some(&0))
63 .cloned()
64 .collect();
65 let repeated = alphabet
66 .symbols()
67 .iter()
68 .filter(|symbol| observed.get(*symbol).is_some_and(|count| *count > 1))
69 .cloned()
70 .collect();
71 let projected = projected_evidence(&rule, &order_positions)?;
72 let ledger = AggregateLedger {
73 alphabet_id: alphabet.id().clone(),
74 rule: rule.kind(),
75 series_len: order.len(),
76 observed,
77 expected,
78 omitted,
79 repeated,
80 projected,
81 };
82 Ok(Self {
83 alphabet,
84 rule,
85 order,
86 ledger,
87 })
88 }
89
90 pub fn alphabet(&self) -> &A {
92 &self.alphabet
93 }
94
95 pub fn rule(&self) -> &AggregateRule {
97 &self.rule
98 }
99
100 pub fn order(&self) -> &[A::Symbol] {
102 &self.order
103 }
104
105 pub fn ledger(&self) -> &AggregateLedger<A::Symbol> {
107 &self.ledger
108 }
109
110 pub fn permutation_rank(&self) -> Result<Nat, SeriesError> {
115 if !self.ledger.is_exhaustive_exactly_once() {
116 return Err(SeriesError::NotPermutation(self.alphabet.id().clone()));
117 }
118 let positions = validate_alphabet(&self.alphabet)?;
119 let permutation = self
120 .order
121 .iter()
122 .map(|symbol| {
123 positions
124 .get(symbol)
125 .copied()
126 .ok_or_else(|| SeriesError::NotPermutation(self.alphabet.id().clone()))
127 })
128 .collect::<Result<Vec<_>, _>>()?;
129 Ok(PermutationSpace::try_new(self.alphabet.symbols().len())?.rank(&permutation)?)
130 }
131
132 pub fn into_parts(self) -> (A, AggregateRule, Vec<A::Symbol>) {
134 (self.alphabet, self.rule, self.order)
135 }
136}
137
138fn validate_rule<A: SerialAlphabet>(
139 alphabet: &A,
140 rule: &AggregateRule,
141 observed: &BTreeMap<A::Symbol, usize>,
142 first_occurrence: &[Option<usize>],
143 order_positions: &[usize],
144 order_len: usize,
145) -> Result<Option<BTreeMap<A::Symbol, usize>>, SeriesError> {
146 match rule {
147 AggregateRule::ExhaustiveExactlyOnce => {
148 require_length(alphabet.symbols().len(), order_len)?;
149 let expected = vec![1; alphabet.symbols().len()];
150 compare_symbol_counts(alphabet, observed, &expected)?;
151 Ok(Some(expected_map(alphabet, &expected)))
152 }
153 AggregateRule::NoRepeat => {
154 reject_repeats(order_positions, first_occurrence)?;
155 Ok(None)
156 }
157 AggregateRule::DeclaredMultiplicity(_) | AggregateRule::DeclaredOmissions(_) => {
158 let declared = rule
159 .declared()
160 .ok_or_else(|| SeriesError::NotPermutation(alphabet.id().clone()))?;
161 declared.validate_for(alphabet)?;
162 let expected = declared.expected();
163 require_length(checked_len(expected)?, order_len)?;
164 compare_symbol_counts(alphabet, observed, expected)?;
165 Ok(Some(expected_map(alphabet, expected)))
166 }
167 AggregateRule::ProjectedAggregate(_) => {
168 let projected = rule
169 .projected()
170 .ok_or_else(|| SeriesError::NotPermutation(alphabet.id().clone()))?;
171 projected.validate_for(alphabet)?;
172 require_length(projected.required_len()?, order_len)?;
173 let evidence = projected_evidence(rule, order_positions)?;
174 for class in evidence {
175 if class.observed != class.expected {
176 return Err(SeriesError::ProjectionMismatch {
177 class_id: class.id,
178 expected: class.expected,
179 found: class.observed,
180 });
181 }
182 }
183 Ok(None)
184 }
185 AggregateRule::FreeOrder => Ok(None),
186 }
187}
188
189fn reject_repeats(
190 order_positions: &[usize],
191 first_occurrence: &[Option<usize>],
192) -> Result<(), SeriesError> {
193 let mut seen = vec![false; first_occurrence.len()];
194 for (position, &alphabet_position) in order_positions.iter().enumerate() {
195 if seen[alphabet_position] {
196 return Err(SeriesError::RepeatedSymbol {
197 position,
198 first: first_occurrence[alphabet_position].unwrap_or(position),
199 });
200 }
201 seen[alphabet_position] = true;
202 }
203 Ok(())
204}
205
206fn compare_symbol_counts<A: SerialAlphabet>(
207 alphabet: &A,
208 observed: &BTreeMap<A::Symbol, usize>,
209 expected: &[usize],
210) -> Result<(), SeriesError> {
211 for (alphabet_position, symbol) in alphabet.symbols().iter().enumerate() {
212 let found = observed.get(symbol).copied().unwrap_or(0);
213 if found != expected[alphabet_position] {
214 return Err(SeriesError::MultiplicityMismatch {
215 alphabet_position,
216 expected: expected[alphabet_position],
217 found,
218 });
219 }
220 }
221 Ok(())
222}
223
224fn expected_map<A: SerialAlphabet>(alphabet: &A, expected: &[usize]) -> BTreeMap<A::Symbol, usize> {
225 alphabet
226 .symbols()
227 .iter()
228 .cloned()
229 .zip(expected.iter().copied())
230 .collect()
231}
232
233fn projected_evidence(
234 rule: &AggregateRule,
235 order_positions: &[usize],
236) -> Result<Vec<ProjectedClassEvidence>, SeriesError> {
237 let Some(projected) = rule.projected() else {
238 return Ok(Vec::new());
239 };
240 let mut observed = vec![0usize; projected.classes().len()];
241 for &position in order_positions {
242 let Some(&class) = projected.class_by_position().get(position) else {
243 return Err(SeriesError::Rule(
244 crate::AggregateRuleError::CardinalityMismatch {
245 expected: projected.class_by_position().len(),
246 found: position.saturating_add(1),
247 },
248 ));
249 };
250 observed[class] += 1;
251 }
252 Ok(projected
253 .classes()
254 .iter()
255 .zip(observed)
256 .map(|(class, observed)| ProjectedClassEvidence {
257 id: class.id().clone(),
258 expected: class.multiplicity(),
259 observed,
260 })
261 .collect())
262}
263
264fn require_length(expected: usize, found: usize) -> Result<(), SeriesError> {
265 if expected == found {
266 Ok(())
267 } else {
268 Err(SeriesError::WrongLength { expected, found })
269 }
270}
271
272fn checked_len(counts: &[usize]) -> Result<usize, SeriesError> {
273 counts.iter().try_fold(0usize, |total, &count| {
274 total.checked_add(count).ok_or(SeriesError::Rule(
275 crate::AggregateRuleError::MultiplicityOverflow,
276 ))
277 })
278}