use crate::{
data::mappings::THREE_BIT_MAPPING,
kmer::{Kmer, KmerEncoder, KmerError, KmerLen, KmerSet, MaxLenToType, SupportedKmerLen},
math::{AnyInt, Uint},
prelude::KmerCounter,
};
use std::hash::{Hash, Hasher, RandomState};
pub(crate) type ThreeBitKmerLen<const MAX_LEN: usize> = KmerLen<MAX_LEN, ThreeBitKmerEncoder<MAX_LEN>>;
pub(crate) type ThreeBitMaxLenToType<const MAX_LEN: usize> = MaxLenToType<MAX_LEN, ThreeBitKmerEncoder<MAX_LEN>>;
pub type ThreeBitKmerSet<const MAX_LEN: usize, S = RandomState> = KmerSet<MAX_LEN, ThreeBitKmerEncoder<MAX_LEN>, S>;
pub type ThreeBitKmerCounter<const MAX_LEN: usize, S = RandomState> = KmerCounter<MAX_LEN, ThreeBitKmerEncoder<MAX_LEN>, S>;
#[derive(Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Hash, Debug)]
#[repr(transparent)]
pub struct ThreeBitEncodedKmer<const MAX_LEN: usize>(ThreeBitMaxLenToType<MAX_LEN>)
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen;
impl<const MAX_LEN: usize, T: Uint> From<T> for ThreeBitEncodedKmer<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen<T = T>,
{
#[inline]
fn from(value: T) -> Self {
Self(value)
}
}
impl<const MAX_LEN: usize, T: Uint> std::fmt::Display for ThreeBitEncodedKmer<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen<T = T>,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut kmer = self.0;
let mut buffer = [0; MAX_LEN];
let mut start = MAX_LEN;
while kmer != T::ZERO && start > 0 {
start -= 1;
let encoded_base = kmer & T::from_literal(0b111);
buffer[start] = ThreeBitKmerEncoder::decode_base(encoded_base);
kmer >>= 3;
}
f.write_str(unsafe { std::str::from_utf8_unchecked(&buffer[start..]) })
}
}
impl<const MAX_LEN: usize, T: Uint> std::fmt::Binary for ThreeBitEncodedKmer<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen<T = T>,
{
#[inline]
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Binary::fmt(&self.0, f)
}
}
#[derive(Copy, Clone, Debug)]
pub struct ThreeBitKmerEncoder<const MAX_LEN: usize>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen, {
kmer_length: usize,
kmer_mask: ThreeBitMaxLenToType<MAX_LEN>,
}
impl<T: Uint, const MAX_LEN: usize> ThreeBitKmerEncoder<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen<T = T>,
{
#[inline]
#[must_use]
pub fn encode_base(base: u8) -> T {
T::from(THREE_BIT_MAPPING[base])
}
#[inline]
#[must_use]
pub fn decode_base(encoded_base: T) -> u8 {
b"000NACGT"[encoded_base.cast_as::<usize>()]
}
}
impl<const MAX_LEN: usize> PartialEq for ThreeBitKmerEncoder<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
#[inline]
fn eq(&self, other: &Self) -> bool {
self.kmer_length == other.kmer_length
}
}
impl<const MAX_LEN: usize> Eq for ThreeBitKmerEncoder<MAX_LEN> where ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen {}
impl<const MAX_LEN: usize> Hash for ThreeBitKmerEncoder<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
#[inline]
fn hash<H: Hasher>(&self, state: &mut H) {
self.kmer_length.hash(state);
}
}
impl<const MAX_LEN: usize, T: Uint> KmerEncoder<MAX_LEN> for ThreeBitKmerEncoder<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen<T = T>,
{
type EncodedKmer = ThreeBitEncodedKmer<MAX_LEN>;
type SeqIter<'a> = ThreeBitKmerIterator<'a, MAX_LEN>;
type SeqIntoIter = ThreeBitKmerIntoIterator<MAX_LEN>;
type SeqIterRev<'a> = ThreeBitKmerIteratorRev<'a, MAX_LEN>;
type SeqIntoIterRev = ThreeBitKmerIntoIteratorRev<MAX_LEN>;
#[inline]
fn new(kmer_length: usize) -> Result<Self, KmerError> {
if kmer_length <= MAX_LEN && kmer_length > 1 {
Ok(Self {
kmer_length,
kmer_mask: (T::ONE << (3 * kmer_length)) - T::ONE,
})
} else {
Err(KmerError::InvalidLength)
}
}
#[inline]
fn kmer_length(&self) -> usize {
self.kmer_length
}
#[inline]
fn encode_kmer(&self, kmer: impl AsRef<[u8]>) -> Self::EncodedKmer {
let mut encoded_kmer = T::ZERO;
for &base in kmer.as_ref() {
encoded_kmer = (encoded_kmer << 3) | Self::encode_base(base);
}
ThreeBitEncodedKmer(encoded_kmer)
}
fn decode_kmer(&self, mut encoded_kmer: Self::EncodedKmer) -> Kmer<MAX_LEN> {
let mut buffer = [0; MAX_LEN];
for i in (0..self.kmer_length).rev() {
let encoded_base = encoded_kmer.0 & T::from_literal(0b111);
buffer[i] = Self::decode_base(encoded_base);
encoded_kmer.0 >>= 3;
}
unsafe { Kmer::new_unchecked(self.kmer_length, buffer) }
}
#[inline]
fn iter_from_sequence<'a, S: AsRef<[u8]> + ?Sized>(&self, seq: &'a S) -> Self::SeqIter<'a> {
ThreeBitKmerIterator::new(self, seq.as_ref())
}
#[inline]
fn iter_consuming_seq<S: Into<Vec<u8>>>(&self, seq: S) -> Self::SeqIntoIter {
ThreeBitKmerIntoIterator::new(self, seq.into())
}
#[inline]
fn iter_from_sequence_rev<'a, S: AsRef<[u8]> + ?Sized>(&self, seq: &'a S) -> Self::SeqIterRev<'a> {
ThreeBitKmerIteratorRev::new(self, seq.as_ref())
}
#[inline]
fn iter_consuming_seq_rev<'a, S: Into<Vec<u8>>>(&self, seq: S) -> Self::SeqIntoIterRev {
ThreeBitKmerIntoIteratorRev::new(self, seq.into())
}
}
pub struct ThreeBitKmerIterator<'a, const MAX_LEN: usize>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen, {
current_kmer: ThreeBitEncodedKmer<MAX_LEN>,
remaining: std::slice::Iter<'a, u8>,
kmer_mask: ThreeBitMaxLenToType<MAX_LEN>,
}
impl<'a, const MAX_LEN: usize> ThreeBitKmerIterator<'a, MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
#[inline]
fn new(encoder: &ThreeBitKmerEncoder<MAX_LEN>, seq: &'a [u8]) -> Self {
if seq.len() < encoder.kmer_length {
ThreeBitKmerIterator {
current_kmer: ThreeBitEncodedKmer(ThreeBitMaxLenToType::<MAX_LEN>::ZERO),
remaining: [].iter(),
kmer_mask: encoder.kmer_mask,
}
} else {
let index = encoder.kmer_length - 1;
let pre_first_kmer = encoder.encode_kmer(&seq[..index]);
ThreeBitKmerIterator {
current_kmer: pre_first_kmer,
remaining: seq[index..].iter(),
kmer_mask: encoder.kmer_mask,
}
}
}
#[inline]
#[must_use]
pub(crate) fn shift_and_add_base(
new_base: u8, kmer: ThreeBitEncodedKmer<MAX_LEN>, kmer_mask: ThreeBitMaxLenToType<MAX_LEN>,
) -> ThreeBitEncodedKmer<MAX_LEN> {
let encoded_base = ThreeBitKmerEncoder::encode_base(new_base);
let shifted_kmer = kmer.0 << 3;
((shifted_kmer | encoded_base) & kmer_mask).into()
}
}
impl<const MAX_LEN: usize> Iterator for ThreeBitKmerIterator<'_, MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
type Item = ThreeBitEncodedKmer<MAX_LEN>;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
let base = *self.remaining.next()?;
self.current_kmer = ThreeBitKmerIterator::shift_and_add_base(base, self.current_kmer, self.kmer_mask);
Some(self.current_kmer)
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
self.remaining.size_hint()
}
}
impl<const MAX_LEN: usize> ExactSizeIterator for ThreeBitKmerIterator<'_, MAX_LEN> where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen
{
}
pub struct ThreeBitKmerIntoIterator<const MAX_LEN: usize>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen, {
seq: Vec<u8>,
index: usize,
current_kmer: ThreeBitEncodedKmer<MAX_LEN>,
kmer_mask: ThreeBitMaxLenToType<MAX_LEN>,
}
impl<const MAX_LEN: usize> ThreeBitKmerIntoIterator<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
#[inline]
fn new(encoder: &ThreeBitKmerEncoder<MAX_LEN>, seq: Vec<u8>) -> Self {
if seq.len() < encoder.kmer_length {
ThreeBitKmerIntoIterator {
index: seq.len(),
seq,
current_kmer: ThreeBitEncodedKmer(ThreeBitMaxLenToType::<MAX_LEN>::ZERO),
kmer_mask: encoder.kmer_mask,
}
} else {
let index = encoder.kmer_length - 1;
let pre_first_kmer = encoder.encode_kmer(&seq[..index]);
ThreeBitKmerIntoIterator {
seq,
index,
current_kmer: pre_first_kmer,
kmer_mask: encoder.kmer_mask,
}
}
}
}
impl<const MAX_LEN: usize> Iterator for ThreeBitKmerIntoIterator<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
type Item = ThreeBitEncodedKmer<MAX_LEN>;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
let base = *self.seq.get(self.index)?;
self.index += 1;
self.current_kmer = ThreeBitKmerIterator::shift_and_add_base(base, self.current_kmer, self.kmer_mask);
Some(self.current_kmer)
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
let size = self.seq.len() - self.index;
(size, Some(size))
}
}
impl<const MAX_LEN: usize> ExactSizeIterator for ThreeBitKmerIntoIterator<MAX_LEN> where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen
{
}
pub struct ThreeBitKmerIteratorRev<'a, const MAX_LEN: usize>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen, {
current_kmer: ThreeBitEncodedKmer<MAX_LEN>,
remaining: std::slice::Iter<'a, u8>,
kmer_length: usize,
}
impl<'a, const MAX_LEN: usize> ThreeBitKmerIteratorRev<'a, MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
fn new(encoder: &ThreeBitKmerEncoder<MAX_LEN>, seq: &'a [u8]) -> Self {
if seq.len() < encoder.kmer_length {
ThreeBitKmerIteratorRev {
current_kmer: ThreeBitEncodedKmer(ThreeBitMaxLenToType::<MAX_LEN>::ZERO),
remaining: [].iter(),
kmer_length: encoder.kmer_length,
}
} else {
let index = seq.len() + 1 - encoder.kmer_length;
let pre_last_kmer = encoder.encode_kmer(&seq[index..]).0 << 3;
ThreeBitKmerIteratorRev {
current_kmer: pre_last_kmer.into(),
remaining: seq[..index].iter(),
kmer_length: encoder.kmer_length,
}
}
}
#[inline]
#[must_use]
pub(crate) fn shift_and_add_base_rev(
new_base: u8, kmer: ThreeBitEncodedKmer<MAX_LEN>, kmer_length: usize,
) -> ThreeBitEncodedKmer<MAX_LEN> {
let encoded_base = ThreeBitKmerEncoder::encode_base(new_base);
let shifted_encoded_base = encoded_base << (3 * (kmer_length - 1));
((kmer.0 >> 3) | shifted_encoded_base).into()
}
}
impl<const MAX_LEN: usize> Iterator for ThreeBitKmerIteratorRev<'_, MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
type Item = ThreeBitEncodedKmer<MAX_LEN>;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
let base = *self.remaining.next_back()?;
self.current_kmer = ThreeBitKmerIteratorRev::shift_and_add_base_rev(base, self.current_kmer, self.kmer_length);
Some(self.current_kmer)
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
self.remaining.size_hint()
}
}
impl<const MAX_LEN: usize> ExactSizeIterator for ThreeBitKmerIteratorRev<'_, MAX_LEN> where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen
{
}
pub struct ThreeBitKmerIntoIteratorRev<const MAX_LEN: usize>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen, {
seq: Vec<u8>,
index: usize,
current_kmer: ThreeBitEncodedKmer<MAX_LEN>,
kmer_length: usize,
}
impl<const MAX_LEN: usize> ThreeBitKmerIntoIteratorRev<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
fn new(encoder: &ThreeBitKmerEncoder<MAX_LEN>, seq: Vec<u8>) -> Self {
if seq.len() < encoder.kmer_length {
ThreeBitKmerIntoIteratorRev {
index: 0,
seq,
current_kmer: ThreeBitEncodedKmer(ThreeBitMaxLenToType::<MAX_LEN>::ZERO),
kmer_length: encoder.kmer_length,
}
} else {
let index = seq.len() + 1 - encoder.kmer_length;
let pre_last_kmer = encoder.encode_kmer(&seq[index..]).0 << 3;
ThreeBitKmerIntoIteratorRev {
index,
seq,
current_kmer: pre_last_kmer.into(),
kmer_length: encoder.kmer_length,
}
}
}
}
impl<const MAX_LEN: usize> Iterator for ThreeBitKmerIntoIteratorRev<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
type Item = ThreeBitEncodedKmer<MAX_LEN>;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
self.index = self.index.checked_sub(1)?;
let base = self.seq[self.index];
self.current_kmer = ThreeBitKmerIteratorRev::shift_and_add_base_rev(base, self.current_kmer, self.kmer_length);
Some(self.current_kmer)
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
(self.index, Some(self.index))
}
}
impl<const MAX_LEN: usize> ExactSizeIterator for ThreeBitKmerIntoIteratorRev<MAX_LEN> where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen
{
}
pub struct ThreeBitOneMismatchIter<const MAX_LEN: usize>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen, {
encoded_kmer: ThreeBitEncodedKmer<MAX_LEN>,
kmer_length: usize,
current_kmer: ThreeBitEncodedKmer<MAX_LEN>,
current_index: usize,
current_base_num: usize,
set_mask_third_bit: ThreeBitMaxLenToType<MAX_LEN>,
not_finished: bool,
}
impl<const MAX_LEN: usize, T: Uint> ThreeBitOneMismatchIter<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen<T = T>,
{
#[inline]
#[must_use]
pub(crate) fn new(encoded_kmer: ThreeBitEncodedKmer<MAX_LEN>, kmer_encoder: &ThreeBitKmerEncoder<MAX_LEN>) -> Self {
Self {
encoded_kmer,
kmer_length: kmer_encoder.kmer_length(),
current_kmer: encoded_kmer,
current_index: 0,
current_base_num: 0,
set_mask_third_bit: T::from_literal(0b100),
not_finished: true,
}
}
}
impl<const MAX_LEN: usize> Iterator for ThreeBitOneMismatchIter<MAX_LEN>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen,
{
type Item = ThreeBitEncodedKmer<MAX_LEN>;
#[inline]
fn next(&mut self) -> Option<Self::Item> {
while self.current_index < self.kmer_length {
if self.current_base_num < 4 {
self.current_base_num += 1;
self.current_kmer = ((self.current_kmer.0 | self.set_mask_third_bit)
- ((self.current_kmer.0 & self.set_mask_third_bit) >> 2))
.into();
return Some(self.current_kmer);
}
self.current_kmer = self.encoded_kmer;
self.current_index += 1;
self.current_base_num = 0;
self.set_mask_third_bit <<= 3;
}
if self.not_finished {
self.not_finished = false;
return Some(self.encoded_kmer);
}
None
}
#[inline]
fn size_hint(&self) -> (usize, Option<usize>) {
let size = usize::from(self.not_finished) + 4 * (self.kmer_length - self.current_index) - self.current_base_num;
(size, Some(size))
}
}
impl<const MAX_LEN: usize> ExactSizeIterator for ThreeBitOneMismatchIter<MAX_LEN> where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen
{
}
pub struct ThreeBitMismatchIter<const MAX_LEN: usize, const N: usize>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen, {
encoded_kmer: ThreeBitMaxLenToType<MAX_LEN>,
current_kmer: ThreeBitMaxLenToType<MAX_LEN>,
current_subset: ThreeBitMaxLenToType<MAX_LEN>,
subset_size: usize,
masks_and_times_mutated: [(ThreeBitMaxLenToType<MAX_LEN>, usize); N],
max_subset: ThreeBitMaxLenToType<MAX_LEN>,
}
impl<const MAX_LEN: usize, const N: usize, T: Uint> ThreeBitMismatchIter<MAX_LEN, N>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen<T = T>,
{
#[inline]
#[must_use]
pub(crate) fn new(encoded_kmer: ThreeBitEncodedKmer<MAX_LEN>, encoder: &ThreeBitKmerEncoder<MAX_LEN>) -> Self {
const { assert!(N >= 2) }
Self {
encoded_kmer: encoded_kmer.0,
current_kmer: encoded_kmer.0,
current_subset: T::ZERO,
subset_size: 0,
masks_and_times_mutated: [(T::from_literal(0b100), usize::MAX); N],
max_subset: T::ONE << encoder.kmer_length(),
}
}
}
impl<const MAX_LEN: usize, const N: usize, T: Uint> Iterator for ThreeBitMismatchIter<MAX_LEN, N>
where
ThreeBitKmerLen<MAX_LEN>: SupportedKmerLen<T = T>,
{
type Item = ThreeBitEncodedKmer<MAX_LEN>;
fn next(&mut self) -> Option<Self::Item> {
let out = ThreeBitEncodedKmer(self.current_kmer);
if self.masks_and_times_mutated[0].1 < 4 {
self.current_kmer = (self.current_kmer | self.masks_and_times_mutated[0].0)
- ((self.current_kmer & self.masks_and_times_mutated[0].0) >> 2);
self.masks_and_times_mutated[0].1 += 1;
} else if let Some(i) = self.masks_and_times_mutated[..self.subset_size]
.iter()
.position(|(_, base_num)| *base_num < 4)
{
for (set_mask_third_bit, base_num) in &mut self.masks_and_times_mutated[..i] {
let set_mask_third_bit = *set_mask_third_bit;
self.current_kmer =
(self.current_kmer | set_mask_third_bit) - ((self.current_kmer & set_mask_third_bit) >> 2);
self.current_kmer =
(self.current_kmer | set_mask_third_bit) - ((self.current_kmer & set_mask_third_bit) >> 2);
*base_num = 1;
}
let set_mask_third_bit = self.masks_and_times_mutated[i].0;
self.current_kmer = (self.current_kmer | set_mask_third_bit) - ((self.current_kmer & set_mask_third_bit) >> 2);
self.masks_and_times_mutated[i].1 += 1;
} else {
if (self.current_subset.count_ones() as usize) < N {
self.current_subset += T::ONE;
} else {
self.current_subset = (self.current_subset | (self.current_subset - T::ONE)) + T::ONE;
}
if self.current_subset >= self.max_subset {
let out = (self.current_subset == self.max_subset).then_some(out);
self.current_subset = self.max_subset + T::ONE;
return out;
}
self.current_kmer = self.encoded_kmer;
self.subset_size = 0;
let mut processed_bases = 0;
let mut temp_subset = self.current_subset;
while temp_subset != T::ZERO {
let num_zeros = temp_subset.trailing_zeros() as usize;
temp_subset >>= num_zeros + 1;
processed_bases += num_zeros;
self.masks_and_times_mutated[self.subset_size].0 = T::from_literal(0b100) << (processed_bases * 3);
processed_bases += 1;
self.subset_size += 1;
}
for (set_mask_third_bit, base_num) in &mut self.masks_and_times_mutated[..self.subset_size] {
self.current_kmer =
(self.current_kmer | *set_mask_third_bit) - ((self.current_kmer & *set_mask_third_bit) >> 2);
*base_num = 1;
}
}
Some(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_valid_decoding() {
assert_eq!(ThreeBitKmerEncoder::<21>::decode_base(0), b'0');
assert_eq!(ThreeBitKmerEncoder::<21>::decode_base(1), b'0');
assert_eq!(ThreeBitKmerEncoder::<21>::decode_base(2), b'0');
assert_eq!(ThreeBitKmerEncoder::<21>::decode_base(3), b'N');
assert_eq!(ThreeBitKmerEncoder::<21>::decode_base(4), b'A');
assert_eq!(ThreeBitKmerEncoder::<21>::decode_base(5), b'C');
assert_eq!(ThreeBitKmerEncoder::<21>::decode_base(6), b'G');
assert_eq!(ThreeBitKmerEncoder::<21>::decode_base(7), b'T');
}
#[test]
#[should_panic(expected = "index out of bounds: the len is 8 but the index is 8")]
fn test_invalid_decoding() {
let _ = ThreeBitKmerEncoder::<21>::decode_base(8);
}
}