const HOME_ROW_STAGGER: f32 = 0.5;
const BOTTOM_ROW_STAGGER: f32 = 1.0;
const TOUCH_SIGMA_KEYS: f32 = 0.5;
const NEIGHBOUR_RADIUS: f32 = 1.2;
fn key_coord(ch: char) -> Option<(f32, f32)> {
const TOP: &[u8] = b"qwertyuiop";
const HOME: &[u8] = b"asdfghjkl";
const BOTTOM: &[u8] = b"zxcvbnm";
let lower = ch.to_ascii_lowercase() as u8;
if let Some(col) = TOP.iter().position(|&k| k == lower) {
return Some((col as f32, 0.0));
}
if let Some(col) = HOME.iter().position(|&k| k == lower) {
return Some((col as f32 + HOME_ROW_STAGGER, 1.0));
}
if let Some(col) = BOTTOM.iter().position(|&k| k == lower) {
return Some((col as f32 + BOTTOM_ROW_STAGGER, 2.0));
}
None
}
fn key_distance(a: char, b: char) -> Option<f32> {
let (ax, ay) = key_coord(a)?;
let (bx, by) = key_coord(b)?;
Some(((ax - bx).powi(2) + (ay - by).powi(2)).sqrt())
}
fn neighbours(ch: char) -> Vec<(char, u16)> {
const LETTERS: &[u8] = b"abcdefghijklmnopqrstuvwxyz";
let Some((cx, cy)) = key_coord(ch) else {
return Vec::new();
};
let uppercase = ch.is_ascii_uppercase();
let mut near = Vec::new();
for &letter in LETTERS {
let candidate = letter as char;
if candidate == ch.to_ascii_lowercase() {
continue;
}
let (nx, ny) = key_coord(candidate).expect("letters have coordinates");
let squared_distance = (cx - nx).powi(2) + (cy - ny).powi(2);
let distance = squared_distance.sqrt();
if distance <= NEIGHBOUR_RADIUS {
let out = if uppercase {
candidate.to_ascii_uppercase()
} else {
candidate
};
let cost = (squared_distance / (2.0 * TOUCH_SIGMA_KEYS * TOUCH_SIGMA_KEYS))
.round()
.max(1.0) as u16;
near.push((out, cost, distance));
}
}
near.sort_by(|a, b| a.2.total_cmp(&b.2));
near.into_iter().map(|(ch, cost, _)| (ch, cost)).collect()
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct KeySlipVariant {
pub text: String,
pub cost: u16,
}
pub fn key_slip_variants(input: &str, max_variants: usize) -> Vec<KeySlipVariant> {
if max_variants == 0 || input.is_empty() {
return Vec::new();
}
let mut scored: Vec<(f32, String, u16)> = Vec::new();
for (byte_index, ch) in input.char_indices() {
if !ch.is_ascii_alphabetic() {
continue;
}
for (near, cost) in neighbours(ch) {
let mut text = String::with_capacity(input.len());
text.push_str(&input[..byte_index]);
text.push(near);
text.push_str(&input[byte_index + ch.len_utf8()..]);
if text == input || scored.iter().any(|(_, existing, _)| *existing == text) {
continue;
}
let distance = key_distance(ch, near).unwrap_or(f32::MAX);
scored.push((distance, text, cost));
}
}
scored.sort_by(|a, b| a.0.total_cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
scored.truncate(max_variants);
scored
.into_iter()
.map(|(_, text, cost)| KeySlipVariant { text, cost })
.collect()
}
const MAX_KEY_SLIP_VARIANTS: usize = 64;
pub fn key_slip_repaired_outputs<T, W>(
input: &str,
baseline_output: &str,
baseline_frequency: Option<u64>,
mut transliterate: T,
mut is_lexicon_word: W,
) -> Vec<super::roman_repair::RomanRepairedOutput>
where
T: FnMut(&str) -> String,
W: FnMut(&str) -> bool,
{
if baseline_frequency.is_some() {
return Vec::new();
}
let mut outputs = Vec::new();
for variant in key_slip_variants(input, MAX_KEY_SLIP_VARIANTS) {
let bangla = transliterate(&variant.text);
if bangla == baseline_output || !is_lexicon_word(&bangla) {
continue;
}
outputs.push(super::roman_repair::RomanRepairedOutput {
roman_input: variant.text,
bangla_output: bangla,
repair_kind: "qwerty_key_slip",
repair_cost: variant.cost,
});
}
outputs
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn neighbours_are_physically_adjacent_keys() {
let near: Vec<char> = neighbours('g').into_iter().map(|(ch, _)| ch).collect();
for expected in ['f', 'h', 't', 'y', 'v', 'b'] {
assert!(near.contains(&expected), "g should neighbour {expected}: {near:?}");
}
for far in ['a', 'p', 'q', 'm', 'l'] {
assert!(!near.contains(&far), "g should not neighbour {far}");
}
}
#[test]
fn horizontal_neighbours_are_cheaper_than_diagonal() {
let f_g = key_distance('f', 'g').unwrap();
let g_t = key_distance('g', 't').unwrap();
assert!(f_g < g_t, "same-row slip should be nearer than diagonal");
assert!((f_g - 1.0).abs() < 1e-6);
}
#[test]
fn case_is_carried_through_as_shift_state() {
let near = neighbours('S');
assert!(near.iter().all(|(ch, _)| ch.is_ascii_uppercase()));
let chars: Vec<char> = near.iter().map(|(ch, _)| *ch).collect();
assert!(chars.contains(&'A') && chars.contains(&'D'));
}
#[test]
fn variants_are_single_substitutions_bounded_and_ordered() {
let variants = key_slip_variants("bangla", 12);
assert!(!variants.is_empty());
for variant in &variants {
assert_eq!(variant.text.chars().count(), "bangla".chars().count());
let diffs = variant
.text
.chars()
.zip("bangla".chars())
.filter(|(a, b)| a != b)
.count();
assert_eq!(diffs, 1, "{} should differ in one position", variant.text);
}
assert!(variants.iter().any(|variant| variant.text == "banhla"));
assert!(variants.len() <= 12);
assert!(variants.windows(2).all(|w| w[0].cost <= w[1].cost));
}
#[test]
fn non_letter_positions_are_left_alone() {
let variants = key_slip_variants("a5", 8);
assert!(variants.iter().all(|variant| variant.text.ends_with('5')));
}
#[test]
#[ignore = "needs the resolved data/autocorrect/models/bn.fst; QWERTY recall/precision probe"]
fn key_slip_recall_and_precision_probe() {
use crate::{
roman_repaired_outputs, FstLexicon, FstRepairedBaseline, FstSuggestOptions,
ObadhEngine, RomanRepairOptions, DEFAULT_ROMAN_REPAIR_BEAM_SIZE,
};
let Ok(bytes) = std::fs::read("data/autocorrect/models/bn.fst") else {
eprintln!("skip: bn.fst not resolved");
return;
};
let lexicon = FstLexicon::from_bytes(bytes).expect("load fst");
let engine = ObadhEngine::new();
let options = FstSuggestOptions {
max_distance: 2,
max_edit_cost: None,
max_candidates: 512,
max_prefix_candidates: 8,
response_candidates: 8,
};
const MIN_FREQ: u64 = 500;
let romans = [
"ami", "tumi", "amar", "tomar", "bangla", "desh", "bhalo", "kemon", "manush",
"kotha", "kaj", "din", "raat", "boi", "naam", "ghor", "jol", "gaan", "chokh",
"hat", "mon", "jibon", "somoy", "bhasha", "chele", "meye", "baba", "bhai", "bon",
"sokal", "ekhon", "tokhon", "jonno", "karon", "kintu", "ebong", "kore", "kori",
"bola", "dekha", "jani", "hobe", "ache", "bikel", "shohor", "gram", "rasta",
"notun", "purano", "boro", "choto",
];
let run = |roman: &str| -> Vec<String> {
let base = engine.transliterate(roman);
let mut reps = roman_repaired_outputs(
roman,
&base,
RomanRepairOptions {
max_repairs: DEFAULT_ROMAN_REPAIR_BEAM_SIZE,
},
|repair| engine.transliterate(repair),
);
reps.extend(key_slip_repaired_outputs(
roman,
&base,
lexicon.exact_frequency(&base),
|repair| engine.transliterate(repair),
|word| lexicon.exact_frequency(word).is_some(),
));
let baselines: Vec<FstRepairedBaseline> = reps
.iter()
.map(|repair| FstRepairedBaseline {
roman_input: &repair.roman_input,
bangla_output: &repair.bangla_output,
repair_kind: repair.repair_kind,
repair_cost: repair.repair_cost,
})
.collect();
lexicon
.suggest_with_repaired_baselines(&base, &baselines, options)
.expect("suggest")
.candidates
.into_iter()
.map(|candidate| candidate.text)
.collect()
};
let (mut valid, mut precision_kept) = (0usize, 0usize);
let (mut trials, mut recall_1, mut recall_5) = (0usize, 0usize, 0usize);
let (mut gate_open, mut gate_open_r1, mut gate_open_r5) = (0usize, 0usize, 0usize);
let mut slip_is_word = 0usize;
for roman in romans {
let expected = engine.transliterate(roman);
match lexicon.exact_frequency(&expected) {
Some(frequency) if frequency >= MIN_FREQ => {}
_ => continue,
}
valid += 1;
if run(roman).first().map(String::as_str) == Some(expected.as_str()) {
precision_kept += 1;
}
for slip in key_slip_variants(roman, 64) {
let base = engine.transliterate(&slip.text);
if base == expected {
continue; }
trials += 1;
let gate_is_open = lexicon.exact_frequency(&base).is_none();
if !gate_is_open {
slip_is_word += 1; }
let ranked = run(&slip.text);
let rank = ranked.iter().position(|text| *text == expected);
if rank == Some(0) {
recall_1 += 1;
}
if rank.is_some_and(|position| position < 5) {
recall_5 += 1;
}
if gate_is_open {
gate_open += 1;
if rank == Some(0) {
gate_open_r1 += 1;
}
if rank.is_some_and(|position| position < 5) {
gate_open_r5 += 1;
}
}
}
}
let pct = |num: usize, den: usize| if den == 0 { 0.0 } else { 100.0 * num as f64 / den as f64 };
eprintln!(
"QWERTY precision: {precision_kept}/{valid} correctly-typed words kept at #1 (target 100%)"
);
eprintln!(
"recall over ALL {trials} injected slips: @1 {:.1}% @5 {:.1}%",
pct(recall_1, trials),
pct(recall_5, trials),
);
eprintln!(
"recall over the {gate_open} gate-OPEN slips (non-word typo): @1 {:.1}% @5 {:.1}%",
pct(gate_open_r1, gate_open),
pct(gate_open_r5, gate_open),
);
eprintln!("({slip_is_word} slips landed on another valid word — left untouched by design)");
assert_eq!(precision_kept, valid, "a correctly-typed word must never be demoted");
}
}