use std::cmp;
use std::collections::{hash_map::Iter, HashMap, VecDeque};
use serde::{Deserialize, Serialize};
#[derive(PartialEq, Clone, Debug, Serialize, Deserialize)]
pub(crate) struct NgramSet<const N: u8> {
map: HashMap<String, u32>,
size: usize,
}
impl<const N: u8> NgramSet<N> {
pub(crate) fn new() -> NgramSet<N> {
NgramSet {
map: HashMap::new(),
size: 0,
}
}
pub(crate) fn from_str(s: &str) -> NgramSet<N> {
let mut set = NgramSet::new();
set.analyze(s);
set
}
fn analyze(&mut self, s: &str) {
let words = s.split(' ');
let mut deque: VecDeque<&str> = VecDeque::with_capacity(N as usize);
for w in words {
deque.push_back(w);
if deque.len() == N as usize {
let parts = deque.iter().copied().collect::<Vec<&str>>();
self.add_gram(parts.join(" "));
deque.pop_front();
}
}
}
fn add_gram(&mut self, gram: String) {
let n = self.map.entry(gram).or_insert(0);
*n += 1;
self.size += 1;
}
fn get(&self, gram: &str) -> u32 {
if let Some(count) = self.map.get(gram) {
*count
} else {
0
}
}
const fn len(&self) -> usize {
self.size
}
const fn is_empty(&self) -> bool {
self.size == 0
}
#[expect(clippy::cast_precision_loss)]
pub(crate) fn dice<const M: u8>(&self, other: &NgramSet<M>) -> f32 {
if M != N {
return 0f32;
}
if self.is_empty() || other.is_empty() {
return 0f32;
}
let mut matches = 0;
if self.len() < other.len() {
for (gram, count) in self {
matches += cmp::min(*count, other.get(gram));
}
} else {
for (gram, count) in other {
matches += cmp::min(*count, self.get(gram));
}
}
(2.0 * matches as f32) / ((self.len() + other.len()) as f32)
}
}
impl<'a, const N: u8> IntoIterator for &'a NgramSet<N> {
type Item = (&'a String, &'a u32);
type IntoIter = Iter<'a, String, u32>;
fn into_iter(self) -> Self::IntoIter {
self.map.iter()
}
}
#[cfg(test)]
#[expect(clippy::float_cmp, reason = "ignore float comparisons in tests")]
mod tests {
use super::*;
#[test]
fn can_construct() {
let set = NgramSet::<2>::new();
assert_eq!(set.size, 0);
}
#[test]
fn no_nan() {
let a = NgramSet::<2>::from_str("");
let b = NgramSet::<2>::from_str("");
let score = a.dice(&b);
assert!(!score.is_nan());
}
#[test]
fn same_size() {
let a = NgramSet::<2>::from_str("");
let b = NgramSet::<3>::from_str("");
let score = a.dice(&b);
assert_eq!(0f32, score);
}
#[test]
fn identical() {
let a = NgramSet::<2>::from_str("one two three apple banana");
let b = NgramSet::<2>::from_str("one two three apple banana");
let score = a.dice(&b);
assert_eq!(1f32, score);
}
}