use std::cmp::Ordering;
pub const NOUN: u8 = 1 << 0;
pub const VERB: u8 = 1 << 1;
pub const ADJ: u8 = 1 << 2;
pub const ADV: u8 = 1 << 3;
const TABLE: &str = include_str!("../wordnet/lemma-pos.txt");
const TABLE_DATA_START: usize = data_start(TABLE.as_bytes());
const fn data_start(bytes: &[u8]) -> usize {
let mut i = 0;
while i < bytes.len() {
if bytes[i] != b'#' {
return i;
}
while i < bytes.len() && bytes[i] != b'\n' {
i += 1;
}
i += 1;
}
bytes.len()
}
fn base_form_candidates(word: &str) -> Vec<String> {
if word.len() < 3 || !word.bytes().all(|b| b.is_ascii_lowercase()) {
return Vec::new();
}
let mut out: Vec<String> = Vec::new();
let mut push = |s: String| {
if s.len() >= 2 && !out.contains(&s) {
out.push(s);
}
};
if let Some(stem) = word.strip_suffix("ies") {
push(format!("{stem}y"));
}
let es_stem = word.strip_suffix("es");
let s_stem = word.strip_suffix('s').filter(|s| !s.ends_with('s'));
let (first, second) = if es_stem.is_some_and(ends_in_sibilant) {
(es_stem, s_stem)
} else {
(s_stem, es_stem)
};
for stem in [first, second].into_iter().flatten() {
push(stem.to_string());
}
if let Some(stem) = word.strip_suffix("ing") {
push(stem.to_string());
push(format!("{stem}e"));
if let Some(undoubled) = undouble(stem) {
push(undoubled);
}
}
if let Some(stem) = word.strip_suffix("ed") {
push(stem.to_string());
push(format!("{stem}e"));
if let Some(undoubled) = undouble(stem) {
push(undoubled);
}
}
out
}
fn ends_in_sibilant(stem: &str) -> bool {
stem.ends_with("ss")
|| stem.ends_with('x')
|| stem.ends_with('z')
|| stem.ends_with("ch")
|| stem.ends_with("sh")
}
fn undouble(stem: &str) -> Option<String> {
let mut chars = stem.chars().rev();
let last = chars.next()?;
if last != chars.next()? {
return None;
}
Some(stem[..stem.len() - last.len_utf8()].to_string())
}
#[derive(Debug, Clone, Copy)]
pub struct WordNetPos {
table: &'static str,
data_start: usize,
}
impl Default for WordNetPos {
fn default() -> Self {
Self::shipped()
}
}
impl WordNetPos {
pub const fn shipped() -> Self {
Self {
table: TABLE,
data_start: TABLE_DATA_START,
}
}
pub fn from_table(table: &'static str) -> Self {
Self {
table,
data_start: data_start(table.as_bytes()),
}
}
pub fn mask(&self, word: &str) -> u8 {
if let Some(m) = self.lookup(word.as_bytes()) {
return m;
}
if word.chars().any(char::is_uppercase) {
let lowered = word.to_lowercase();
if let Some(m) = self.lookup(lowered.as_bytes()) {
return m;
}
return self.inflected_mask(&lowered);
}
self.inflected_mask(word)
}
fn inflected_mask(&self, word: &str) -> u8 {
for candidate in base_form_candidates(word) {
if let Some(m) = self.lookup(candidate.as_bytes()) {
return m;
}
}
0
}
fn lookup(&self, needle: &[u8]) -> Option<u8> {
let bytes = self.table.as_bytes();
let mut lo = self.data_start;
let mut hi = bytes.len();
while lo < hi {
let mut start = lo + (hi - lo) / 2;
while start > lo && bytes[start - 1] != b'\n' {
start -= 1;
}
let mut end = start;
while end < hi && bytes[end] != b'\n' {
end += 1;
}
let line = &bytes[start..end];
let tab = line.iter().position(|b| *b == b'\t')?;
match line[..tab].cmp(needle) {
Ordering::Less => lo = end + 1,
Ordering::Greater => hi = start,
Ordering::Equal => {
return std::str::from_utf8(&line[tab + 1..])
.ok()?
.trim()
.parse::<u8>()
.ok();
}
}
}
None
}
pub fn is_known(&self, word: &str) -> bool {
self.mask(word) != 0
}
pub fn is_noun(&self, word: &str) -> bool {
self.mask(word) & NOUN != 0
}
pub fn is_adjective_only(&self, word: &str) -> bool {
let m = self.mask(word);
m & ADJ != 0 && m & NOUN == 0
}
pub fn lemma_count(&self) -> usize {
self.table[self.data_start..]
.lines()
.filter(|l| !l.is_empty())
.count()
}
}
#[cfg(test)]
mod tests {
use super::*;
const TINY: &str = "# notice line\n# another\nalpha\t1\nbeta\t4\ndelta\t2\nomega\t15\n";
#[test]
fn data_start_skips_the_whole_header() {
assert_eq!(
&TINY[data_start(TINY.as_bytes())..],
"alpha\t1\nbeta\t4\ndelta\t2\nomega\t15\n"
);
}
#[test]
fn data_start_handles_no_header() {
assert_eq!(data_start(b"alpha\t1\n"), 0);
assert_eq!(data_start(b"# only header\n"), 14);
}
#[test]
fn lookup_finds_the_first_and_last_records() {
let wn = WordNetPos::from_table(TINY);
assert_eq!(wn.mask("alpha"), 1);
assert_eq!(wn.mask("beta"), 4);
assert_eq!(wn.mask("delta"), 2);
assert_eq!(wn.mask("omega"), 15);
}
#[test]
fn lookup_misses_outside_the_table_range() {
let wn = WordNetPos::from_table(TINY);
for w in ["aardvark", "zulu", "carrot", "epsilon", "alph"] {
assert_eq!(wn.mask(w), 0, "{w} should not be found");
}
}
#[test]
fn lookup_retries_an_inflected_form() {
let wn = WordNetPos::from_table(TINY);
assert_eq!(wn.mask("alphas"), wn.mask("alpha"));
assert_eq!(wn.mask("alph"), 0, "a prefix is still a miss");
}
#[test]
fn mask_resolves_regular_inflections_to_their_base_form() {
let wn = WordNetPos::shipped();
assert_eq!(wn.mask("containing"), wn.mask("contain"));
assert_eq!(
wn.mask("containing") & NOUN,
0,
"a participle is not a noun"
);
assert_eq!(wn.mask("parsing"), wn.mask("parse"));
assert_eq!(wn.mask("committing"), wn.mask("commit"));
assert_eq!(wn.mask("parsers"), wn.mask("parser"));
assert!(wn.is_noun("parsers"));
assert_eq!(wn.mask("libraries"), wn.mask("library"));
assert_eq!(wn.mask("indexed"), wn.mask("index"));
}
#[test]
fn mask_leaves_non_words_and_names_unknown() {
let wn = WordNetPos::shipped();
for w in [
"rustc",
"librs",
"tantivy",
"redb",
"trusty-memory",
"crates/trusty-search/src/allowlist/tests.rs",
"budget_tokens",
] {
assert_eq!(wn.mask(w), 0, "{w} must stay unknown");
assert!(!wn.is_adjective_only(w), "{w} must fail open");
}
}
#[test]
fn base_forms_cover_the_regular_inflections() {
assert!(base_form_candidates("parsers").contains(&"parser".to_string()));
assert!(base_form_candidates("libraries").contains(&"library".to_string()));
assert!(base_form_candidates("boxes").contains(&"box".to_string()));
assert!(base_form_candidates("containing").contains(&"contain".to_string()));
assert!(base_form_candidates("parsing").contains(&"parse".to_string()));
assert!(base_form_candidates("stopping").contains(&"stop".to_string()));
assert!(base_form_candidates("indexed").contains(&"index".to_string()));
assert!(!base_form_candidates("class").contains(&"clas".to_string()));
}
#[test]
fn es_and_s_are_ordered_by_the_sibilant_rule() {
let wn = WordNetPos::shipped();
for (inflected, base) in [
("notes", "note"),
("sites", "site"),
("writes", "write"),
("rides", "ride"),
("envelopes", "envelope"),
("uses", "use"),
("houses", "house"),
("releases", "release"),
("cases", "case"),
] {
assert_eq!(
wn.mask(inflected),
wn.mask(base),
"{inflected} must resolve to {base}"
);
}
for (inflected, base) in [
("attaches", "attach"),
("passes", "pass"),
("boxes", "box"),
("dishes", "dish"),
("matches", "match"),
("indexes", "index"),
("classes", "class"),
("buses", "bus"),
] {
assert_eq!(
wn.mask(inflected),
wn.mask(base),
"{inflected} must resolve to {base}"
);
}
assert_ne!(wn.mask("not"), wn.mask("note"));
assert_ne!(wn.mask("attach"), wn.mask("attache"));
}
#[test]
fn base_forms_skip_non_words() {
for w in ["trusty-memory", "budget_tokens", "src/main.rs", "c#", "ab"] {
assert!(
base_form_candidates(w).is_empty(),
"{w} should generate no candidates"
);
}
}
#[test]
fn lookup_tolerates_a_bad_line() {
let wn = WordNetPos::from_table("alpha\t1\nbroken-line\nomega\t15\n");
assert_eq!(wn.mask("nonsense"), 0);
}
#[test]
fn shipped_table_answers_the_four_pos_classes() {
let wn = WordNetPos::shipped();
assert_eq!(wn.lemma_count(), 83_253);
assert!(wn.is_noun("compiler"));
assert!(wn.mask("run") & VERB != 0);
assert!(wn.mask("hard") & ADJ != 0);
assert!(wn.mask("quickly") & ADV != 0);
}
#[test]
fn the_shipped_table_is_sorted_and_parseable() {
let wn = WordNetPos::shipped();
let mut prev: &str = "";
let mut n = 0usize;
for line in wn.table[wn.data_start..].lines() {
assert!(
!line.is_empty(),
"blank data line after {prev:?} — it aborts any lookup that bisects onto it"
);
let (lemma, mask) = line.split_once('\t').expect("every data line has a tab");
assert!(
lemma.as_bytes() > prev.as_bytes(),
"table out of order at {lemma:?} (after {prev:?}) — binary search is invalid"
);
assert!(
!lemma.contains('_'),
"multi-word lemma {lemma:?} is dead weight"
);
let m: u8 = mask.parse().expect("mask parses");
assert!(
m > 0 && m <= (NOUN | VERB | ADJ | ADV),
"bad mask {m} for {lemma:?}"
);
prev = lemma;
n += 1;
}
assert_eq!(n, wn.lemma_count());
}
#[test]
fn multiword_lemmas_are_absent() {
let wn = WordNetPos::shipped();
assert_eq!(wn.mask("hot_dog"), 0);
assert!(wn.is_noun("dog"));
}
#[test]
fn mask_returns_zero_for_unknown_words() {
let wn = WordNetPos::shipped();
for w in ["rustc", "librs", "tantivy", "redb", "trusty-memory"] {
assert_eq!(wn.mask(w), 0, "{w} should be unknown to WordNet");
assert!(!wn.is_adjective_only(w), "{w} must fail open");
}
}
#[test]
fn mask_reports_every_pos_for_a_four_way_lemma() {
let wn = WordNetPos::shipped();
assert_eq!(wn.mask("fast"), NOUN | VERB | ADJ | ADV);
}
#[test]
fn adjective_only_catches_hard_and_spares_fast() {
let wn = WordNetPos::shipped();
assert!(wn.is_adjective_only("hard"));
assert!(!wn.is_adjective_only("fast"));
assert!(!wn.is_adjective_only("parser"));
}
#[test]
fn mask_is_case_insensitive() {
let wn = WordNetPos::shipped();
assert_eq!(wn.mask("Compiler"), wn.mask("compiler"));
assert_eq!(wn.mask("HARD"), wn.mask("hard"));
}
#[test]
fn independent_handles_agree() {
let a = WordNetPos::shipped();
let b = WordNetPos::default();
for w in ["compiler", "hard", "fast", "unknownium"] {
assert_eq!(a.mask(w), b.mask(w));
}
}
}