1use crate::alphabet::validate_alphabet;
4use crate::{
5 AggregateRule, AggregateRuleKind, BlockPartition, OrdinalMap, ProjectedClassSpec,
6 RelaxedInvariant, SerialAlphabet, Series, SeriesTransformError, SymbolBijectionError,
7 TransformCertificate, TransformedSeries,
8};
9
10#[derive(Clone, Debug, PartialEq, Eq)]
12pub struct SymbolBijection<A: SerialAlphabet> {
13 source: A,
14 target: A,
15 source_to_target: OrdinalMap,
16}
17
18impl<A: SerialAlphabet> SymbolBijection<A> {
19 pub fn try_new<I>(source: A, target: A, pairs: I) -> Result<Self, SymbolBijectionError>
21 where
22 I: IntoIterator<Item = (A::Symbol, A::Symbol)>,
23 {
24 let source_positions = validate_alphabet(&source)?;
25 let target_positions = validate_alphabet(&target)?;
26 if source.symbols().len() != target.symbols().len() {
27 return Err(SymbolBijectionError::CardinalityMismatch {
28 source_cardinality: source.symbols().len(),
29 target_cardinality: target.symbols().len(),
30 });
31 }
32
33 let cardinality = source.symbols().len();
34 let mut source_to_target = vec![None; cardinality];
35 let mut target_sources = vec![None; cardinality];
36 for (source_symbol, target_symbol) in pairs {
37 let Some(&source_position) = source_positions.get(&source_symbol) else {
38 return Err(SymbolBijectionError::ForeignSourceSymbol {
39 alphabet_id: source.id().clone(),
40 });
41 };
42 let Some(&target_position) = target_positions.get(&target_symbol) else {
43 return Err(SymbolBijectionError::ForeignTargetSymbol {
44 alphabet_id: target.id().clone(),
45 });
46 };
47 if source_to_target[source_position]
48 .replace(target_position)
49 .is_some()
50 {
51 return Err(SymbolBijectionError::DuplicateSource {
52 position: source_position,
53 });
54 }
55 if target_sources[target_position]
56 .replace(source_position)
57 .is_some()
58 {
59 return Err(SymbolBijectionError::DuplicateTarget {
60 position: target_position,
61 });
62 }
63 }
64
65 let mut complete = Vec::with_capacity(cardinality);
66 for (position, target_position) in source_to_target.into_iter().enumerate() {
67 let Some(target_position) = target_position else {
68 return Err(SymbolBijectionError::MissingSource { position });
69 };
70 complete.push(target_position);
71 }
72 for (position, source_position) in target_sources.into_iter().enumerate() {
73 if source_position.is_none() {
74 return Err(SymbolBijectionError::MissingTarget { position });
75 }
76 }
77 Self::from_ordinals(source, target, complete)
78 }
79
80 pub fn cyclic(alphabet: A, steps: usize) -> Result<Self, SymbolBijectionError> {
82 validate_alphabet(&alphabet)?;
83 let cardinality = alphabet.symbols().len();
84 let shift = steps % cardinality;
85 let source_to_target = (0..cardinality)
86 .map(|source| (source + shift) % cardinality)
87 .collect();
88 Self::from_ordinals(alphabet.clone(), alphabet, source_to_target)
89 }
90
91 fn from_ordinals(
92 source: A,
93 target: A,
94 source_to_target: Vec<usize>,
95 ) -> Result<Self, SymbolBijectionError> {
96 validate_alphabet(&source)?;
97 validate_alphabet(&target)?;
98 if source.symbols().len() != target.symbols().len() {
99 return Err(SymbolBijectionError::CardinalityMismatch {
100 source_cardinality: source.symbols().len(),
101 target_cardinality: target.symbols().len(),
102 });
103 }
104 if source_to_target.len() != source.symbols().len() {
105 return Err(SymbolBijectionError::CardinalityMismatch {
106 source_cardinality: source.symbols().len(),
107 target_cardinality: source_to_target.len(),
108 });
109 }
110 Ok(Self {
111 source,
112 target,
113 source_to_target: OrdinalMap::try_new(source_to_target)?,
114 })
115 }
116
117 pub fn source(&self) -> &A {
119 &self.source
120 }
121
122 pub fn target(&self) -> &A {
124 &self.target
125 }
126
127 pub fn source_to_target(&self) -> &[usize] {
129 self.source_to_target.output_to_input()
130 }
131
132 pub fn map_symbol(&self, symbol: &A::Symbol) -> Result<A::Symbol, SymbolBijectionError> {
134 let Some(source_position) = self
135 .source
136 .symbols()
137 .iter()
138 .position(|candidate| candidate == symbol)
139 else {
140 return Err(SymbolBijectionError::ForeignSourceSymbol {
141 alphabet_id: self.source.id().clone(),
142 });
143 };
144 let Some(&target_position) = self.source_to_target().get(source_position) else {
145 return Err(SymbolBijectionError::MissingSource {
146 position: source_position,
147 });
148 };
149 self.target.symbols().get(target_position).cloned().ok_or(
150 SymbolBijectionError::MissingTarget {
151 position: target_position,
152 },
153 )
154 }
155
156 pub fn is_identity(&self) -> bool {
158 self.source == self.target && self.source_to_target.is_identity()
159 }
160
161 pub fn inverse(&self) -> Result<Self, SymbolBijectionError> {
163 Self::from_ordinals(
164 self.target.clone(),
165 self.source.clone(),
166 self.source_to_target.inverse()?.output_to_input().to_vec(),
167 )
168 }
169
170 pub fn compose(&self, next: &Self) -> Result<Self, SeriesTransformError> {
172 if self.target != next.source {
173 return Err(SeriesTransformError::CompositionAlphabetMismatch {
174 first_target: self.target.id().clone(),
175 second_source: next.source.id().clone(),
176 });
177 }
178 let mut composed = Vec::with_capacity(self.source.symbols().len());
179 for (source_position, &intermediate) in self.source_to_target().iter().enumerate() {
180 let Some(&target_position) = next.source_to_target().get(intermediate) else {
181 return Err(SymbolBijectionError::MissingSource {
182 position: source_position,
183 }
184 .into());
185 };
186 composed.push(target_position);
187 }
188 Ok(Self::from_ordinals(
189 self.source.clone(),
190 next.target.clone(),
191 composed,
192 )?)
193 }
194
195 pub fn canonical_form(&self) -> String {
197 format!(
198 "bijection/v1:{}->{}:{}",
199 self.source.id(),
200 self.target.id(),
201 self.source_to_target.canonical_form()
202 )
203 }
204}
205
206#[derive(Clone, Debug, PartialEq, Eq)]
208pub struct SeriesTransform<A: SerialAlphabet> {
209 order_map: Option<OrdinalMap>,
210 relabeling: Option<SymbolBijection<A>>,
211}
212
213impl<A: SerialAlphabet> SeriesTransform<A> {
214 pub fn identity(cardinality: usize) -> Self {
216 Self::ordinal_permutation(OrdinalMap::identity(cardinality))
217 }
218
219 pub fn retrograde(cardinality: usize) -> Self {
221 Self::ordinal_permutation(OrdinalMap::retrograde(cardinality))
222 }
223
224 pub fn rotation(cardinality: usize, steps: usize) -> Self {
226 Self::ordinal_permutation(OrdinalMap::rotation(cardinality, steps))
227 }
228
229 pub fn block_partition(partition: BlockPartition) -> Self {
231 Self::ordinal_permutation(partition.order_map().clone())
232 }
233
234 pub fn ordinal_permutation(order_map: OrdinalMap) -> Self {
236 Self {
237 order_map: Some(order_map),
238 relabeling: None,
239 }
240 }
241
242 pub fn cyclic_relabeling(alphabet: A, steps: usize) -> Result<Self, SymbolBijectionError> {
244 Ok(Self::bijection(SymbolBijection::cyclic(alphabet, steps)?))
245 }
246
247 pub fn bijection(relabeling: SymbolBijection<A>) -> Self {
249 Self {
250 order_map: None,
251 relabeling: Some(relabeling),
252 }
253 }
254
255 pub fn order_map(&self) -> Option<&OrdinalMap> {
257 self.order_map.as_ref()
258 }
259
260 pub fn relabeling(&self) -> Option<&SymbolBijection<A>> {
262 self.relabeling.as_ref()
263 }
264
265 pub fn compose(&self, next: &Self) -> Result<Self, SeriesTransformError> {
267 let order_map = match (&self.order_map, &next.order_map) {
268 (Some(first), Some(second)) => Some(first.compose(second)?),
269 (Some(first), None) => Some(first.clone()),
270 (None, Some(second)) => Some(second.clone()),
271 (None, None) => None,
272 };
273 let relabeling = match (&self.relabeling, &next.relabeling) {
274 (Some(first), Some(second)) => Some(first.compose(second)?),
275 (Some(first), None) => Some(first.clone()),
276 (None, Some(second)) => Some(second.clone()),
277 (None, None) => None,
278 };
279 Ok(Self {
280 order_map,
281 relabeling,
282 })
283 }
284
285 pub fn inverse(&self) -> Result<Self, SeriesTransformError> {
287 Ok(Self {
288 order_map: self
289 .order_map
290 .as_ref()
291 .map(OrdinalMap::inverse)
292 .transpose()?,
293 relabeling: self
294 .relabeling
295 .as_ref()
296 .map(SymbolBijection::inverse)
297 .transpose()?,
298 })
299 }
300
301 pub fn canonical_form(&self) -> String {
303 let order = self
304 .order_map
305 .as_ref()
306 .map_or_else(|| "retain".to_owned(), OrdinalMap::canonical_form);
307 let relabeling = self
308 .relabeling
309 .as_ref()
310 .map_or_else(|| "retain".to_owned(), SymbolBijection::canonical_form);
311 format!("series-transform/v1;order={order};symbols={relabeling}")
312 }
313}
314
315impl<A: SerialAlphabet> Series<A> {
316 pub fn apply(
318 &self,
319 operation: &SeriesTransform<A>,
320 ) -> Result<TransformedSeries<A>, SeriesTransformError> {
321 let order_map = operation
322 .order_map
323 .clone()
324 .unwrap_or_else(|| OrdinalMap::identity(self.order().len()));
325 let ordered = order_map.apply(self.order())?;
326
327 let (target_alphabet, target_rule, target_order) =
328 if let Some(relabeling) = &operation.relabeling {
329 if self.alphabet() != relabeling.source() {
330 return Err(SeriesTransformError::SourceAlphabetMismatch {
331 expected: relabeling.source().id().clone(),
332 found: self.alphabet().id().clone(),
333 });
334 }
335 let mapped = ordered
336 .iter()
337 .map(|symbol| relabeling.map_symbol(symbol))
338 .collect::<Result<Vec<_>, _>>()?;
339 let rule = remap_rule(self.rule(), self.alphabet(), relabeling)?;
340 (relabeling.target().clone(), rule, mapped)
341 } else {
342 (self.alphabet().clone(), self.rule().clone(), ordered)
343 };
344
345 let series = Series::try_new(target_alphabet, target_rule, target_order)?;
346 let mut relaxed_invariants = Vec::new();
347 if !order_map.is_identity() {
348 relaxed_invariants.push(RelaxedInvariant::SourceOrder);
349 }
350 if let Some(relabeling) = &operation.relabeling
351 && !relabeling.is_identity()
352 {
353 relaxed_invariants.push(RelaxedInvariant::SymbolIdentity);
354 if relabeling.source().id() != relabeling.target().id() {
355 relaxed_invariants.push(RelaxedInvariant::AlphabetIdentity);
356 }
357 }
358 relaxed_invariants.sort_unstable();
359
360 let certificate = TransformCertificate {
361 source_alphabet: self.alphabet().id().clone(),
362 target_alphabet: series.alphabet().id().clone(),
363 aggregate_preserved: true,
364 order_map,
365 inverse: Some(operation.inverse()?),
366 relaxed_invariants,
367 };
368 Ok(TransformedSeries {
369 series,
370 certificate,
371 })
372 }
373}
374
375fn remap_rule<A: SerialAlphabet>(
376 rule: &AggregateRule,
377 source: &A,
378 relabeling: &SymbolBijection<A>,
379) -> Result<AggregateRule, SeriesTransformError> {
380 let target = relabeling.target();
381 match rule.kind() {
382 AggregateRuleKind::ExhaustiveExactlyOnce => Ok(AggregateRule::exhaustive_exactly_once()),
383 AggregateRuleKind::NoRepeat => Ok(AggregateRule::no_repeat()),
384 AggregateRuleKind::FreeOrder => Ok(AggregateRule::free_order()),
385 AggregateRuleKind::DeclaredMultiplicity => {
386 let counts =
387 rule.declared_counts(source)?
388 .ok_or(SeriesTransformError::RuleKindMismatch(
389 AggregateRuleKind::DeclaredMultiplicity,
390 ))?;
391 let mapped = counts
392 .into_iter()
393 .map(|(symbol, count)| Ok((relabeling.map_symbol(&symbol)?, count)))
394 .collect::<Result<Vec<_>, SeriesTransformError>>()?;
395 Ok(AggregateRule::declared_multiplicity(target, mapped)?)
396 }
397 AggregateRuleKind::DeclaredOmissions => {
398 let counts =
399 rule.declared_counts(source)?
400 .ok_or(SeriesTransformError::RuleKindMismatch(
401 AggregateRuleKind::DeclaredOmissions,
402 ))?;
403 let omitted = counts
404 .into_iter()
405 .filter(|(_, count)| *count == 0)
406 .map(|(symbol, _)| relabeling.map_symbol(&symbol))
407 .collect::<Result<Vec<_>, _>>()?;
408 Ok(AggregateRule::declared_omissions(target, omitted)?)
409 }
410 AggregateRuleKind::ProjectedAggregate => {
411 let classes =
412 rule.projected_classes(source)?
413 .ok_or(SeriesTransformError::RuleKindMismatch(
414 AggregateRuleKind::ProjectedAggregate,
415 ))?;
416 let mapped = classes
417 .into_iter()
418 .map(|class| {
419 let symbols = class
420 .symbols
421 .iter()
422 .map(|symbol| relabeling.map_symbol(symbol))
423 .collect::<Result<Vec<_>, _>>()?;
424 Ok(ProjectedClassSpec::new(
425 class.id,
426 symbols,
427 class.multiplicity,
428 ))
429 })
430 .collect::<Result<Vec<_>, SymbolBijectionError>>()?;
431 Ok(AggregateRule::projected_aggregate(target, mapped)?)
432 }
433 }
434}