use std::ops::{Mul, MulAssign};
use auto_impl_ops::auto_ops;
use either::Either;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Perm {
data: Either<usize, Vec<usize>>,
}
impl Perm {
pub(crate) fn new(data: Vec<usize>) -> Self {
assert!(is_valid_perm(&data), "not a valid permutation: {:?}", data);
Self::new_unchecked(data)
}
fn new_unchecked(data: Vec<usize>) -> Self {
Self { data: Either::Right(data) }
}
pub fn from_indices<I>(images: I) -> Self
where I: IntoIterator<Item = usize> {
Self::new(images.into_iter().collect())
}
pub fn id(n: usize) -> Self {
Self { data: Either::Left(n) }
}
pub fn len(&self) -> usize {
match &self.data {
Either::Left(n) => *n,
Either::Right(v) => v.len(),
}
}
pub fn at(&self, i: usize) -> usize {
match &self.data {
Either::Left(_) => i,
Either::Right(v) => v[i],
}
}
pub fn is_id(&self) -> bool {
match &self.data {
Either::Left(_) => true,
Either::Right(v) => v.iter().enumerate().all(|(i, &x)| i == x),
}
}
pub fn inv(&self) -> Self {
match &self.data {
Either::Left(n) => Self::id(*n),
Either::Right(v) => {
let mut inv = vec![0; v.len()];
for (i, &j) in v.iter().enumerate() {
inv[j] = i;
}
Self { data: Either::Right(inv) }
}
}
}
pub fn apply_to<R>(&self, y: Vec<R>) -> Vec<R> {
assert_eq!(y.len(), self.len());
match &self.data {
Either::Left(_) => y,
Either::Right(v) => {
let n = v.len();
let mut result: Vec<R> = Vec::with_capacity(n);
let dst = result.as_mut_ptr();
for (i, x) in y.into_iter().enumerate() {
unsafe { dst.add(v[i]).write(x); }
}
unsafe { result.set_len(n); }
result
}
}
}
pub fn apply_inv_to<R>(&self, y: Vec<R>) -> Vec<R> {
assert_eq!(y.len(), self.len());
match &self.data {
Either::Left(_) => y,
Either::Right(v) => {
let mut y = std::mem::ManuallyDrop::new(y);
let src = y.as_mut_ptr();
let cap = y.capacity();
let result: Vec<R> = v.iter().map(|&k| unsafe { src.add(k).read() }).collect();
unsafe { drop(Vec::from_raw_parts(src, 0, cap)); }
result
}
}
}
pub fn shift(self, r: usize) -> Self {
match self.data {
Either::Left(n) => Self::id(n + r),
Either::Right(v) => {
let mut data: Vec<usize> = (0..r).collect();
data.extend(v.into_iter().map(|x| x + r));
Self::new(data)
}
}
}
pub fn extend(self, r: usize) -> Self {
match self.data {
Either::Left(n) => Self::id(n + r),
Either::Right(mut v) => {
let m = v.len();
v.extend(m..m + r);
Self::new(v)
}
}
}
pub fn forward_indices<I>(n: usize, prefix: I) -> Self
where I: IntoIterator<Item = usize> {
let mut data = vec![0usize; n];
let mut taken = vec![false; n];
let mut k = 0;
for i in prefix {
assert!(i < n, "index {i} out of range 0..{n}");
assert!(!taken[i], "duplicate index {i} in prefix");
data[i] = k;
taken[i] = true;
k += 1;
}
if k == 0 {
return Self::id(n);
}
let mut pos = k;
for j in 0..n {
if !taken[j] {
data[j] = pos;
pos += 1;
}
}
Self::new_unchecked(data)
}
}
#[auto_ops]
impl Mul<&Perm> for &Perm {
type Output = Perm;
fn mul(self, rhs: &Perm) -> Perm {
assert_eq!(self.len(), rhs.len(), "permutations must have the same length");
match (&self.data, &rhs.data) {
(Either::Left(_), _) => rhs.clone(),
(_, Either::Left(_)) => self.clone(),
(Either::Right(_), Either::Right(qv)) => {
Perm::new(qv.iter().map(|&i| self.at(i)).collect())
}
}
}
}
fn is_valid_perm(v: &[usize]) -> bool {
let n = v.len();
let mut seen = vec![false; n];
for &x in v {
if x >= n || seen[x] {
return false;
}
seen[x] = true;
}
true
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn id() {
let p = Perm::id(4);
assert!(p.is_id());
assert_eq!(p.len(), 4);
for i in 0..4 {
assert_eq!(p.at(i), i);
}
}
#[test]
fn new_accepts_valid() {
let p = Perm::new(vec![2, 0, 1]);
assert_eq!(p.at(0), 2);
assert_eq!(p.at(1), 0);
assert_eq!(p.at(2), 1);
}
#[test]
#[should_panic]
fn new_rejects_duplicate() {
let _ = Perm::new(vec![0, 0, 1]);
}
#[test]
#[should_panic]
fn new_rejects_out_of_range() {
let _ = Perm::new(vec![0, 1, 5]);
}
#[test]
fn is_id_true_false() {
assert!(Perm::id(3).is_id());
assert!(Perm::new(vec![0, 1, 2]).is_id());
assert!(!Perm::new(vec![1, 0]).is_id());
}
#[test]
fn accessors() {
let v = vec![2, 0, 1, 3];
let p = Perm::new(v.clone());
assert_eq!(p.len(), 4);
for (i, &expected) in v.iter().enumerate() {
assert_eq!(p.at(i), expected);
}
}
#[test]
fn inv_of_id() {
let id = Perm::id(5);
assert_eq!(id.inv(), id);
}
#[test]
fn inv_of_inv() {
let p = Perm::new(vec![2, 0, 3, 1]);
assert_eq!(p.inv().inv(), p);
}
#[test]
fn inv_roundtrip() {
let p = Perm::new(vec![2, 0, 3, 1]);
let pi = p.inv();
for i in 0..p.len() {
assert_eq!(pi.at(p.at(i)), i);
assert_eq!(p.at(pi.at(i)), i);
}
}
#[test]
fn mul_formula() {
let p = Perm::new(vec![2, 0, 3, 1]);
let q = Perm::new(vec![1, 2, 0, 3]);
let pq = &p * &q;
for i in 0..4 {
assert_eq!(pq.at(i), p.at(q.at(i)));
}
}
#[test]
fn mul_with_id() {
let p = Perm::new(vec![2, 0, 3, 1]);
let id = Perm::id(4);
assert_eq!(&p * &id, p);
assert_eq!(&id * &p, p);
}
#[test]
fn mul_id_with_id() {
let id = Perm::id(4);
assert_eq!(&id * &id, id);
}
#[test]
fn mul_by_inverse_is_id() {
let p = Perm::new(vec![2, 0, 3, 1]);
assert!((&p * &p.inv()).is_id());
assert!((&p.inv() * &p).is_id());
}
#[test]
fn mul_associative() {
let p = Perm::new(vec![2, 0, 3, 1]);
let q = Perm::new(vec![1, 3, 0, 2]);
let r = Perm::new(vec![3, 1, 2, 0]);
assert_eq!(&(&p * &q) * &r, &p * &(&q * &r));
}
#[test]
#[should_panic]
fn mul_panics_on_dim_mismatch() {
let p = Perm::id(3);
let q = Perm::id(4);
let _ = &p * &q;
}
#[test]
fn apply_to_consumes() {
let p = Perm::from_indices([1, 2, 0]);
assert_eq!(p.apply_to(vec![10, 20, 30]), vec![30, 10, 20]);
}
#[test]
fn apply_to_by_id() {
let id = Perm::id(4);
assert_eq!(id.apply_to(vec![1, 2, 3, 4]), vec![1, 2, 3, 4]);
}
#[test]
fn apply_inv_to_consumes() {
let p = Perm::from_indices([1, 2, 0]);
assert_eq!(p.apply_inv_to(vec![10, 20, 30]), vec![20, 30, 10]);
}
#[test]
fn apply_inv_to_inverts_apply_to() {
let p = Perm::from_indices([2, 0, 3, 1]);
let y = vec![10, 20, 30, 40];
let permuted = p.apply_to(y.clone());
assert_eq!(p.apply_inv_to(permuted), y);
}
#[test]
#[should_panic]
fn apply_to_panics_on_len_mismatch() {
let p = Perm::id(3);
let _ = p.apply_to(vec![1, 2]);
}
#[test]
fn shift_basic() {
let p = Perm::from_indices([1, 2, 0]);
let s = p.shift(2);
assert_eq!(s.len(), 5);
for (i, &x) in [0, 1, 3, 4, 2].iter().enumerate() {
assert_eq!(s.at(i), x);
}
}
#[test]
fn shift_zero() {
let p = Perm::from_indices([2, 0, 1]);
let s = p.clone().shift(0);
assert_eq!(s, p);
}
#[test]
fn shift_of_id() {
let s = Perm::id(3).shift(2);
assert!(s.is_id());
assert_eq!(s.len(), 5);
}
#[test]
fn extend_basic() {
let p = Perm::from_indices([1, 2, 0]);
let s = p.extend(2);
assert_eq!(s.len(), 5);
for (i, &x) in [1, 2, 0, 3, 4].iter().enumerate() {
assert_eq!(s.at(i), x);
}
}
#[test]
fn extend_zero() {
let p = Perm::from_indices([2, 0, 1]);
let s = p.clone().extend(0);
assert_eq!(s, p);
}
#[test]
fn extend_of_id() {
let s = Perm::id(3).extend(2);
assert!(s.is_id());
assert_eq!(s.len(), 5);
}
#[test]
fn forward_indices_basic() {
let p = Perm::forward_indices(5, [3, 1]);
let expected = [2, 1, 3, 0, 4];
for (i, &x) in expected.iter().enumerate() {
assert_eq!(p.at(i), x);
}
}
#[test]
fn forward_indices_empty_prefix() {
let p = Perm::forward_indices(4, std::iter::empty());
assert!(p.is_id());
}
#[test]
fn forward_indices_full_prefix() {
let p = Perm::forward_indices(4, [2, 0, 3, 1]);
assert_eq!(p.at(2), 0);
assert_eq!(p.at(0), 1);
assert_eq!(p.at(3), 2);
assert_eq!(p.at(1), 3);
}
#[test]
fn mul_all_ref_variants() {
let p = Perm::new(vec![2, 0, 3, 1]);
let q = Perm::new(vec![1, 2, 0, 3]);
let expected = &p * &q;
assert_eq!(p.clone() * q.clone(), expected);
assert_eq!(p.clone() * &q, expected);
assert_eq!(&p * q.clone(), expected);
assert_eq!(&p * &q, expected);
}
}