use std::fmt::{self, Write as _};
use crate::MAX_SPARSE_GRAM_SIZE;
const MULTIPLICATIVE_HASH: u64 = 0x9E37_79B9_7F4A_7C15;
#[derive(Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[repr(transparent)]
pub struct NGram(pub(crate) u32);
impl NGram {
const LEN_BIAS: u32 = 2;
const LEN_BITS: u32 = 3;
const LEN_MASK: u32 = (1 << Self::LEN_BITS) - 1;
const PAYLOAD_BITS: u32 = 24;
const PAYLOAD_MASK: u32 = (1 << Self::PAYLOAD_BITS) - 1;
pub(crate) const BITS: u32 = Self::PAYLOAD_BITS + Self::LEN_BITS;
pub(crate) const MASK: u32 = (1 << Self::BITS) - 1;
pub fn from_bytes(src: &[u8]) -> Self {
debug_assert!(
(Self::LEN_BIAS as usize..=MAX_SPARSE_GRAM_SIZE).contains(&src.len()),
"ngram length {} out of range [{}, {}]",
src.len(),
Self::LEN_BIAS,
MAX_SPARSE_GRAM_SIZE,
);
let payload = if src.len() <= 3 {
let mut p = 0u32;
for &byte in src {
p = (p << 8) | byte as u32;
}
p
} else {
let mut buf = [0u8; 8];
buf[..src.len()].copy_from_slice(src);
let product = u64::from_le_bytes(buf).wrapping_mul(MULTIPLICATIVE_HASH);
(product >> (u64::BITS - Self::PAYLOAD_BITS)) as u32
};
Self::pack(src.len(), payload)
}
#[inline]
pub(crate) fn from_window(value: u64, len: usize) -> Self {
debug_assert!(
(Self::LEN_BIAS as usize..=MAX_SPARSE_GRAM_SIZE).contains(&len),
"ngram length {len} out of range [{}, {}]",
Self::LEN_BIAS,
MAX_SPARSE_GRAM_SIZE,
);
let payload = if len <= 3 {
(value >> (u64::BITS - len as u32 * 8)) as u32
} else {
let product = value.swap_bytes().wrapping_mul(MULTIPLICATIVE_HASH);
(product >> (u64::BITS - Self::PAYLOAD_BITS)) as u32
};
Self::pack(len, payload)
}
#[inline]
fn pack(len: usize, payload: u32) -> Self {
let packed = ((len as u32 - Self::LEN_BIAS) << Self::PAYLOAD_BITS) | payload;
Self(mix27(packed))
}
#[inline]
pub fn len(&self) -> usize {
let packed = unmix27(self.0);
(((packed >> Self::PAYLOAD_BITS) & Self::LEN_MASK) + Self::LEN_BIAS) as usize
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline]
pub fn as_u32(&self) -> u32 {
self.0
}
}
fn mix27(mut x: u32) -> u32 {
debug_assert!(x <= NGram::MASK, "mix27 input must be a 27-bit value");
x ^= x >> 15;
x = x.wrapping_mul(0x2c1b_3c6d) & NGram::MASK;
x ^= x >> 12;
x = x.wrapping_mul(0x297a_2d39) & NGram::MASK;
x ^= x >> 15;
x
}
fn unmix27(mut x: u32) -> u32 {
debug_assert!(x <= NGram::MASK, "unmix27 input must be a 27-bit value");
x ^= x >> 15;
x = x.wrapping_mul(0x4f0_b109) & NGram::MASK; x ^= x >> 12;
x ^= x >> 24;
x = x.wrapping_mul(0x4ea_2d65) & NGram::MASK; x ^= x >> 15;
x
}
impl fmt::Debug for NGram {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let packed = unmix27(self.0);
let len = (((packed >> Self::PAYLOAD_BITS) & Self::LEN_MASK) + Self::LEN_BIAS) as usize;
let mut s = String::new();
if len <= 3 {
let payload = packed & Self::PAYLOAD_MASK;
let bytes = [(payload >> 16) as u8, (payload >> 8) as u8, payload as u8];
for &byte in &bytes[3 - len..3] {
if byte.is_ascii_graphic() || byte == b' ' {
s.push(byte as char);
} else {
write!(s, "\\x{byte:02x}")?;
}
}
} else {
write!(s, "{:#08x}", packed & Self::PAYLOAD_MASK)?;
}
write!(f, "NGram('{s}', len={len})")
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mix27_is_invertible() {
for x in (0..=NGram::MASK).step_by(97) {
assert_eq!(unmix27(mix27(x)), x, "round-trip failed for {x:#x}");
}
for x in [0, 1, 2, NGram::MASK - 1, NGram::MASK] {
assert_eq!(unmix27(mix27(x)), x, "round-trip failed for {x:#x}");
}
}
#[test]
fn test_from_bytes_roundtrip() {
for len in 2..=MAX_SPARSE_GRAM_SIZE {
let bytes = vec![b'a'; len];
assert_eq!(
NGram::from_bytes(&bytes).len(),
len,
"len mismatch for {len}"
);
}
}
#[test]
fn test_equal_content_equal_ngram() {
assert_eq!(NGram::from_bytes(b"abc"), NGram::from_bytes(b"abc"));
assert_eq!(NGram::from_bytes(b"abcdef"), NGram::from_bytes(b"abcdef"));
}
#[test]
fn test_short_grams_are_lossless() {
use std::collections::HashSet;
let mut seen = HashSet::new();
for a in 0u8..64 {
for b in 0u8..64 {
assert!(seen.insert(NGram::from_bytes(&[a, b])), "bigram collision");
for c in 0u8..8 {
assert!(
seen.insert(NGram::from_bytes(&[a, b, c])),
"trigram collision"
);
}
}
}
}
#[test]
fn test_same_content_different_length() {
let a = NGram::from_bytes(b"ab");
let b = NGram::from_bytes(b"abc");
assert_ne!(a, b);
assert_ne!(a.len(), b.len());
}
#[test]
fn from_window_matches_from_bytes() {
for len in (NGram::LEN_BIAS as usize)..=MAX_SPARSE_GRAM_SIZE {
let bytes: Vec<u8> = (0..len as u8).map(|i| b'a' + i).collect();
let mut buf = [0u8; 8];
buf[..len].copy_from_slice(&bytes);
let window = u64::from_be_bytes(buf);
assert_eq!(
NGram::from_window(window, len),
NGram::from_bytes(&bytes),
"mismatch for {len}-byte gram",
);
}
}
#[test]
fn test_default_is_not_empty() {
assert!(!NGram::default().is_empty());
assert_eq!(NGram::default().len(), 2);
}
}