use vers_vecs::BitVec;
use crate::error::Error;
use crate::{text::Text, Character};
pub fn count_chars<C, T>(text: &Text<C, T>) -> Vec<usize>
where
C: Character,
T: AsRef<[C]>,
{
let max_character = text.max_character();
let text = text.text();
let mut occs = vec![0; max_character.into_usize() + 1];
for &c in text.iter() {
occs[c.into_usize()] += 1;
}
occs
}
pub fn get_bucket_start_pos(occs: &[usize]) -> Vec<usize> {
let mut sum = 0;
let mut buckets = vec![0; occs.len()];
for (&occ, b) in occs.iter().zip(buckets.iter_mut()) {
*b = sum;
sum += occ
}
buckets
}
pub fn get_bucket_end_pos(occs: &[usize]) -> Vec<usize> {
let mut sum = 0;
let mut buckets = vec![0; occs.len()];
for (&occ, b) in occs.iter().zip(buckets.iter_mut()) {
sum += occ;
*b = sum;
}
buckets
}
fn get_types<C, T>(text: &Text<C, T>) -> (BitVec, Vec<usize>)
where
C: Character,
T: AsRef<[C]>,
{
let text = text.text();
let n = text.len();
let mut types = BitVec::from_zeros(n);
types.set(n - 1, 1).unwrap();
if n == 1 {
return (types, vec![]);
}
let mut lms = vec![n - 1];
let mut prev_is_s_type = false;
for i in (0..(n - 1)).rev() {
let is_s_type = text[i].into_u64() < text[i + 1].into_u64()
|| (text[i].into_u64() == text[i + 1].into_u64() && prev_is_s_type);
if is_s_type {
types.set(i, 1).unwrap();
} else if prev_is_s_type {
lms.push(i + 1);
}
prev_is_s_type = is_s_type;
}
(types, lms)
}
fn is_lms(types: &BitVec, i: usize) -> bool {
i > 0 && i < usize::MAX && types.is_bit_set(i).unwrap() && !types.is_bit_set(i - 1).unwrap()
}
fn induced_sort<C, T>(text: &Text<C, T>, types: &BitVec, occs: &[usize], sa: &mut [usize])
where
C: Character,
T: AsRef<[C]>,
{
let text = text.text();
let n = text.len();
let mut bucket_start_pos = get_bucket_start_pos(occs);
for i in 0..n {
let j = sa[i];
if 0 < j && j < usize::MAX && !types.is_bit_set(j - 1).unwrap() {
let c = text[j - 1].into_usize();
let p = bucket_start_pos[c];
sa[p] = j - 1;
bucket_start_pos[c] += 1;
}
}
let mut bucket_end_pos = get_bucket_end_pos(occs);
for i in (0..n).rev() {
let j = sa[i];
if j != 0 && j != usize::MAX && types.is_bit_set(j - 1).unwrap() {
let c = text[j - 1].into_usize();
let p = bucket_end_pos[c] - 1;
sa[p] = j - 1;
bucket_end_pos[c] -= 1;
}
}
}
pub fn build_suffix_array<C, T>(text: &Text<C, T>) -> Result<Vec<usize>, Error>
where
C: Character,
T: AsRef<[C]>,
{
let n = text.text().len();
if n == 0 {
return Ok(vec![]);
}
if n == 1 {
return Ok(vec![0]);
}
let first_char = text.text()[0];
if first_char.into_u64() == 0 {
return Err(Error::InvalidText(
"the given text must not start with zero character",
));
}
let last_non_zero_char = text.text().iter().rposition(|&c| c.into_u64() != 0);
if last_non_zero_char != Some(text.text().len() - 2) {
return Err(Error::InvalidText(
"the given text must end with exactly one zero character",
));
}
let mut sa = vec![usize::MAX; n];
sais_sub(text, &mut sa);
Ok(sa)
}
#[allow(clippy::cognitive_complexity)]
fn sais_sub<C, T>(text: &Text<C, T>, sa: &mut [usize])
where
C: Character,
T: AsRef<[C]>,
{
let n = text.text().len();
let (types, lms) = get_types(text);
let lms_len = lms.len();
let occs = count_chars(text);
let mut bucket_end_pos = get_bucket_end_pos(&occs);
for &i in lms.iter().rev() {
let c = text.text()[i].into_usize();
let k = bucket_end_pos[c] - 1;
sa[k] = i;
bucket_end_pos[c] = k;
}
induced_sort(text, &types, &occs, sa);
let mut k = 0;
for i in 0..n {
let p = sa[i];
if is_lms(&types, p) {
sa[k] = p;
k += 1;
if k == lms_len {
break;
}
}
}
let mut name = 1;
{
let (sa_lms, names) = sa.split_at_mut(lms_len);
for n in names.iter_mut() {
*n = usize::MAX;
}
names[sa_lms[0] / 2] = 0; if lms_len <= 1 {
debug_assert!(lms_len != 0);
} else {
names[sa_lms[1] / 2] = 1; for i in 2..lms_len {
let p = sa_lms[i - 1];
let q = sa_lms[i];
let mut d = 1;
let mut same = text.text()[p].into_u64() == text.text()[q].into_u64()
&& types.is_bit_set(p) == types.is_bit_set(q);
while same {
if text.text()[p + d].into_u64() != text.text()[q + d].into_u64()
|| types.is_bit_set(p + d) != types.is_bit_set(q + d)
{
same = false;
break;
} else if is_lms(&types, p + d) && is_lms(&types, q + d) {
break;
}
d += 1;
}
if !same {
name += 1;
}
names[q / 2] = name;
}
}
for s in sa_lms.iter_mut() {
*s = usize::MAX;
}
}
let mut i = sa.len() - 1;
let mut j = 0;
while j < lms_len {
if sa[i] < usize::MAX {
sa[sa.len() - 1 - j] = sa[i];
j += 1;
}
i -= 1;
}
{
let (sa1, s1) = sa.split_at_mut(sa.len() - lms_len);
if name < lms_len {
sais_sub(&Text::with_max_character(&s1, name), sa1);
} else {
for (i, &s) in s1.iter().enumerate() {
sa1[s] = i;
}
}
let p1 = s1;
for (j, i) in lms.into_iter().rev().enumerate() {
p1[j] = i;
}
for i in 0..lms_len {
sa1[i] = p1[sa1[i]];
}
}
for i in &mut sa[lms_len..] {
*i = usize::MAX;
}
let mut bucket_end_pos = get_bucket_end_pos(&occs);
for i in (0..lms_len).rev() {
let j = sa[i];
sa[i] = usize::MAX;
let c = if j == n {
0
} else {
text.text()[j].into_usize()
};
let k = bucket_end_pos[c] - 1;
sa[k] = j;
bucket_end_pos[c] = k;
}
induced_sort(text, &types, &occs, sa);
}
#[cfg(test)]
mod tests {
use super::*;
use num_traits::Zero;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
fn marks_to_lms(s: &str) -> Vec<usize> {
s.as_bytes()
.iter()
.enumerate()
.filter(|(_, &c)| c == b'*')
.map(|(i, _)| i)
.rev()
.collect::<Vec<_>>()
}
#[test]
fn test_get_types() {
let text = "mmiissiissiippii\0";
let n = text.len();
let types_expected = "LLSSLLSSLLSSLLLLS";
let lms_expected = marks_to_lms(" * * * *");
let (types, lms) = get_types(&Text::new(text));
let types_actual = (0..n)
.map(|i| {
if types.is_bit_set(i).unwrap() {
'S'
} else {
'L'
}
})
.collect::<String>();
assert_eq!(types_expected, types_actual);
assert_eq!(lms_expected, lms);
}
#[test]
fn test_get_types_zeros() {
let text = "m\0\0a\0";
let n = text.len();
let types_expected = "LSSLS".to_string();
let lms_expected = marks_to_lms(" * *");
let (types, lms) = get_types(&Text::new(&text));
let types_actual = (0..n)
.map(|i| {
if types.is_bit_set(i).unwrap() {
'S'
} else {
'L'
}
})
.collect::<String>();
assert_eq!(types_expected, types_actual);
assert_eq!(lms_expected, lms);
}
#[test]
fn test_get_bucket_start_pos() {
let text = "mmiissiissiippii\0";
let occs = count_chars(&Text::new(text));
let bucket_start_pos = get_bucket_start_pos(&occs);
let sa_expected = vec![(b'\0', 0), (b'i', 1), (b'm', 9), (b'p', 11), (b's', 13)];
for (c, expected) in sa_expected {
let actual = bucket_start_pos[c as usize];
assert_eq!(
actual, expected,
"bucket_start_pos['{}'] should be {} but {}",
c as char, expected, actual
);
}
}
#[test]
fn test_get_bucket_end_pos() {
let text = "mmiissiissiippii\0";
let sa_expected = vec![(b'\0', 1), (b'i', 9), (b'm', 11), (b'p', 13), (b's', 17)];
let occs = count_chars(&Text::new(text));
let bucket_end_pos = get_bucket_end_pos(&occs);
for (c, expected) in sa_expected {
let actual = bucket_end_pos[c as usize];
assert_eq!(
actual, expected,
"bucket_end_pos['{}'] should be {} but {}",
c as char, expected, actual
);
}
}
#[test]
fn test_error_no_trailing_zero() {
let text = "nozero".to_string().into_bytes();
assert!(matches!(
build_suffix_array(&Text::new(text)),
Err(Error::InvalidText(_))
));
}
#[test]
fn test_error_too_many_trailing_zero() {
let text = "toomanyzeros\0\0".to_string().into_bytes();
assert!(matches!(
build_suffix_array(&Text::new(text)),
Err(Error::InvalidText(_))
));
}
#[test]
fn test_error_starting_with_zero() {
let text = b"\0starting_with_zero\0".to_vec();
assert!(matches!(
build_suffix_array(&Text::new(text)),
Err(Error::InvalidText(_))
));
}
#[test]
fn test_length_1() {
let text = &[0u8];
let sa_actual = build_suffix_array(&Text::new(text)).unwrap();
let sa_expected = build_expected_suffix_array(text);
assert_eq!(sa_actual, sa_expected);
}
#[test]
fn test_length_2() {
let text = &[3u8, 0];
let sa_actual = build_suffix_array(&Text::new(text)).unwrap();
let sa_expected = build_expected_suffix_array(text);
assert_eq!(sa_actual, sa_expected);
}
#[test]
fn test_length_4() {
let text = &[3u8, 2, 1, 0];
let sa_actual = build_suffix_array(&Text::new(text)).unwrap();
let sa_expected = build_expected_suffix_array(text);
assert_eq!(sa_actual, sa_expected);
}
#[test]
fn test_nulls() {
let text = b"mm\0ii\0s\0sii\0ssii\0ppii\0".to_vec();
let sa_actual = build_suffix_array(&Text::new(&text)).unwrap();
let sa_expected = build_expected_suffix_array(text);
assert_eq!(sa_actual, sa_expected);
}
#[test]
fn test_small() {
let mut text = "mmiissiissiippii".to_string().into_bytes();
text.push(0);
let sa_actual = build_suffix_array(&Text::new(&text)).unwrap();
let sa_expected = build_expected_suffix_array(&text);
assert_eq!(sa_actual, sa_expected, "text: {:?}", text);
}
#[test]
fn test_rand_alphabets() {
let len = 1000;
let mut rng: StdRng = SeedableRng::from_seed([0; 32]);
for _ in 0..1000 {
let text = build_text(|| rng.gen::<u8>() % (b'z' - b'a') + b'a', len);
let sa_actual = build_suffix_array(&Text::new(&text)).unwrap();
let sa_expected = build_expected_suffix_array(&text);
assert_eq!(sa_actual, sa_expected, "text: {:?}", text);
}
}
#[test]
fn test_rand_binary_alphabets() {
let len = 1000;
let prob = 1.0 / 4.0;
let mut rng: StdRng = SeedableRng::from_seed([0; 32]);
for _ in 0..1000 {
let text = build_text(|| if rng.gen_bool(prob) { b'a' } else { b'b' }, len);
let sa_actual = build_suffix_array(&Text::new(&text)).unwrap();
let sa_expected = build_expected_suffix_array(&text);
assert_eq!(sa_actual, sa_expected, "text: {:?}", text);
}
}
#[test]
fn test_rand_binary_zero_one() {
let len = 1000;
let mut rng: StdRng = SeedableRng::from_seed([0; 32]);
for _ in 0..1000 {
let text = build_text(|| rng.gen::<u8>() % 2, len);
let sa_actual = build_suffix_array(&Text::new(&text)).unwrap();
let sa_expected = build_expected_suffix_array(&text);
assert_eq!(sa_actual, sa_expected, "text: {:?}", text);
}
}
#[test]
fn test_rand_bytes() {
let len = 1000;
let mut rng: StdRng = SeedableRng::from_seed([0; 32]);
for _ in 0..1000 {
let text = build_text(|| rng.gen::<u8>(), len);
let sa_actual = build_suffix_array(&Text::new(&text)).unwrap();
let sa_expected = build_expected_suffix_array(&text);
assert_eq!(sa_actual, sa_expected, "text: {:?}", text);
}
}
fn build_text<T: Zero + Clone, F: FnMut() -> T>(mut gen: F, len: usize) -> Vec<T> {
let mut text = vec![T::zero(); len];
let mut prev_zero = true;
for t in text.iter_mut().take(len - 1) {
let mut c = gen();
if prev_zero {
while c.is_zero() {
c = gen();
}
}
prev_zero = c.is_zero();
*t = c;
}
while text[len - 2].is_zero() {
text[len - 2] = gen();
}
text
}
fn build_expected_suffix_array<C, T>(text: T) -> Vec<usize>
where
C: Character + Ord,
T: AsRef<[C]>,
{
let text = text.as_ref();
let n = text.len();
let suffixes = (0..n).map(|i| &text[i..n]).collect::<Vec<_>>();
let mut sa = (0..(suffixes.len())).collect::<Vec<_>>();
sa.sort_by_key(|i| &suffixes[*i]);
sa
}
}