#[cfg(feature = "rand")]
use rand::rngs::StdRng;
#[cfg(feature = "rand")]
use rand::seq::SliceRandom;
#[cfg(feature = "rand")]
use rand::{Rng, SeedableRng};
use std::collections::HashMap;
pub struct TypoGenerator {
rng: StdRng,
alphabet: Vec<char>,
}
impl TypoGenerator {
pub fn new(seed: u64) -> Self {
Self {
rng: StdRng::seed_from_u64(seed),
alphabet: "abcdefghijklmnopqrstuvwxyz".chars().collect(),
}
}
pub fn generate_typos(&mut self, word: &str, distance: usize, count: usize) -> Vec<String> {
let mut typos = Vec::with_capacity(count);
for _ in 0..count {
let typo = self.generate_single_typo(word, distance);
typos.push(typo);
}
typos
}
pub fn all_distance_1(&self, word: &str) -> Vec<String> {
let chars: Vec<char> = word.chars().collect();
let mut typos = Vec::new();
for i in 0..chars.len() {
let mut deleted = chars.clone();
deleted.remove(i);
typos.push(deleted.iter().collect());
}
for i in 0..=chars.len() {
for &c in &self.alphabet {
let mut inserted = chars.clone();
inserted.insert(i, c);
typos.push(inserted.iter().collect());
}
}
for i in 0..chars.len() {
for &c in &self.alphabet {
if c != chars[i] {
let mut substituted = chars.clone();
substituted[i] = c;
typos.push(substituted.iter().collect());
}
}
}
for i in 0..chars.len().saturating_sub(1) {
let mut transposed = chars.clone();
transposed.swap(i, i + 1);
typos.push(transposed.iter().collect());
}
typos
}
fn generate_single_typo(&mut self, word: &str, distance: usize) -> String {
let mut result = word.to_string();
for _ in 0..distance {
result = self.apply_random_edit(&result);
}
result
}
fn apply_random_edit(&mut self, word: &str) -> String {
if word.is_empty() {
return self.alphabet[self.rng.gen_range(0..self.alphabet.len())].to_string();
}
let chars: Vec<char> = word.chars().collect();
let edit_type = self.rng.gen_range(0..4);
match edit_type {
0 => self.apply_deletion(&chars),
1 => self.apply_insertion(&chars),
2 => self.apply_substitution(&chars),
_ => self.apply_transposition(&chars),
}
}
fn apply_deletion(&mut self, chars: &[char]) -> String {
if chars.is_empty() {
return String::new();
}
let pos = self.rng.gen_range(0..chars.len());
let mut result = chars.to_vec();
result.remove(pos);
result.iter().collect()
}
fn apply_insertion(&mut self, chars: &[char]) -> String {
let pos = self.rng.gen_range(0..=chars.len());
let new_char = self.alphabet[self.rng.gen_range(0..self.alphabet.len())];
let mut result = chars.to_vec();
result.insert(pos, new_char);
result.iter().collect()
}
fn apply_substitution(&mut self, chars: &[char]) -> String {
if chars.is_empty() {
return String::new();
}
let pos = self.rng.gen_range(0..chars.len());
let new_char = self.alphabet[self.rng.gen_range(0..self.alphabet.len())];
let mut result = chars.to_vec();
result[pos] = new_char;
result.iter().collect()
}
fn apply_transposition(&mut self, chars: &[char]) -> String {
if chars.len() < 2 {
return chars.iter().collect();
}
let pos = self.rng.gen_range(0..chars.len() - 1);
let mut result = chars.to_vec();
result.swap(pos, pos + 1);
result.iter().collect()
}
}
#[derive(Debug, Clone)]
pub struct QueryWorkload {
pub queries: Vec<(String, usize)>,
}
impl QueryWorkload {
pub fn from_frequencies(
frequencies: &HashMap<String, usize>,
total_tokens: usize,
num_queries: usize,
seed: u64,
) -> Self {
let mut rng = StdRng::seed_from_u64(seed);
let mut words: Vec<_> = frequencies.iter().collect();
words.sort_unstable_by(|a, b| b.1.cmp(a.1));
let mut cumulative = Vec::with_capacity(words.len());
let mut sum = 0;
for (word, &freq) in &words {
sum += freq;
cumulative.push((word.as_str(), sum));
}
let mut queries = Vec::with_capacity(num_queries);
for _ in 0..num_queries {
let sample = rng.gen_range(0..total_tokens);
let idx = cumulative
.binary_search_by(|&(_, cum)| cum.cmp(&sample))
.unwrap_or_else(|i| i);
let word = cumulative[idx.min(cumulative.len() - 1)].0;
let freq = frequencies[word];
queries.push((word.to_string(), freq));
}
Self { queries }
}
pub fn uniform(words: &[String], num_queries: usize, seed: u64) -> Self {
let mut rng = StdRng::seed_from_u64(seed);
let mut queries = Vec::with_capacity(num_queries);
for _ in 0..num_queries {
let word = words
.choose(&mut rng)
.expect("uniform query gen requires non-empty word list");
queries.push((word.clone(), 1));
}
Self { queries }
}
pub fn query_strings(&self) -> Vec<&str> {
self.queries.iter().map(|(s, _)| s.as_str()).collect()
}
pub fn unique_queries(&self) -> Vec<&str> {
let mut unique: Vec<_> = self.queries.iter().map(|(s, _)| s.as_str()).collect();
unique.sort_unstable();
unique.dedup();
unique
}
pub fn stats(&self) -> WorkloadStats {
let unique = self.unique_queries().len();
let total = self.queries.len();
let frequencies: Vec<_> = self.queries.iter().map(|(_, f)| *f).collect();
let min_freq = *frequencies.iter().min().unwrap_or(&0);
let max_freq = *frequencies.iter().max().unwrap_or(&0);
let avg_freq = if !frequencies.is_empty() {
frequencies.iter().sum::<usize>() as f64 / frequencies.len() as f64
} else {
0.0
};
WorkloadStats {
total_queries: total,
unique_queries: unique,
min_frequency: min_freq,
max_frequency: max_freq,
avg_frequency: avg_freq,
}
}
}
#[derive(Debug, Clone)]
pub struct WorkloadStats {
pub total_queries: usize,
pub unique_queries: usize,
pub min_frequency: usize,
pub max_frequency: usize,
pub avg_frequency: f64,
}
#[cfg(all(test, feature = "rand"))]
mod tests {
use super::*;
#[test]
fn test_typo_generator_distance_1() {
let mut gen = TypoGenerator::new(42);
let typos = gen.generate_typos("test", 1, 10);
assert_eq!(typos.len(), 10);
for typo in &typos {
assert!(typo.len() >= 3 && typo.len() <= 5);
}
}
#[test]
fn test_typo_generator_all_distance_1() {
let gen = TypoGenerator::new(42);
let typos = gen.all_distance_1("ab");
assert_eq!(typos.len(), 131);
assert!(typos.contains(&"a".to_string())); assert!(typos.contains(&"b".to_string())); assert!(typos.contains(&"ba".to_string())); assert!(typos.contains(&"aab".to_string())); assert!(typos.contains(&"xb".to_string())); }
#[test]
fn test_query_workload_uniform() {
let words = vec!["hello".to_string(), "world".to_string(), "test".to_string()];
let workload = QueryWorkload::uniform(&words, 100, 42);
assert_eq!(workload.queries.len(), 100);
let unique = workload.unique_queries();
assert!(unique.len() <= 3);
}
#[test]
fn test_query_workload_from_frequencies() {
let mut frequencies = HashMap::new();
frequencies.insert("the".to_string(), 100);
frequencies.insert("quick".to_string(), 10);
frequencies.insert("fox".to_string(), 1);
let total = 111;
let workload = QueryWorkload::from_frequencies(&frequencies, total, 1000, 42);
assert_eq!(workload.queries.len(), 1000);
let the_count = workload.queries.iter().filter(|(w, _)| w == "the").count();
let fox_count = workload.queries.iter().filter(|(w, _)| w == "fox").count();
assert!(the_count > fox_count * 10);
}
#[test]
fn test_workload_stats() {
let queries = vec![
("the".to_string(), 100),
("the".to_string(), 100),
("quick".to_string(), 10),
];
let workload = QueryWorkload { queries };
let stats = workload.stats();
assert_eq!(stats.total_queries, 3);
assert_eq!(stats.unique_queries, 2);
assert_eq!(stats.min_frequency, 10);
assert_eq!(stats.max_frequency, 100);
}
}