use std::collections::{HashMap, HashSet, VecDeque};
use rusqlite::{params, Connection};
use crate::error::Result;
pub fn tokenize(text: &str) -> Vec<String> {
text.split(|c: char| !c.is_alphanumeric())
.filter(|s| !s.is_empty())
.map(|s| s.to_lowercase())
.collect()
}
#[cfg(test)]
mod c5a_tests {
use super::*;
#[test]
fn possessive_resolves_to_the_bare_entity() {
let toks = tokenize("What is Taylor's role?");
assert!(toks.contains(&"taylor".to_string()), "{toks:?}");
assert_eq!(
tokenize("What is Taylor's role?")
.iter()
.filter(|t| *t == "taylor")
.count(),
tokenize("What is Taylor s role?")
.iter()
.filter(|t| *t == "taylor")
.count(),
"possessive and plain forms must tokenize alike"
);
}
#[test]
fn apostrophe_names_match_symmetrically() {
let text_tokens = tokenize("A meeting with O'Brien about the launch");
assert!(entity_matches_text("O'Brien", &text_tokens));
assert!(entity_matches_text("o brien", &text_tokens));
}
#[test]
fn contractions_stop_being_coherent_tokens() {
assert_eq!(tokenize("Don't"), vec!["don", "t"]);
}
}
pub fn entity_matches_text(entity: &str, text_tokens: &[String]) -> bool {
let entity_tokens = tokenize(entity);
if entity_tokens.is_empty() {
return false;
}
if entity_tokens.len() == 1 {
text_tokens.iter().any(|t| t == &entity_tokens[0])
} else {
text_tokens
.windows(entity_tokens.len())
.any(|window| window.iter().zip(entity_tokens.iter()).all(|(w, e)| w == e))
}
}
const ENTITY_STOPWORDS: &[&str] = &[
"The",
"A",
"An",
"I",
"We",
"You",
"He",
"She",
"It",
"They",
"This",
"That",
"These",
"Those",
"My",
"Your",
"His",
"Her",
"Its",
"Our",
"Their",
"But",
"And",
"Or",
"So",
"If",
"When",
"Where",
"What",
"Who",
"Why",
"How",
"Is",
"Are",
"Was",
"Were",
"Be",
"Been",
"Being",
"Have",
"Has",
"Had",
"Do",
"Does",
"Did",
"Of",
"In",
"On",
"At",
"To",
"For",
"With",
"From",
"By",
"As",
"Than",
"Then",
"Also",
"Just",
"Only",
"Very",
"Much",
"Not",
"No",
"Nor",
"Most",
"More",
"Less",
"Some",
"Any",
"All",
"Each",
"Every",
"Both",
"Such",
"Same",
"Other",
"Another",
"Yet",
"Still",
"Because",
"While",
"After",
"Before",
"During",
"Since",
"Until",
"Between",
"Through",
"About",
"Into",
"Over",
"Under",
"Again",
"Once",
"Here",
"There",
"Now",
"Thus",
"However",
"Therefore",
"Note",
"See",
"Can",
"Could",
"Will",
"Would",
"Should",
"May",
"Might",
"Must",
"Let",
"Get",
"Got",
];
const AMBIGUOUS_COMMON_ENTITIES: &[&str] = &[
"January",
"February",
"March",
"April",
"May",
"June",
"July",
"August",
"September",
"October",
"November",
"December",
"Monday",
"Tuesday",
"Wednesday",
"Thursday",
"Friday",
"Saturday",
"Sunday",
];
fn is_entity_stopword(tok: &str) -> bool {
ENTITY_STOPWORDS.iter().any(|s| s.eq_ignore_ascii_case(tok))
|| AMBIGUOUS_COMMON_ENTITIES
.iter()
.any(|s| s.eq_ignore_ascii_case(tok))
}
const MAX_ENTITY_TOKENS: usize = 6;
const MAX_ALLCAPS_TOKENS: usize = 2;
fn is_all_caps_token(tok: &str) -> bool {
tok.chars().any(|c| c.is_alphabetic())
&& tok.chars().all(|c| !c.is_alphabetic() || c.is_uppercase())
}
fn is_prose_run(chunk: &[String]) -> bool {
if chunk.len() > MAX_ENTITY_TOKENS {
return true;
}
chunk.iter().filter(|t| is_all_caps_token(t)).count() > MAX_ALLCAPS_TOKENS
}
pub fn is_rejected_entity_name(name: &str) -> bool {
let toks: Vec<String> = name.split_whitespace().map(|s| s.to_string()).collect();
if toks.is_empty() {
return true;
}
if !name.chars().any(|c| c.is_alphabetic()) {
return true;
}
if toks.iter().all(|t| is_entity_stopword(t)) {
return true;
}
is_prose_run(&toks)
}
fn strip_code(text: &str) -> std::borrow::Cow<'_, str> {
if !text.contains('`') {
return std::borrow::Cow::Borrowed(text);
}
let mut out = String::with_capacity(text.len());
let mut rest = text;
while let Some(t) = rest.find('`') {
out.push_str(&rest[..t]);
out.push(' ');
let after = &rest[t..];
if let Some(body) = after.strip_prefix("```") {
match body.find("```") {
Some(end) => rest = &body[end + 3..],
None => return std::borrow::Cow::Owned(out), }
} else {
let body = &after[1..];
match body.find('`') {
Some(end) => rest = &body[end + 1..],
None => {
out.push_str(body);
return std::borrow::Cow::Owned(out);
}
}
}
}
out.push_str(rest);
std::borrow::Cow::Owned(out)
}
pub const COMMON_WORD_SEED: &[&str] = &[
"about",
"above",
"actually",
"add",
"added",
"adding",
"after",
"again",
"against",
"ago",
"all",
"almost",
"already",
"also",
"although",
"always",
"another",
"anyway",
"apparently",
"around",
"ask",
"asked",
"back",
"basically",
"because",
"before",
"began",
"begin",
"behind",
"below",
"besides",
"better",
"between",
"big",
"both",
"bring",
"build",
"builder",
"built",
"call",
"called",
"came",
"can",
"cannot",
"capability",
"certainly",
"change",
"changed",
"check",
"checked",
"clearly",
"close",
"closed",
"code",
"come",
"coming",
"common",
"compare",
"consider",
"critically",
"current",
"currently",
"day",
"days",
"decide",
"decided",
"default",
"definitely",
"delete",
"deleted",
"did",
"different",
"do",
"does",
"doing",
"done",
"down",
"during",
"each",
"early",
"easy",
"efficient",
"either",
"else",
"end",
"enough",
"especially",
"even",
"eventually",
"ever",
"every",
"everything",
"exactly",
"example",
"except",
"expected",
"fail",
"failed",
"failing",
"fails",
"far",
"fast",
"few",
"final",
"finally",
"find",
"first",
"fix",
"fixed",
"fixing",
"follow",
"following",
"found",
"from",
"full",
"further",
"general",
"generally",
"get",
"gets",
"getting",
"give",
"given",
"go",
"going",
"good",
"got",
"great",
"had",
"happens",
"hard",
"has",
"have",
"having",
"hence",
"here",
"high",
"hopefully",
"how",
"however",
"idea",
"ideally",
"idempotent",
"if",
"important",
"instead",
"into",
"issue",
"just",
"keep",
"key",
"kind",
"large",
"last",
"later",
"least",
"less",
"let",
"lets",
"like",
"likely",
"line",
"link",
"linked",
"little",
"long",
"look",
"looked",
"looking",
"low",
"made",
"main",
"make",
"makes",
"making",
"many",
"may",
"maybe",
"mean",
"means",
"meanwhile",
"might",
"more",
"moreover",
"most",
"mostly",
"move",
"moved",
"much",
"must",
"near",
"need",
"needed",
"needs",
"never",
"new",
"next",
"nice",
"no",
"nope",
"normally",
"not",
"note",
"nothing",
"now",
"obviously",
"of",
"off",
"often",
"ok",
"okay",
"old",
"on",
"once",
"one",
"only",
"open",
"opened",
"option",
"or",
"other",
"otherwise",
"our",
"out",
"over",
"overall",
"own",
"part",
"pass",
"passed",
"past",
"per",
"perhaps",
"plan",
"please",
"point",
"possible",
"possibly",
"previous",
"previously",
"probably",
"problem",
"put",
"quick",
"quickly",
"quite",
"rather",
"ready",
"real",
"really",
"reason",
"recent",
"recently",
"remove",
"removed",
"result",
"results",
"right",
"run",
"running",
"runs",
"said",
"same",
"saw",
"say",
"says",
"second",
"see",
"seems",
"seen",
"set",
"several",
"should",
"show",
"shows",
"similar",
"simple",
"simply",
"since",
"small",
"so",
"some",
"something",
"sometimes",
"soon",
"start",
"started",
"starting",
"still",
"stop",
"stopped",
"such",
"sure",
"take",
"taken",
"target",
"test",
"tested",
"testing",
"tests",
"than",
"that",
"then",
"there",
"therefore",
"these",
"thing",
"things",
"think",
"this",
"those",
"though",
"three",
"through",
"thus",
"time",
"today",
"together",
"tomorrow",
"tonight",
"too",
"took",
"total",
"tried",
"true",
"try",
"trying",
"turn",
"two",
"under",
"unfortunately",
"unless",
"until",
"up",
"update",
"updated",
"upon",
"use",
"used",
"using",
"usually",
"very",
"want",
"wanted",
"way",
"weaker",
"well",
"went",
"were",
"what",
"whatever",
"when",
"whenever",
"where",
"whether",
"which",
"while",
"why",
"will",
"with",
"within",
"without",
"work",
"worked",
"working",
"works",
"would",
"wrong",
"yes",
"yesterday",
"yet",
"you",
"your",
"false",
"nil",
"none",
"null",
"undefined",
];
pub const COMMON_WORD_MIN_LOWER: i64 = 3;
pub const COMMON_WORD_LOWER_RATIO: i64 = 2;
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub struct CaseStats {
pub lower_n: i64,
pub cap_mid_n: i64,
pub cap_start_n: i64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum TokenCase {
Lower,
CapStart,
CapMid,
}
pub fn token_case_observations(text: &str) -> Vec<(String, TokenCase)> {
let stripped = strip_code(text);
let mut seen: std::collections::HashSet<(String, TokenCase)> = std::collections::HashSet::new();
let mut out = Vec::new();
for segment in stripped.split(|c: char| {
matches!(
c,
'.' | '!'
| '?'
| ':'
| ';'
| '\n'
| '('
| '['
| '"'
| '\u{201c}'
| '\u{201d}'
| '\u{2014}'
| '\u{2013}'
| '*'
| '\u{2022}'
| '>'
| '|'
)
}) {
let mut first = true;
for word in segment
.split(|c: char| !c.is_alphanumeric() && c != '\'')
.filter(|s| !s.is_empty())
{
let word = word.trim_matches('\'');
if word.chars().count() < 2 || !word.chars().all(|c| c.is_alphabetic() || c == '\'') {
if !word.is_empty() {
first = false;
}
continue;
}
let class = if word.chars().next().is_some_and(|c| c.is_uppercase()) {
if first {
TokenCase::CapStart
} else {
TokenCase::CapMid
}
} else {
TokenCase::Lower
};
first = false;
let key = (word.to_lowercase(), class);
if seen.insert(key.clone()) {
out.push(key);
}
}
}
out
}
pub fn is_common_word(token: &str, stats: Option<CaseStats>) -> bool {
if let Some(s) = stats {
if s.cap_mid_n >= COMMON_WORD_MIN_LOWER
&& s.cap_mid_n > s.lower_n
&& s.cap_mid_n >= s.cap_start_n
{
return false;
}
if s.lower_n >= COMMON_WORD_MIN_LOWER && s.lower_n >= COMMON_WORD_LOWER_RATIO * s.cap_mid_n
{
return true;
}
}
let lower = token.to_lowercase();
COMMON_WORD_SEED.contains(&lower.as_str())
}
pub fn admit_entity_with<F>(name: &str, lookup: F) -> bool
where
F: Fn(&str) -> Option<CaseStats>,
{
if !admit_entity(name) {
return false;
}
let toks: Vec<&str> = name.split_whitespace().collect();
if toks.len() != 1 {
return true;
}
let tok = toks[0];
!is_common_word(tok, lookup(&tok.to_lowercase()))
}
pub const ENTITY_MAX_CHARS: usize = 40;
pub const ENTITY_MAX_WORDS: usize = 4;
pub const ACRONYM_MAX_CHARS: usize = 6;
pub const ACRONYM_RUN_TOKEN_MAX_CHARS: usize = 5;
pub fn admit_entity(name: &str) -> bool {
let name = name.trim();
if name.ends_with("'s") || name.ends_with("\u{2019}s") || name.ends_with('\'') {
return false;
}
if is_rejected_entity_name(name) {
return false;
}
let mut toks: Vec<&str> = name.split_whitespace().collect();
while toks.first().is_some_and(|t| is_entity_stopword(t)) {
toks.remove(0);
}
while toks
.last()
.is_some_and(|t| is_entity_stopword(t) && t.chars().count() > 1)
{
toks.pop();
}
if toks.is_empty() || !toks.iter().any(|t| t.chars().any(|c| c.is_alphabetic())) {
return false;
}
if toks.len() > ENTITY_MAX_WORDS || name.chars().count() >= ENTITY_MAX_CHARS {
return false;
}
let caps: Vec<bool> = toks.iter().map(|t| is_all_caps_token(t)).collect();
if toks.len() == 1 && caps[0] && toks[0].chars().count() > ACRONYM_MAX_CHARS {
return false;
}
if toks.len() > 1
&& caps.iter().all(|&c| c)
&& toks
.iter()
.any(|t| t.chars().count() > ACRONYM_RUN_TOKEN_MAX_CHARS)
{
return false;
}
true
}
pub const VALUE_OBJECT_RELS: &[&str] = &["runs", "born_in", "founded_in", "released"];
pub fn relation_admits_value_object(rel_type: &str, dst: &str) -> bool {
!is_value_object(dst) || VALUE_OBJECT_RELS.contains(&rel_type)
}
pub fn is_value_object(name: &str) -> bool {
let name = name.trim();
if name.is_empty() || !name.chars().any(|c| c.is_ascii_digit()) {
return false;
}
let mut prev_sep = true;
for c in name.chars() {
if c.is_ascii_digit() {
prev_sep = false;
} else if (c == '.' || c == '-') && !prev_sep {
prev_sep = true;
} else {
return false;
}
}
!prev_sep
}
pub fn extract_value_candidates(text: &str) -> Vec<String> {
let stripped = strip_code(text);
let mut out: Vec<String> = Vec::new();
for word in stripped
.split(|c: char| {
c.is_whitespace() || matches!(c, ',' | ';' | ':' | '(' | ')' | '[' | ']' | '"' | '\'')
})
.filter(|s| !s.is_empty())
{
let w = word.trim_end_matches(|c: char| c == '.' || c == '!' || c == '?');
if !is_value_object(w) {
continue;
}
if !out.iter().any(|o| o == w) {
out.push(w.to_string());
}
}
out
}
pub fn extract_heuristic_entities(text: &str) -> Vec<String> {
extract_heuristic_entities_with(text, |_| None)
}
pub fn extract_heuristic_entities_with<F>(text: &str, lookup: F) -> Vec<String>
where
F: Fn(&str) -> Option<CaseStats>,
{
let stripped = strip_code(text);
extract_heuristic_entities_inner(stripped.as_ref(), &lookup)
}
fn extract_heuristic_entities_inner(
text: &str,
lookup: &dyn Fn(&str) -> Option<CaseStats>,
) -> Vec<String> {
let mut entities: Vec<String> = Vec::new();
for segment in text.split(|c: char| {
matches!(
c,
':' | ';' | ',' | '!' | '?' | '\n' | '(' | ')' | '[' | ']' | '"'
)
}) {
extract_entities_from_segment(segment, &mut entities, lookup);
}
let mut seen: std::collections::HashSet<String> = std::collections::HashSet::new();
entities.retain(|e| seen.insert(e.clone()));
entities
}
fn extract_entities_from_segment(
text: &str,
entities: &mut Vec<String>,
lookup: &dyn Fn(&str) -> Option<CaseStats>,
) {
let mut chunk: Vec<String> = Vec::new();
let flush = |chunk: &mut Vec<String>, out: &mut Vec<String>| {
while !chunk.is_empty() && is_entity_stopword(&chunk[0]) {
chunk.remove(0);
}
while let Some(last) = chunk.last() {
if is_entity_stopword(last) && last.chars().count() > 1 {
chunk.pop();
} else {
break;
}
}
if !chunk.is_empty() && !is_prose_run(chunk) {
let candidate = chunk.join(" ");
let alpha_chars = candidate.chars().filter(|c| c.is_alphanumeric()).count();
if alpha_chars >= 2 && admit_entity_with(&candidate, lookup) {
out.push(candidate);
}
}
chunk.clear();
};
for word in text
.split(|c: char| !c.is_alphanumeric() && c != '\'')
.filter(|s| !s.is_empty())
{
let possessive = word
.strip_suffix("'s")
.or_else(|| word.strip_suffix("'S"))
.or_else(|| word.strip_suffix('\''))
.filter(|bare| !bare.is_empty());
let entity_word = possessive.unwrap_or(word);
if !entity_word.chars().any(|c| c.is_alphabetic()) {
flush(&mut chunk, entities);
continue;
}
let first = entity_word.chars().next().unwrap();
let starts_upper = first.is_uppercase();
let is_all_caps = entity_word.len() > 1
&& entity_word
.chars()
.all(|c| !c.is_alphabetic() || c.is_uppercase());
let joins_chunk = if chunk.is_empty() {
starts_upper || is_all_caps
} else {
starts_upper || is_all_caps || (entity_word.len() == 1 && first.is_ascii_uppercase())
};
if joins_chunk {
chunk.push(entity_word.to_string());
if possessive.is_some() {
flush(&mut chunk, entities);
}
} else {
flush(&mut chunk, entities);
}
}
flush(&mut chunk, entities);
}
#[derive(Debug, Clone)]
pub struct RelationCandidate {
pub src: String,
pub rel_type: String,
pub dst: String,
pub polarity: i32, pub modality: String, pub confidence_band: String, }
const RELATION_PATTERNS: &[(&[&str], &str)] = &[
(
&["is the ceo of", "is ceo of", "serves as ceo of"],
"ceo_of",
),
(
&["is the cto of", "is cto of", "serves as cto of"],
"cto_of",
),
(
&["is the cfo of", "is cfo of", "serves as cfo of"],
"cfo_of",
),
(
&["is the founder of", "is founder of", "co-founded"],
"founded",
),
(&["founded"], "founded"),
(&["leads", "heads", "manages", "directs"], "leads"),
(&["runs", "is running", "now runs"], "runs"),
(
&[
"works at",
"works for",
"employed at",
"employed by",
"joined",
],
"works_at",
),
(&["was born in", "born in"], "born_in"),
(
&[
"is headquartered in",
"headquartered in",
"is based in",
"based in",
"located in",
],
"headquartered_in",
),
(&["is married to", "married to", "wed to"], "married_to"),
(
&["acquired", "bought", "purchased", "took over"],
"acquired",
),
(
&[
"is a subsidiary of",
"subsidiary of",
"is owned by",
"owned by",
],
"subsidiary_of",
),
(&["speaks", "is fluent in"], "speaks"),
(
&["is a member of", "member of", "belongs to", "part of"],
"member_of",
),
(&["reports to"], "reports_to"),
];
const REVERSE_ROLE_PATTERNS: &[(&str, &str)] = &[
("ceo", "ceo_of"),
("cto", "cto_of"),
("cfo", "cfo_of"),
("founder", "founded"),
("president", "leads"),
("director", "leads"),
("head", "leads"),
];
const ANCHORED_RELATION_PATTERNS: &[(&[&str], &str)] = &[
(
&[
"lives in",
"live in",
"living in",
"now lives in",
"resides in",
"reside in",
"residing in",
"moved to",
"has moved to",
"relocated to",
"has relocated to",
],
"lives_in",
),
(
&["hometown is", "grew up in", "originally from"],
"hometown",
),
];
fn contains_word_phrase(hay: &str, phrase: &str) -> bool {
if phrase.is_empty() {
return false;
}
let mut start = 0;
while let Some(idx) = hay[start..].find(phrase) {
let at = start + idx;
let end = at + phrase.len();
if boundary_before(hay, at) && boundary_after(hay, end) {
return true;
}
start = at + hay[at..].chars().next().map_or(1, char::len_utf8);
}
false
}
fn ends_with_word_phrase(hay: &str, phrase: &str) -> bool {
let hay = hay.trim_end();
if phrase.is_empty() || !hay.ends_with(phrase) {
return false;
}
boundary_before(hay, hay.len() - phrase.len())
}
fn contains_phrase_end_bounded(hay: &str, phrase: &str) -> bool {
if phrase.is_empty() {
return false;
}
let mut start = 0;
while let Some(idx) = hay[start..].find(phrase) {
let at = start + idx;
if boundary_after(hay, at + phrase.len()) {
return true;
}
start = at + hay[at..].chars().next().map_or(1, char::len_utf8);
}
false
}
fn strip_trailing_articles(window: &str) -> String {
let mut toks: Vec<&str> = window.split_whitespace().collect();
while matches!(toks.last().copied(), Some("the" | "a" | "an")) {
toks.pop();
}
toks.join(" ")
}
fn boundary_before(hay: &str, at: usize) -> bool {
at == 0
|| !hay[..at]
.chars()
.next_back()
.is_some_and(char::is_alphanumeric)
}
fn boundary_after(hay: &str, end: usize) -> bool {
end >= hay.len() || !hay[end..].chars().next().is_some_and(char::is_alphanumeric)
}
pub fn extract_learned_relations(
text: &str,
entities: &[String],
templates: &[(String, String)],
) -> Vec<RelationCandidate> {
if templates.is_empty() {
return vec![];
}
let mut candidates = Vec::new();
for w in between_windows(text, entities) {
if w.has_inner_entity {
continue;
}
for (phrase, rel_type) in templates {
if ends_with_word_phrase(&strip_trailing_articles(&w.between_stripped), phrase) {
candidates.push(RelationCandidate {
src: w.entity_a.to_string(),
rel_type: rel_type.clone(),
dst: w.entity_b.to_string(),
polarity: w.polarity,
modality: w.modality.to_string(),
confidence_band: "medium".to_string(),
});
break;
}
}
}
let mut seen = std::collections::HashSet::new();
candidates.retain(|c| seen.insert((c.src.clone(), c.rel_type.clone(), c.dst.clone())));
candidates
}
struct BetweenWindow<'a> {
entity_a: &'a str,
entity_b: &'a str,
between_stripped: String,
polarity: i32,
modality: &'static str,
has_inner_entity: bool,
}
fn between_windows<'a>(text: &str, entities: &'a [String]) -> Vec<BetweenWindow<'a>> {
let mut out = Vec::new();
if entities.len() < 2 {
return out;
}
let text_lower = text.to_lowercase();
let mut entity_positions: Vec<(usize, &str)> = Vec::new();
for entity in entities {
let entity_lower = entity.to_lowercase();
if let Some(pos) = text_lower.find(&entity_lower) {
entity_positions.push((pos, entity.as_str()));
}
}
entity_positions.sort_by_key(|(pos, _)| *pos);
for i in 0..entity_positions.len() {
for j in (i + 1)..entity_positions.len() {
let (pos_a, entity_a) = entity_positions[i];
let (pos_b, entity_b) = entity_positions[j];
if pos_b - pos_a > 150 {
continue;
}
let between_start = pos_a + entity_a.to_lowercase().len();
let between_end = pos_b;
if between_start >= between_end || between_end > text_lower.len() {
continue;
}
let between = text_lower[between_start..between_end].trim();
if between.is_empty() {
continue;
}
let has_negation = NEGATION_CUES
.iter()
.any(|cue| between.split_whitespace().any(|w| w == *cue));
let between_stripped: String = between
.split_whitespace()
.filter(|w| !NEGATION_CUES.contains(w))
.collect::<Vec<_>>()
.join(" ");
let modality = if MODALITY_CUES.iter().any(|cue| between.contains(cue)) {
"reported"
} else {
"asserted"
};
let has_inner_entity = entity_positions[i + 1..j]
.iter()
.any(|(p, e)| *p > pos_a && *p + e.len() <= pos_b);
out.push(BetweenWindow {
entity_a,
entity_b,
between_stripped,
polarity: if has_negation { -1 } else { 1 },
modality,
has_inner_entity,
});
}
}
out
}
pub fn extract_heuristic_relations(text: &str, entities: &[String]) -> Vec<RelationCandidate> {
if entities.len() < 2 {
return vec![];
}
let text_lower = text.to_lowercase();
let mut candidates: Vec<RelationCandidate> = Vec::new();
let mut entity_positions: Vec<(usize, &str)> = Vec::new();
for entity in entities {
let entity_lower = entity.to_lowercase();
if let Some(pos) = text_lower.find(&entity_lower) {
entity_positions.push((pos, entity.as_str()));
}
}
entity_positions.sort_by_key(|(pos, _)| *pos);
for i in 0..entity_positions.len() {
for j in (i + 1)..entity_positions.len() {
let (pos_a, entity_a) = entity_positions[i];
let (pos_b, entity_b) = entity_positions[j];
if pos_b - pos_a > 150 {
continue;
}
let between_start = pos_a + entity_a.to_lowercase().len();
let between_end = pos_b;
if between_start >= between_end || between_end > text_lower.len() {
continue;
}
let between = text_lower[between_start..between_end].trim();
if between.is_empty() {
continue;
}
let has_negation = NEGATION_CUES
.iter()
.any(|cue| between.split_whitespace().any(|w| w == *cue));
let polarity = if has_negation { -1 } else { 1 };
let between_stripped: String = between
.split_whitespace()
.filter(|w| !NEGATION_CUES.contains(w))
.collect::<Vec<_>>()
.join(" ");
let modality = if MODALITY_CUES.iter().any(|cue| between.contains(cue)) {
"reported"
} else {
"asserted"
};
let inner_entities: Vec<&str> = entity_positions[i + 1..j]
.iter()
.filter(|(p, e)| *p > pos_a && *p + e.len() <= pos_b)
.map(|(_, e)| *e)
.collect();
let window_anchored = strip_trailing_articles(&between_stripped);
for (patterns, rel_type) in RELATION_PATTERNS {
for pattern in *patterns {
let inner_ok = inner_entities
.iter()
.all(|e| pattern.contains(&e.to_lowercase()));
if inner_ok && ends_with_word_phrase(&window_anchored, pattern) {
candidates.push(RelationCandidate {
src: entity_a.to_string(),
rel_type: rel_type.to_string(),
dst: entity_b.to_string(),
polarity,
modality: modality.to_string(),
confidence_band: "medium".to_string(),
});
break; }
}
}
for (patterns, rel_type) in ANCHORED_RELATION_PATTERNS {
for pattern in *patterns {
if inner_entities.is_empty() && ends_with_word_phrase(&window_anchored, pattern)
{
candidates.push(RelationCandidate {
src: entity_a.to_string(),
rel_type: rel_type.to_string(),
dst: entity_b.to_string(),
polarity,
modality: modality.to_string(),
confidence_band: "medium".to_string(),
});
break;
}
}
}
for (role_keyword, rel_type) in REVERSE_ROLE_PATTERNS {
let possessive = format!("'s {}", role_keyword);
let possessive2 = format!("s {}", role_keyword);
if contains_phrase_end_bounded(&between_stripped, &possessive)
|| contains_phrase_end_bounded(&between_stripped, &possessive2)
{
candidates.push(RelationCandidate {
src: entity_b.to_string(), rel_type: rel_type.to_string(),
dst: entity_a.to_string(), polarity,
modality: modality.to_string(),
confidence_band: "medium".to_string(),
});
break;
}
}
}
}
let mut seen = std::collections::HashSet::new();
candidates.retain(|c| seen.insert((c.src.clone(), c.rel_type.clone(), c.dst.clone())));
candidates
}
const NEGATION_CUES: &[&str] = &[
"not", "no", "never", "denied", "refuted", "isn't", "wasn't", "aren't", "weren't", "doesn't",
"didn't", "disputes", "denies",
];
pub fn negation_cue(word: &str) -> bool {
NEGATION_CUES.contains(&word)
}
const TEMPORAL_CUES: &[&str] = &[
"was",
"were",
"until",
"before",
"after",
"since",
"during",
"former",
"current",
"currently",
"previously",
"recently",
"now",
"then",
"later",
"earlier",
"ago",
"yesterday",
"tomorrow",
];
const MODALITY_CUES: &[&str] = &[
"may",
"might",
"allegedly",
"reportedly",
"rumor",
"rumored",
"said",
"claims",
"according",
"stated",
"announced",
];
const COMPOUND_MARKERS: &[&str] = &[
"; ",
", then ",
", subsequently ",
" but ",
" however ",
" although ",
];
#[derive(Debug, Clone, Default)]
pub struct TextFeatures {
pub char_length: usize,
pub sentence_count: usize,
pub entity_count: usize,
pub negation_cue_count: usize,
pub temporal_cue_count: usize,
pub modality_cue_count: usize,
pub has_compound_markers: bool,
pub likely_assertion: bool,
}
pub fn analyze_text_features(text: &str, extracted_entities: &[String]) -> TextFeatures {
let lower = text.to_lowercase();
let tokens: Vec<&str> = text
.split(|c: char| !c.is_alphanumeric() && c != '\'')
.filter(|s| !s.is_empty())
.collect();
let tokens_lower: Vec<String> = tokens.iter().map(|t| t.to_lowercase()).collect();
let sentence_count = text
.chars()
.filter(|c| matches!(c, '.' | '!' | '?'))
.count()
.max(1);
let negation_cue_count = tokens_lower
.iter()
.filter(|t| NEGATION_CUES.contains(&t.as_str()))
.count();
let temporal_cue_count = tokens_lower
.iter()
.filter(|t| TEMPORAL_CUES.contains(&t.as_str()))
.count();
let modality_cue_count = tokens_lower
.iter()
.filter(|t| MODALITY_CUES.contains(&t.as_str()))
.count();
let has_compound_markers = COMPOUND_MARKERS.iter().any(|m| lower.contains(m));
let likely_assertion =
!text.trim_end().ends_with('?') && tokens.len() >= 2 && modality_cue_count == 0;
TextFeatures {
char_length: text.chars().count(),
sentence_count,
entity_count: extracted_entities.len(),
negation_cue_count,
temporal_cue_count,
modality_cue_count,
has_compound_markers,
likely_assertion,
}
}
const TECH_BLOCKLIST: &[&str] = &[
"faiss",
"onnx",
"scann",
"redis",
"kafka",
"docker",
"kubernetes",
"react",
"python",
"rust",
"java",
"swift",
"flutter",
"pytorch",
"tensorflow",
"numpy",
"pandas",
"spark",
"hadoop",
"nginx",
"postgres",
"mysql",
"sqlite",
"graphql",
"grpc",
"oauth",
"jwt",
"html",
"css",
"api",
"sdk",
"ml",
"ai",
"gpu",
"cpu",
"ram",
"ssd",
"aws",
"gcp",
"claude",
"openai",
"anthropic",
"gemini",
"llama",
"ollama",
];
const NON_PERSON_PREFIXES: &[&str] = &[
"project",
"team",
"company",
"group",
"department",
"org",
"the",
"operation",
"task",
"plan",
"system",
"service",
"app",
"tool",
"code",
"server",
"client",
"api",
"db",
"database",
"agent",
"model",
"version",
"release",
"build",
"deploy",
"config",
];
pub fn classify_entity_type(name: &str) -> &'static str {
let trimmed = name.trim();
if trimmed.is_empty() {
return "unknown";
}
let lower = trimmed.to_lowercase();
if TECH_BLOCKLIST.contains(&lower.as_str()) {
return "tech";
}
if trimmed.len() > 1
&& trimmed
.chars()
.all(|c| c.is_uppercase() || !c.is_alphabetic())
{
return "tech";
}
if trimmed.contains(' ') {
let words: Vec<&str> = trimmed.split_whitespace().collect();
if words.len() == 2
&& words
.iter()
.all(|w| w.chars().next().map(|c| c.is_uppercase()).unwrap_or(false))
{
let first_lower = words[0].to_lowercase();
if NON_PERSON_PREFIXES.contains(&first_lower.as_str()) {
return "unknown";
}
if words
.iter()
.any(|w| TECH_BLOCKLIST.contains(&w.to_lowercase().as_str()))
{
return "tech";
}
return "person";
}
}
"unknown"
}
const PERSON_PERSON_RELS: &[&str] = &[
"married_to",
"mother_of",
"father_of",
"daughter_of",
"son_of",
"sister_of",
"brother_of",
"sibling_of",
"parent_of",
"child_of",
"knows",
"friends_with",
"met",
"dating",
"engaged_to",
"mentors",
"mentored_by",
"reports_to",
"manages",
"colleagues",
"roommate",
"neighbor",
"called",
"texted",
"messaged",
"date_night",
];
const PLACE_DST_RELS: &[&str] = &[
"lives_in",
"born_in",
"grew_up_in",
"located_in",
"based_in",
"visited",
"moved_to",
"traveled_to",
"from",
];
const ORG_DST_RELS: &[&str] = &[
"works_at",
"works_for",
"employed_at",
"employed_by",
"studied_at",
"attended",
"enrolled_in",
"graduated_from",
"member_of",
"belongs_to",
"founded",
];
const TECH_DST_RELS: &[&str] = &[
"built_with",
"uses",
"depends_on",
"integrates",
"requires",
"written_in",
"coded_in",
"implemented_with",
"powered_by",
"runs_on",
"compiled_with",
];
const INFRA_DST_RELS: &[&str] = &[
"deployed_on",
"hosted_on",
"deployed_to",
"hosted_at",
"runs_on_infra",
"served_by",
];
const PERSON_PROJECT_RELS: &[&str] = &[
"works_on",
"contributes_to",
"maintains",
"leads",
"created",
"built",
"designed",
"architected",
"owns",
];
const PROJECT_PROJECT_RELS: &[&str] = &[
"depends_on_project",
"extends",
"forks",
"replaces",
"supersedes",
"derived_from",
];
const EVENT_DST_RELS: &[&str] = &[
"attended_event",
"participated_in",
"scheduled_for",
"presented_at",
"spoke_at",
];
const CONCEPT_DST_RELS: &[&str] = &[
"interested_in",
"studies",
"researches",
"specializes_in",
"expert_in",
"learning",
"teaches",
];
pub fn classify_with_relationship(
src: &str,
dst: &str,
rel_type: &str,
) -> (&'static str, &'static str) {
let rel_lower = rel_type.to_lowercase();
let rel = rel_lower.as_str();
if PERSON_PERSON_RELS.contains(&rel) {
return ("person", "person");
}
if PLACE_DST_RELS.contains(&rel) {
return ("person", "place");
}
if ORG_DST_RELS.contains(&rel) {
return ("person", "organization");
}
if TECH_DST_RELS.contains(&rel) {
let src_type = classify_entity_type(src);
return (
if src_type == "unknown" {
"project"
} else {
src_type
},
"tech",
);
}
if INFRA_DST_RELS.contains(&rel) {
let src_type = classify_entity_type(src);
return (
if src_type == "unknown" {
"project"
} else {
src_type
},
"infrastructure",
);
}
if PERSON_PROJECT_RELS.contains(&rel) {
return ("person", "project");
}
if PROJECT_PROJECT_RELS.contains(&rel) {
return ("project", "project");
}
if EVENT_DST_RELS.contains(&rel) {
return (classify_entity_type(src), "event");
}
if CONCEPT_DST_RELS.contains(&rel) {
return ("person", "concept");
}
(classify_entity_type(src), classify_entity_type(dst))
}
pub fn entities_for_memories(conn: &Connection, rids: &[&str]) -> Result<Vec<String>> {
if rids.is_empty() {
return Ok(vec![]);
}
let placeholders: String = (0..rids.len())
.map(|i| format!("?{}", i + 1))
.collect::<Vec<_>>()
.join(",");
let sql = format!(
"SELECT DISTINCT entity_name FROM memory_entities WHERE memory_rid IN ({placeholders})"
);
let mut stmt = conn.prepare(&sql)?;
let param_values: Vec<Box<dyn rusqlite::types::ToSql>> = rids
.iter()
.map(|r| Box::new(r.to_string()) as Box<dyn rusqlite::types::ToSql>)
.collect();
let params_ref: Vec<&dyn rusqlite::types::ToSql> =
param_values.iter().map(|p| p.as_ref()).collect();
let entities = stmt
.query_map(params_ref.as_slice(), |row| row.get(0))?
.collect::<std::result::Result<Vec<String>, _>>()?;
Ok(entities)
}
pub fn memories_for_entities(conn: &Connection, entity_names: &[&str]) -> Result<HashSet<String>> {
if entity_names.is_empty() {
return Ok(HashSet::new());
}
let placeholders: String = (0..entity_names.len())
.map(|i| format!("?{}", i + 1))
.collect::<Vec<_>>()
.join(",");
let sql = format!(
"SELECT DISTINCT memory_rid FROM memory_entities WHERE entity_name IN ({placeholders})"
);
let mut stmt = conn.prepare(&sql)?;
let param_values: Vec<Box<dyn rusqlite::types::ToSql>> = entity_names
.iter()
.map(|e| Box::new(e.to_string()) as Box<dyn rusqlite::types::ToSql>)
.collect();
let params_ref: Vec<&dyn rusqlite::types::ToSql> =
param_values.iter().map(|p| p.as_ref()).collect();
let rids = stmt
.query_map(params_ref.as_slice(), |row| row.get(0))?
.collect::<std::result::Result<HashSet<String>, _>>()?;
Ok(rids)
}
pub fn expand_entities_nhop(
conn: &Connection,
seeds: &[&str],
max_hops: u8,
max_entities: usize,
) -> Result<Vec<(String, u8, f64)>> {
let mut result: Vec<(String, u8, f64)> = Vec::new();
let mut visited: HashMap<String, (u8, f64)> = HashMap::new();
for s in seeds {
visited.insert(s.to_string(), (0, 1.0));
result.push((s.to_string(), 0, 1.0));
}
let mut frontier: VecDeque<(String, u8, f64)> =
seeds.iter().map(|s| (s.to_string(), 0u8, 1.0f64)).collect();
while let Some((entity, hops, weight)) = frontier.pop_front() {
if hops >= max_hops || result.len() >= max_entities {
break;
}
let mut stmt = conn.prepare(
"SELECT src, dst, weight FROM edges WHERE (src = ?1 OR dst = ?1) AND tombstoned = 0",
)?;
let neighbors: Vec<(String, f64)> = stmt
.query_map(params![entity], |row| {
let src: String = row.get(0)?;
let dst: String = row.get(1)?;
let w: f64 = row.get(2)?;
let neighbor = if src == entity { dst } else { src };
Ok((neighbor, w))
})?
.collect::<std::result::Result<Vec<_>, _>>()?;
for (neighbor, edge_weight) in neighbors {
if visited.contains_key(&neighbor) {
continue;
}
if result.len() >= max_entities {
break;
}
let cumulative = weight * edge_weight;
let next_hops = hops + 1;
visited.insert(neighbor.clone(), (next_hops, cumulative));
result.push((neighbor.clone(), next_hops, cumulative));
if next_hops < max_hops {
frontier.push_back((neighbor, next_hops, cumulative));
}
}
}
Ok(result)
}
pub fn graph_proximity(
conn: &Connection,
memory_rid: &str,
expanded_entities: &HashMap<String, (u8, f64)>,
) -> Result<f64> {
let mem_entities: Vec<String> = conn
.prepare("SELECT entity_name FROM memory_entities WHERE memory_rid = ?1")?
.query_map(params![memory_rid], |row| row.get(0))?
.collect::<std::result::Result<Vec<_>, _>>()?;
let mut max_proximity = 0.0f64;
for entity in &mem_entities {
if let Some(&(hops, weight)) = expanded_entities.get(entity) {
let prox = weight / f64::powf(2.0, hops as f64);
if prox > max_proximity {
max_proximity = prox;
}
}
}
Ok(max_proximity)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::YantrikDB;
#[test]
fn test_extract_heuristic_entities_basic_names() {
let got = extract_heuristic_entities("Alice Chen is the CEO of Acme Corp");
assert!(got.contains(&"Alice Chen".to_string()), "got: {:?}", got);
assert!(got.contains(&"Acme Corp".to_string()), "got: {:?}", got);
assert!(got.contains(&"CEO".to_string()), "got: {:?}", got);
}
#[test]
fn test_extract_heuristic_entities_strips_sentence_start() {
let got = extract_heuristic_entities("The database backend is PostgreSQL");
assert_eq!(got, vec!["PostgreSQL".to_string()]);
}
#[test]
fn test_extract_heuristic_entities_multi_word_place() {
let got = extract_heuristic_entities("Acme is headquartered in San Francisco");
assert!(got.contains(&"Acme".to_string()), "got: {:?}", got);
assert!(got.contains(&"San Francisco".to_string()), "got: {:?}", got);
}
#[test]
fn test_extract_heuristic_entities_single_letter_suffix() {
let got = extract_heuristic_entities("Series A funding was 20 million dollars");
assert!(got.contains(&"Series A".to_string()), "got: {:?}", got);
}
#[test]
fn test_extract_heuristic_entities_dedupe() {
let got = extract_heuristic_entities("Alice met Alice at the cafe");
let alice_count = got.iter().filter(|e| *e == "Alice").count();
assert_eq!(alice_count, 1);
}
#[test]
fn test_extract_heuristic_entities_empty_on_lowercase() {
let got = extract_heuristic_entities("the quick brown fox jumps over the lazy dog");
assert!(got.is_empty(), "got: {:?}", got);
}
#[test]
fn test_extract_relations_ceo_of() {
let entities = vec!["Alice Chen".to_string(), "Acme Corp".to_string()];
let rels = extract_heuristic_relations("Alice Chen is the CEO of Acme Corp", &entities);
assert_eq!(rels.len(), 1, "got: {:?}", rels);
assert_eq!(rels[0].src, "Alice Chen");
assert_eq!(rels[0].rel_type, "ceo_of");
assert_eq!(rels[0].dst, "Acme Corp");
assert_eq!(rels[0].polarity, 1);
}
#[test]
fn test_extract_relations_works_at() {
let entities = vec!["Bob".to_string(), "Google".to_string()];
let rels = extract_heuristic_relations("Bob works at Google as an engineer", &entities);
assert!(
rels.iter().any(|r| r.rel_type == "works_at"),
"got: {:?}",
rels
);
}
#[test]
fn test_extract_relations_headquartered() {
let entities = vec!["Acme".to_string(), "San Francisco".to_string()];
let rels = extract_heuristic_relations("Acme is headquartered in San Francisco", &entities);
assert!(
rels.iter().any(|r| r.rel_type == "headquartered_in"),
"got: {:?}",
rels
);
}
#[test]
fn test_extract_relations_negation_detected() {
let entities = vec!["Alice".to_string(), "Acme".to_string()];
let rels = extract_heuristic_relations("Alice is not the CEO of Acme", &entities);
assert_eq!(rels.len(), 1);
assert_eq!(rels[0].polarity, -1, "negation should set polarity to -1");
}
#[test]
fn test_extract_relations_no_match_unrelated() {
let entities = vec!["Alice".to_string(), "Bob".to_string()];
let rels = extract_heuristic_relations("Alice and Bob went for coffee", &entities);
assert!(
rels.is_empty(),
"should not extract relation from unrelated text, got: {:?}",
rels
);
}
#[test]
fn test_extract_relations_multiple_pairs() {
let entities = vec![
"Alice".to_string(),
"Acme".to_string(),
"San Francisco".to_string(),
];
let rels = extract_heuristic_relations(
"Alice is the CEO of Acme which is headquartered in San Francisco",
&entities,
);
assert!(
rels.len() >= 2,
"should find CEO + headquartered, got: {:?}",
rels
);
}
#[test]
fn test_extract_relations_lives_in_is_anchored_to_the_next_entity() {
let entities = vec![
"Pranab".to_string(),
"Berlin".to_string(),
"Maria".to_string(),
];
let rels = extract_heuristic_relations("Pranab lives in Berlin with Maria", &entities);
let lives: Vec<_> = rels.iter().filter(|r| r.rel_type == "lives_in").collect();
assert_eq!(lives.len(), 1, "got: {:?}", rels);
assert_eq!(lives[0].src, "Pranab");
assert_eq!(lives[0].dst, "Berlin");
assert_eq!(lives[0].polarity, 1);
}
#[test]
fn test_extract_relations_moved_to_shares_the_lives_in_key() {
let entities = vec!["Alice Moreau".to_string(), "Munich".to_string()];
let rels = extract_heuristic_relations("Alice Moreau moved to Munich last year", &entities);
assert_eq!(rels.len(), 1, "got: {:?}", rels);
assert_eq!(rels[0].rel_type, "lives_in");
assert_eq!(rels[0].dst, "Munich");
}
#[test]
fn test_extract_relations_lives_in_negation() {
let entities = vec!["Pranab".to_string(), "Berlin".to_string()];
let rels = extract_heuristic_relations("Pranab does not live in Berlin", &entities);
assert_eq!(rels.len(), 1, "got: {:?}", rels);
assert_eq!(rels[0].rel_type, "lives_in");
assert_eq!(rels[0].polarity, -1);
}
#[test]
fn test_extract_relations_hometown() {
let entities = vec!["Pranab".to_string(), "Kolkata".to_string()];
for text in ["Pranab's hometown is Kolkata", "Pranab grew up in Kolkata"] {
let rels = extract_heuristic_relations(text, &entities);
assert_eq!(rels.len(), 1, "{text}: {:?}", rels);
assert_eq!(rels[0].rel_type, "hometown", "{text}");
assert_eq!(rels[0].dst, "Kolkata", "{text}");
}
}
#[test]
fn test_extract_relations_headquartered_does_not_mint_reverse_leads() {
let entities = vec!["Fennwick Labs".to_string(), "Berlin".to_string()];
let rels =
extract_heuristic_relations("Fennwick Labs is headquartered in Berlin", &entities);
assert!(
rels.iter().all(|r| r.rel_type == "headquartered_in"),
"got: {:?}",
rels
);
assert_eq!(rels.len(), 1, "got: {:?}", rels);
}
#[test]
fn test_extract_relations_patterns_match_whole_words_only() {
let entities = vec!["Acme".to_string(), "Globex".to_string()];
let rels =
extract_heuristic_relations("Acme dismissed unfounded rumors about Globex", &entities);
assert!(
rels.is_empty(),
"'unfounded' must not mint founded, got: {:?}",
rels
);
let rels = extract_heuristic_relations("Acme founded Globex", &entities);
assert_eq!(rels.len(), 1, "got: {:?}", rels);
assert_eq!(rels[0].rel_type, "founded");
}
#[test]
fn test_extract_relations_possessive_role_still_matches() {
let entities = vec!["Acme".to_string(), "Alice".to_string()];
let rels = extract_heuristic_relations("Acme's CEO, Alice, spoke first", &entities);
assert!(
rels.iter()
.any(|r| r.rel_type == "ceo_of" && r.src == "Alice" && r.dst == "Acme"),
"got: {:?}",
rels
);
}
#[test]
fn test_extract_learned_relations_is_anchored_and_labelled() {
let entities = vec!["Dana".to_string(), "Priya".to_string(), "Acme".to_string()];
let templates = vec![("mentors".to_string(), "mentors".to_string())];
let rels = extract_learned_relations(
"Dana mentors Priya at Acme this quarter",
&entities,
&templates,
);
assert_eq!(rels.len(), 1, "got: {:?}", rels);
assert_eq!(
(
rels[0].src.as_str(),
rels[0].rel_type.as_str(),
rels[0].dst.as_str()
),
("Dana", "mentors", "Priya")
);
let rels = extract_learned_relations(
"Dana does not mentor Priya",
&entities,
&[("mentor".into(), "mentors".into())],
);
assert_eq!(rels.len(), 1);
assert_eq!(rels[0].polarity, -1);
assert!(extract_learned_relations("Dana mentors Priya", &entities, &[]).is_empty());
}
#[test]
fn test_extract_relations_runs_is_a_version_relation_not_leadership() {
let entities = vec!["CT128".to_string(), "Yantrikdb".to_string()];
let rels = extract_heuristic_relations("CT128 runs Yantrikdb in production", &entities);
assert_eq!(rels.len(), 1, "got: {:?}", rels);
assert_eq!(rels[0].rel_type, "runs");
assert!(rels.iter().all(|r| r.rel_type != "leads"));
let entities = vec!["Alice".to_string(), "Acme".to_string()];
let rels = extract_heuristic_relations("Alice leads Acme", &entities);
assert_eq!(rels[0].rel_type, "leads");
}
#[test]
fn test_forward_patterns_are_anchored_to_the_adjacent_pair() {
let entities = vec![
"Pranab".to_string(),
"Materializer".to_string(),
"UTC".to_string(),
];
let rels = extract_heuristic_relations(
"Pranab confirmed the Materializer runs the loop every tick at UTC midnight",
&entities,
);
assert!(
!rels.iter().any(|r| r.src == "Pranab" && r.dst == "UTC"),
"no claim may bridge Pranab and UTC across Materializer: {:?}",
rels
);
let entities = vec!["Alice".to_string(), "Acme".to_string()];
let rels = extract_heuristic_relations("Alice works at the Acme office", &entities);
assert!(rels.iter().any(|r| r.rel_type == "works_at"), "{:?}", rels);
let rels =
extract_heuristic_relations("Alice works at home and later visited Acme", &entities);
assert!(
rels.is_empty(),
"verb not adjacent to the object: {:?}",
rels
);
let entities = vec![
"Alice Chen".to_string(),
"CEO".to_string(),
"Acme Corp".to_string(),
];
let rels = extract_heuristic_relations("Alice Chen is the CEO of Acme Corp", &entities);
assert!(
rels.iter()
.any(|r| r.rel_type == "ceo_of" && r.src == "Alice Chen" && r.dst == "Acme Corp"),
"{:?}",
rels
);
}
#[test]
fn test_extract_relations_needs_two_entities() {
let entities = vec!["Alice".to_string()];
let rels = extract_heuristic_relations("Alice is the CEO", &entities);
assert!(
rels.is_empty(),
"cannot extract relation with only one entity"
);
}
#[test]
fn test_analyze_text_features_basic_assertion() {
let entities = vec!["Alice Chen".to_string(), "Acme Corp".to_string()];
let f = analyze_text_features("Alice Chen is the CEO of Acme Corp", &entities);
assert_eq!(f.entity_count, 2);
assert_eq!(f.negation_cue_count, 0);
assert_eq!(f.modality_cue_count, 0);
assert!(f.likely_assertion);
assert!(!f.has_compound_markers);
}
#[test]
fn test_analyze_text_features_negation() {
let f = analyze_text_features("Alice is not the CEO of Acme", &[]);
assert_eq!(f.negation_cue_count, 1);
}
#[test]
fn test_analyze_text_features_temporal() {
let f = analyze_text_features("Alice was previously the CEO before 2024", &[]);
assert!(f.temporal_cue_count >= 2, "got: {}", f.temporal_cue_count);
}
#[test]
fn test_analyze_text_features_modality_suppresses_assertion() {
let f = analyze_text_features("Alice may become CEO allegedly", &[]);
assert!(f.modality_cue_count >= 2);
assert!(!f.likely_assertion);
}
#[test]
fn test_analyze_text_features_compound() {
let f = analyze_text_features("Alice was CEO until 2024; then Bob took over", &[]);
assert!(f.has_compound_markers);
}
#[test]
fn test_analyze_text_features_question_not_assertion() {
let f = analyze_text_features("Who is the CEO of Acme?", &[]);
assert!(!f.likely_assertion);
}
#[test]
fn test_extract_heuristic_entities_distinct_people() {
let a = extract_heuristic_entities("Alice Chen is the CEO of Acme Corp");
let b = extract_heuristic_entities("Sarah Kim is the CTO of Acme Corp");
let a_set: std::collections::HashSet<_> = a.iter().collect();
let b_set: std::collections::HashSet<_> = b.iter().collect();
assert!(a_set.contains(&"Alice Chen".to_string()));
assert!(b_set.contains(&"Sarah Kim".to_string()));
assert!(!a_set.contains(&"Sarah Kim".to_string()));
assert!(!b_set.contains(&"Alice Chen".to_string()));
}
fn setup_db() -> YantrikDB {
let db = YantrikDB::new(":memory:", 4).unwrap();
db.relate("Alice", "Bob", "knows", 1.0).unwrap();
db.relate("Bob", "Charlie", "knows", 0.8).unwrap();
db.relate("Alice", "ProjectX", "works_on", 1.0).unwrap();
db.relate("Dave", "ProjectX", "works_on", 0.9).unwrap();
let emb = vec![1.0f32, 0.0, 0.0, 0.0];
let r1 = db
.record(
"Alice discussed the plan",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&emb,
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
let r2 = db
.record(
"Bob reviewed the code",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&emb,
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
let r3 = db
.record(
"Charlie deployed to production",
"episodic",
0.5,
0.0,
604800.0,
&serde_json::json!({}),
&emb,
"default",
0.8,
"general",
"user",
None,
)
.unwrap();
db.link_memory_entity(&r1, "Alice").unwrap();
db.link_memory_entity(&r1, "ProjectX").unwrap();
db.link_memory_entity(&r2, "Bob").unwrap();
db.link_memory_entity(&r3, "Charlie").unwrap();
db
}
#[test]
fn test_entities_for_memories() {
let db = setup_db();
let rid: String = db
.conn()
.query_row(
"SELECT rid FROM memories ORDER BY created_at LIMIT 1",
[],
|row| row.get(0),
)
.unwrap();
let entities = entities_for_memories(&*db.conn(), &[&rid]).unwrap();
assert!(entities.contains(&"Alice".to_string()));
assert!(entities.contains(&"ProjectX".to_string()));
}
#[test]
fn test_memories_for_entities() {
let db = setup_db();
let rids = memories_for_entities(&*db.conn(), &["Alice"]).unwrap();
assert_eq!(rids.len(), 1); }
#[test]
fn test_expand_1hop() {
let db = setup_db();
let expanded = expand_entities_nhop(&*db.conn(), &["Alice"], 1, 30).unwrap();
let names: HashSet<String> = expanded.iter().map(|(n, _, _)| n.clone()).collect();
assert!(names.contains("Alice"));
assert!(names.contains("Bob"));
assert!(names.contains("ProjectX"));
}
#[test]
fn test_expand_2hop() {
let db = setup_db();
let expanded = expand_entities_nhop(&*db.conn(), &["Alice"], 2, 30).unwrap();
let names: HashSet<String> = expanded.iter().map(|(n, _, _)| n.clone()).collect();
assert!(names.contains("Charlie"));
assert!(names.contains("Dave"));
}
#[test]
fn test_expand_budget_limit() {
let db = setup_db();
let expanded = expand_entities_nhop(&*db.conn(), &["Alice"], 2, 3).unwrap();
assert!(expanded.len() <= 3);
}
#[test]
fn test_no_tombstoned_edges() {
let db = setup_db();
db.conn()
.execute(
"UPDATE claims SET tombstoned = 1 WHERE src = 'Alice' AND dst = 'Bob'",
[],
)
.unwrap();
let expanded = expand_entities_nhop(&*db.conn(), &["Alice"], 1, 30).unwrap();
let names: HashSet<String> = expanded.iter().map(|(n, _, _)| n.clone()).collect();
assert!(!names.contains("Bob"));
assert!(names.contains("ProjectX"));
}
#[test]
fn test_graph_proximity_score() {
let db = setup_db();
let rid: String = db
.conn()
.query_row(
"SELECT rid FROM memories ORDER BY created_at LIMIT 1",
[],
|row| row.get(0),
)
.unwrap();
let mut expanded = HashMap::new();
expanded.insert("Alice".to_string(), (0u8, 1.0f64));
expanded.insert("ProjectX".to_string(), (1u8, 1.0f64));
let prox = graph_proximity(&*db.conn(), &rid, &expanded).unwrap();
assert!((prox - 1.0).abs() < 1e-10);
}
#[test]
fn test_tokenize_basic() {
let tokens = tokenize("What is Sarah working on?");
assert_eq!(tokens, vec!["what", "is", "sarah", "working", "on"]);
}
#[test]
fn test_tokenize_splits_apostrophes() {
let tokens = tokenize("daughter's school play");
assert_eq!(tokens, vec!["daughter", "s", "school", "play"]);
}
#[test]
fn test_entity_matches_single_word() {
let tokens = tokenize("Sarah discussed the plan with Mike");
assert!(entity_matches_text("Sarah", &tokens));
assert!(entity_matches_text("Mike", &tokens));
assert!(!entity_matches_text("Sara", &tokens)); }
#[test]
fn test_entity_matches_multi_word() {
let tokens = tokenize("The data pipeline crashed during migration");
assert!(entity_matches_text("data pipeline", &tokens));
assert!(!entity_matches_text("data migration", &tokens)); }
#[test]
fn test_entity_no_substring_false_positive() {
let tokens = tokenize("The database was updated successfully");
assert!(!entity_matches_text("data", &tokens));
}
#[test]
fn test_entity_matches_case_insensitive() {
let tokens = tokenize("We evaluated FAISS for vector search");
assert!(entity_matches_text("FAISS", &tokens));
assert!(entity_matches_text("faiss", &tokens));
}
#[test]
fn test_classify_name_only_ambiguous() {
assert_eq!(classify_entity_type("Sarah"), "unknown");
assert_eq!(classify_entity_type("Bangalore"), "unknown");
assert_eq!(classify_entity_type("Flipkart"), "unknown");
}
#[test]
fn test_classify_name_multi_word_person() {
assert_eq!(classify_entity_type("Sarah Chen"), "person");
assert_eq!(classify_entity_type("Priya Sharma"), "person");
}
#[test]
fn test_classify_tech_blocklist() {
assert_eq!(classify_entity_type("FAISS"), "tech");
assert_eq!(classify_entity_type("ONNX"), "tech");
assert_eq!(classify_entity_type("Redis"), "tech");
assert_eq!(classify_entity_type("Python"), "tech");
}
#[test]
fn test_classify_tech_allcaps() {
assert_eq!(classify_entity_type("GPU"), "tech");
assert_eq!(classify_entity_type("API"), "tech");
}
#[test]
fn test_classify_unknown() {
assert_eq!(classify_entity_type("recommendation engine"), "unknown");
assert_eq!(classify_entity_type("data pipeline"), "unknown");
assert_eq!(classify_entity_type("sleep patterns"), "unknown");
}
#[test]
fn test_classify_with_rel_person_person() {
let (s, d) = classify_with_relationship("Arjun", "Priya", "married_to");
assert_eq!(s, "person");
assert_eq!(d, "person");
}
#[test]
fn test_classify_with_rel_person_place() {
let (s, d) = classify_with_relationship("Priya", "Bangalore", "lives_in");
assert_eq!(s, "person");
assert_eq!(d, "place");
}
#[test]
fn test_classify_with_rel_person_org() {
let (s, d) = classify_with_relationship("Priya", "Flipkart", "works_at");
assert_eq!(s, "person");
assert_eq!(d, "organization");
}
#[test]
fn test_classify_with_rel_tech_dst() {
let (s, d) = classify_with_relationship("FAISS", "data pipeline", "uses");
assert_eq!(s, "tech");
assert_eq!(d, "tech");
}
#[test]
fn test_classify_with_rel_built_with() {
let (s, d) = classify_with_relationship("MyApp", "React", "built_with");
assert_eq!(s, "project");
assert_eq!(d, "tech");
}
#[test]
fn test_classify_with_rel_deployed_on() {
let (s, d) = classify_with_relationship("MyApp", "AWS", "deployed_on");
assert_eq!(s, "project");
assert_eq!(d, "infrastructure");
}
#[test]
fn test_classify_with_rel_works_on() {
let (s, d) = classify_with_relationship("Pranab", "YantrikDB", "works_on");
assert_eq!(s, "person");
assert_eq!(d, "project");
}
#[test]
fn test_classify_with_rel_fallback() {
let (s, d) = classify_with_relationship("FAISS", "data pipeline", "related_to");
assert_eq!(s, "tech");
assert_eq!(d, "unknown");
}
}
#[cfg(test)]
mod code_stripping_tests {
use super::*;
#[test]
fn code_identifiers_do_not_become_entities() {
let text = "Alice deployed the service.\n\n```python\n\
@app.route('/login', methods=['GET', 'POST'])\n\
def login():\n data = LoginSchema(String)\n```\n\
She reported it to Acme Corp.";
let got = extract_heuristic_entities(text);
for bad in ["GET", "POST", "String", "LoginSchema"] {
assert!(
!got.iter().any(|e| e.contains(bad)),
"code identifier {bad:?} leaked into entities: {got:?}"
);
}
assert!(
got.iter().any(|e| e == "Alice"),
"lost prose entity: {got:?}"
);
assert!(
got.iter().any(|e| e.contains("Acme")),
"lost prose entity after the block: {got:?}"
);
}
#[test]
fn inline_spans_are_stripped_without_welding_neighbours() {
let got = extract_heuristic_entities("Bob set `MAX_RETRIES` Carol reviewed it");
assert!(!got.iter().any(|e| e.contains("MAX_RETRIES")), "{got:?}");
assert!(got.iter().any(|e| e == "Bob"), "{got:?}");
assert!(got.iter().any(|e| e == "Carol"), "{got:?}");
assert!(!got.iter().any(|e| e == "Bob Carol"), "welded: {got:?}");
}
#[test]
fn inline_span_before_a_fence_is_still_stripped() {
let text = "`GET` Alice then
```python
class User: pass
```
done";
let got = extract_heuristic_entities(text);
assert!(
!got.iter().any(|e| e.contains("GET")),
"inline span leaked: {got:?}"
);
assert!(
!got.iter().any(|e| e.contains("User")),
"fence leaked: {got:?}"
);
assert!(got.iter().any(|e| e == "Alice"), "prose lost: {got:?}");
}
#[test]
fn text_without_backticks_is_unchanged() {
let plain = "Alice Chen is the CEO of Acme Corp";
assert_eq!(
extract_heuristic_entities(plain),
extract_heuristic_entities_inner(plain, &|_| None),
"no-backtick path must be byte-identical to the pre-change behavior"
);
assert!(matches!(strip_code(plain), std::borrow::Cow::Borrowed(_)));
}
#[test]
fn unterminated_markers_do_not_drop_prose() {
let got = extract_heuristic_entities("Dave noted ` then Erin shipped it");
assert!(
got.iter().any(|e| e == "Erin"),
"prose lost after stray tick: {got:?}"
);
}
}
#[cfg(test)]
mod stopword_hygiene_tests {
use super::*;
#[test]
fn all_caps_function_words_are_not_entities() {
for text in [
"AT the meeting we shipped it",
"THE release went out",
"DID the migration finish",
"NOT a real entity here",
] {
let ents = extract_heuristic_entities(text);
for bad in ["AT", "THE", "DID", "NOT"] {
assert!(
!ents.iter().any(|e| e == bad),
"{bad:?} became an entity from {text:?} -> {ents:?}"
);
}
}
}
#[test]
fn mixed_case_function_words_are_not_entities() {
let ents = extract_heuristic_entities("aT tHe meeting, dId anything ship");
assert!(
!ents.iter().any(|e| e.eq_ignore_ascii_case("at")
|| e.eq_ignore_ascii_case("the")
|| e.eq_ignore_ascii_case("did")),
"mixed-case function word survived: {ents:?}"
);
}
#[test]
fn newly_listed_function_words_are_not_entities() {
let ents = extract_heuristic_entities("Most of it shipped. Not all. More later.");
for bad in ["Most", "Not", "More"] {
assert!(
!ents.iter().any(|e| e == bad),
"{bad:?} became an entity -> {ents:?}"
);
}
}
#[test]
fn bare_month_names_are_not_entities() {
let ents = extract_heuristic_entities("June was busy. We shipped in March.");
for bad in ["June", "March"] {
assert!(
!ents.iter().any(|e| e == bad),
"{bad:?} became an entity -> {ents:?}"
);
}
}
#[test]
fn real_entities_still_extracted() {
let ents = extract_heuristic_entities(
"At Yantrik Systems we met Alice Chen about the Boston office.",
);
for good in ["Yantrik Systems", "Alice Chen", "Boston"] {
assert!(
ents.iter()
.any(|e| e.contains(good) || good.contains(e.as_str())),
"real entity {good:?} was lost -> {ents:?}"
);
}
}
#[test]
fn all_caps_acronyms_survive() {
let ents = extract_heuristic_entities("The NASA contract and the HNSW index shipped.");
assert!(
ents.iter().any(|e| e.contains("NASA")),
"NASA was stripped as if it were a function word -> {ents:?}"
);
}
}
#[cfg(test)]
mod prose_run_tests {
use super::*;
#[test]
fn all_caps_headings_are_not_entities() {
for text in [
"THINGS I MISSED THAT CODEX FOUND BY READING THE CODE follow.",
"USER MUST UPDATE MCP CONFIG before restarting.",
"REAL ESTATE TAX ANALYSIS was attached.",
"HERMES REMOTE DESKTOP LIVE VERIFICATION PASSED today.",
] {
for e in extract_heuristic_entities(text) {
let caps = e
.split_whitespace()
.filter(|t| is_all_caps_token(t))
.count();
assert!(
caps <= MAX_ALLCAPS_TOKENS,
"heading became entity {e:?} from {text:?}"
);
}
}
}
#[test]
fn overlong_capitalized_runs_are_not_entities() {
let ents =
extract_heuristic_entities("Recall Return Unrelated Records Root Cause Found Today");
assert!(
ents.iter()
.all(|e| e.split_whitespace().count() <= MAX_ENTITY_TOKENS),
"overlong run survived -> {ents:?}"
);
}
#[test]
fn short_acronyms_and_names_survive() {
let ents = extract_heuristic_entities(
"NASA and IBM Watson met Alice Chen at Yantrik Systems in San Francisco.",
);
for good in [
"NASA",
"IBM Watson",
"Alice Chen",
"Yantrik Systems",
"San Francisco",
] {
assert!(
ents.iter().any(|e| e.contains(good)),
"real entity {good:?} lost -> {ents:?}"
);
}
}
#[test]
fn two_token_all_caps_names_survive() {
let ents = extract_heuristic_entities("The NASA JPL team shipped it.");
assert!(
ents.iter().any(|e| e.contains("NASA JPL")),
"two-token acronym name lost -> {ents:?}"
);
}
}
#[cfg(test)]
mod possessive_entity_tests {
use super::*;
#[test]
fn possessives_are_canonicalized_before_becoming_entities() {
let ents = extract_heuristic_entities(
"Pranab's benchmark compared Reddit's API with Sol's Q2 plan.",
);
for canonical in ["Pranab", "Reddit", "Sol", "Q2"] {
assert!(
ents.iter().any(|e| e == canonical),
"canonical {canonical:?} missing from {ents:?}"
);
}
assert!(
ents.iter()
.all(|e| !e.ends_with("'s") && !e.ends_with('\'')),
"possessive phantom survived: {ents:?}"
);
}
#[test]
fn apostrophes_inside_names_are_preserved() {
let ents = extract_heuristic_entities("O'Brien met D'Arcy about O'Brien's release.");
assert!(ents.iter().any(|e| e == "O'Brien"), "{ents:?}");
assert!(ents.iter().any(|e| e == "D'Arcy"), "{ents:?}");
assert!(!ents.iter().any(|e| e == "O'Brien's"), "{ents:?}");
}
#[test]
fn capitalized_contractions_do_not_create_bare_phantoms() {
let ents = extract_heuristic_entities("Let's begin. It's ready. What's next?");
for bad in ["Let", "It", "What"] {
assert!(!ents.iter().any(|e| e == bad), "{bad:?} survived: {ents:?}");
}
}
}
#[cfg(test)]
mod entity_admission_tests {
use super::*;
#[test]
fn measured_junk_classes_are_refused() {
for bad in [
"2026", "0.19.0", "15", "STRATEGIC POINT", "MASTERING", "NOT 1348", "Recall Return Unrelated Records Root", "A Very Long Capitalized Phrase That Is Clearly A Sentence Not A Name",
] {
assert!(!admit_entity(bad), "{bad:?} was admitted");
}
}
#[test]
fn real_names_and_acronyms_are_admitted() {
for good in [
"Alice Chen",
"Fennwick Labs",
"San Francisco",
"NASA",
"HNSW",
"FAISS",
"CT128",
"ONNX",
"NASA JPL",
"Series A",
"Q2",
"Indian Institute",
"O'Brien",
"Yantrikdb",
] {
assert!(admit_entity(good), "{good:?} was refused");
}
}
#[test]
fn possessive_stragglers_are_refused_as_nodes() {
assert!(!admit_entity("Pranab\u{2019}s"));
assert!(!admit_entity("Pranab's"));
assert!(admit_entity("Pranab"));
}
#[test]
fn numbers_are_values_not_entities_but_still_relation_objects() {
let text = "CT128 runs 0.19.0 in production since 2026.";
let ents = extract_heuristic_entities(text);
assert!(ents.iter().any(|e| e == "CT128"), "{ents:?}");
assert!(
!ents.iter().any(|e| e == "0.19.0" || e == "2026"),
"value minted as entity: {ents:?}"
);
let values = extract_value_candidates(text);
assert_eq!(values, vec!["0.19.0".to_string(), "2026".to_string()]);
for good in ["1985", "0.19.0", "3.6", "2026-08-01", "12"] {
assert!(is_value_object(good), "{good:?} refused");
}
for bad in [
"67%", "2+", "24/7", "*/5", "~12", "+4.6%", "~121-127", "1.", "-3", "v2", "",
] {
assert!(!is_value_object(bad), "{bad:?} admitted as a value");
}
let mut cands = ents.clone();
cands.extend(values);
let rels = extract_heuristic_relations(text, &cands);
assert!(
rels.iter()
.any(|r| r.src == "CT128" && r.rel_type == "runs" && r.dst == "0.19.0"),
"runs claim lost its value object: {rels:?}"
);
}
#[test]
fn values_are_objects_only_for_relations_that_can_take_one() {
assert!(relation_admits_value_object("runs", "0.19.0"));
assert!(relation_admits_value_object("born_in", "1985"));
assert!(relation_admits_value_object("leads", "Acme")); assert!(!relation_admits_value_object("leads", "2"));
assert!(!relation_admits_value_object("works_at", "2026-08-11"));
assert!(!relation_admits_value_object("ceo_of", "42"));
}
#[test]
fn shouted_headings_never_reach_the_entity_list() {
let ents = extract_heuristic_entities(
"STRATEGIC POINT: MASTERING the release. The NASA JPL team shipped it.",
);
assert!(
!ents
.iter()
.any(|e| e.contains("STRATEGIC") || e == "MASTERING"),
"{ents:?}"
);
assert!(ents.iter().any(|e| e == "NASA JPL"), "{ents:?}");
}
}
#[cfg(test)]
mod common_word_tests {
use super::*;
#[test]
fn observations_classify_by_position_and_case_once_per_memory() {
let obs = token_case_observations(
"Critically, the build failed. Alice Moreau fixed it; make it green. Make it so!",
);
let has = |t: &str, c: TokenCase| obs.contains(&(t.to_string(), c));
assert!(has("critically", TokenCase::CapStart), "{obs:?}");
assert!(has("alice", TokenCase::CapStart), "{obs:?}");
assert!(has("moreau", TokenCase::CapMid), "{obs:?}");
assert!(
has("make", TokenCase::Lower) && has("make", TokenCase::CapStart),
"{obs:?}"
);
assert!(!has("make", TokenCase::CapMid), "{obs:?}");
assert_eq!(
obs.iter().filter(|(t, _)| t == "it").count(),
1,
"deduplicated per class"
);
}
#[test]
fn seed_refuses_sentence_starters_and_stats_override_both_ways() {
for w in [
"Critically",
"Failed",
"Idempotent",
"Lets",
"Make",
"Trying",
"Target",
] {
assert!(
is_common_word(w, None),
"{w} should be a common word by seed"
);
}
for w in ["Pranab", "Fennwick", "Berlin", "Yantrikdb"] {
assert!(!is_common_word(w, None), "{w} is a name");
}
assert!(is_common_word(
"Recall",
Some(CaseStats {
lower_n: 500,
cap_mid_n: 40,
cap_start_n: 0,
})
));
assert!(!is_common_word(
"Python",
Some(CaseStats {
lower_n: 200,
cap_mid_n: 150,
cap_start_n: 33,
})
));
assert!(!is_common_word(
"Target",
Some(CaseStats {
lower_n: 2,
cap_mid_n: 9,
cap_start_n: 1,
})
));
assert!(is_common_word(
"Make",
Some(CaseStats {
lower_n: 1,
cap_mid_n: 0,
cap_start_n: 0,
})
));
assert!(!is_common_word(
"Gizmo",
Some(CaseStats {
lower_n: 2,
cap_mid_n: 0,
cap_start_n: 0,
})
));
assert!(is_common_word(
"Gizmo",
Some(CaseStats {
lower_n: 4,
cap_mid_n: 1,
cap_start_n: 0,
})
));
}
#[test]
fn admission_with_stats_only_touches_single_token_names_and_never_acronyms() {
let none = |_: &str| None;
assert!(!admit_entity_with("Critically", none));
assert!(admit_entity_with("Alice Moreau", none));
assert!(
admit_entity_with("API", none),
"cold store, not a seed word"
);
assert!(
!admit_entity_with("CODE", none),
"cold store, shouted seed word"
);
assert!(admit_entity_with("Pranab", none));
let stats = |t: &str| match t {
"class" => Some(CaseStats {
lower_n: 450,
cap_mid_n: 57,
cap_start_n: 5,
}),
"api" => Some(CaseStats {
lower_n: 366,
cap_mid_n: 650,
cap_start_n: 28,
}),
_ => None,
};
assert!(!admit_entity_with("CLASS", stats));
assert!(admit_entity_with("API", stats));
assert!(is_common_word(
"Critically",
Some(CaseStats {
lower_n: 0,
cap_mid_n: 3,
cap_start_n: 8
})
));
assert!(is_common_word(
"FIX",
Some(CaseStats {
lower_n: 1083,
cap_mid_n: 232,
cap_start_n: 424
})
));
assert!(!is_common_word(
"Pranab",
Some(CaseStats {
lower_n: 108,
cap_mid_n: 1474,
cap_start_n: 903
})
));
assert!(!is_common_word(
"UTC",
Some(CaseStats {
lower_n: 4,
cap_mid_n: 952,
cap_start_n: 8
})
));
assert!(
is_common_word("None", None),
"literal values are seed words"
);
let obs2 = token_case_observations("we shipped it (Critically, twice) \u{2014} Finally.");
assert!(
obs2.contains(&("critically".to_string(), TokenCase::CapStart)),
"{obs2:?}"
);
assert!(
obs2.contains(&("finally".to_string(), TokenCase::CapStart)),
"{obs2:?}"
);
let learned = |t: &str| {
if t == "gizmo" {
Some(CaseStats {
lower_n: 6,
cap_mid_n: 0,
cap_start_n: 0,
})
} else {
None
}
};
assert!(!admit_entity_with("Gizmo", learned));
assert!(admit_entity_with("Gizmo Labs", learned));
}
#[test]
fn seed_is_lowercase_and_has_no_duplicates() {
let mut seen = std::collections::HashSet::new();
for w in COMMON_WORD_SEED {
assert_eq!(*w, w.to_lowercase(), "{w}");
assert!(seen.insert(*w), "duplicate {w}");
}
}
}