use std::collections::HashMap;
use crate::types::entity::{Entity, EntityCategory};
pub const GLINER_WINDOW_TOKENS: usize = 512;
pub const GLINER_WINDOW_OVERLAP_TOKENS: usize = 64;
#[derive(Debug, Clone, PartialEq, Eq)]
#[cfg_attr(alef, alef(skip))]
pub struct TextWindow {
pub byte_offset: usize,
pub text: String,
}
fn is_word_char(character: char) -> bool {
character.is_alphanumeric() || character == '_'
}
fn token_spans(text: &str) -> Vec<(usize, usize)> {
let mut spans: Vec<(usize, usize)> = Vec::new();
let mut characters = text.char_indices().peekable();
while let Some((index, character)) = characters.next() {
if character.is_whitespace() {
continue;
}
if !is_word_char(character) {
spans.push((index, index + character.len_utf8()));
continue;
}
let mut end = index + character.len_utf8();
while let Some(&(next_index, next_character)) = characters.peek() {
if !is_word_char(next_character) {
break;
}
end = next_index + next_character.len_utf8();
characters.next();
}
spans.push((index, end));
}
spans
}
#[cfg_attr(alef, alef(skip))]
pub fn split_into_windows(text: &str, window_tokens: usize, overlap_tokens: usize) -> Vec<TextWindow> {
let tokens = token_spans(text);
if tokens.is_empty() {
return Vec::new();
}
let window_tokens = window_tokens.max(1);
let overlap_tokens = overlap_tokens.min(window_tokens - 1);
let step = window_tokens - overlap_tokens;
let mut windows = Vec::new();
let mut first_token = 0usize;
loop {
let last_token = (first_token + window_tokens).min(tokens.len());
let start_byte = tokens[first_token].0;
let end_byte = tokens[last_token - 1].1;
windows.push(TextWindow {
byte_offset: start_byte,
text: text[start_byte..end_byte].to_string(),
});
if last_token >= tokens.len() {
break;
}
first_token += step;
}
windows
}
#[cfg_attr(alef, alef(skip))]
pub fn merge_windowed_entities(window_offsets: &[usize], per_window: Vec<Vec<Entity>>) -> Vec<Entity> {
let mut all: Vec<Entity> = Vec::new();
for (offset, entities) in window_offsets.iter().zip(per_window) {
for mut entity in entities {
let start = offset.saturating_add(entity.start as usize);
let end = offset.saturating_add(entity.end as usize);
let (Ok(start), Ok(end)) = (u32::try_from(start), u32::try_from(end)) else {
continue;
};
entity.start = start;
entity.end = end;
all.push(entity);
}
}
all.sort_by(|a, b| {
a.start
.cmp(&b.start)
.then((b.end.saturating_sub(b.start)).cmp(&(a.end.saturating_sub(a.start))))
});
let mut furthest_end: HashMap<EntityCategory, u32> = HashMap::new();
let mut kept: Vec<Entity> = Vec::with_capacity(all.len());
for entity in all {
if let Some(&previous_end) = furthest_end.get(&entity.category)
&& entity.start < previous_end
{
continue;
}
furthest_end.insert(entity.category.clone(), entity.end);
kept.push(entity);
}
kept
}
#[cfg_attr(alef, alef(skip))]
pub fn entities_for_every_occurrence(
text: &str,
mention: &str,
category: EntityCategory,
confidence: Option<f32>,
) -> Vec<Entity> {
if mention.is_empty() {
return Vec::new();
}
let mut out = Vec::new();
for (start, matched) in text.match_indices(mention) {
let (Ok(start), Ok(end)) = (u32::try_from(start), u32::try_from(start + matched.len())) else {
continue;
};
out.push(Entity {
category: category.clone(),
text: mention.to_string(),
start,
end,
confidence,
});
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn person(start: u32, end: u32, text: &str) -> Entity {
Entity {
category: EntityCategory::Person,
text: text.to_string(),
start,
end,
confidence: None,
}
}
#[test]
fn token_spans_splits_words_and_punctuation() {
assert_eq!(token_spans("Ada, Bob"), vec![(0, 3), (3, 4), (5, 8)]);
}
#[test]
fn token_spans_is_empty_for_whitespace_only_text() {
assert!(token_spans(" \n\t ").is_empty());
}
#[test]
fn split_into_windows_returns_single_window_when_text_fits() {
let windows = split_into_windows("Ada Lovelace works here", 512, 64);
assert_eq!(windows.len(), 1);
assert_eq!(windows[0].byte_offset, 0);
assert_eq!(windows[0].text, "Ada Lovelace works here");
}
#[test]
fn split_into_windows_covers_every_token_past_the_budget() {
let text = (0..50).map(|index| format!("w{index}")).collect::<Vec<_>>().join(" ");
let windows = split_into_windows(&text, 10, 2);
assert!(windows.len() > 1, "long text must be windowed");
assert_eq!(windows[0].byte_offset, 0);
let last = windows.last().expect("at least one window");
assert!(
last.text.ends_with("w49"),
"final window must reach the end of the source: {}",
last.text
);
for window in &windows {
assert_eq!(
&text[window.byte_offset..window.byte_offset + window.text.len()],
window.text,
"window byte_offset must locate its own text in the source"
);
}
}
#[test]
fn split_into_windows_overlaps_adjacent_windows() {
let text = (0..30).map(|index| format!("w{index}")).collect::<Vec<_>>().join(" ");
let windows = split_into_windows(&text, 10, 3);
assert!(windows.len() >= 3);
assert!(
windows[1].byte_offset < windows[0].byte_offset + windows[0].text.len(),
"window 1 must start before window 0 ends"
);
}
#[test]
fn merge_windowed_entities_shifts_into_source_coordinates() {
let merged = merge_windowed_entities(&[0, 100], vec![vec![person(0, 3, "Ada")], vec![person(4, 7, "Bob")]]);
assert_eq!(merged.len(), 2);
assert_eq!((merged[0].start, merged[0].end), (0, 3));
assert_eq!((merged[1].start, merged[1].end), (104, 107));
}
#[test]
fn merge_windowed_entities_collapses_overlap_duplicates() {
let merged = merge_windowed_entities(
&[0, 90],
vec![vec![person(100, 103, "Ada")], vec![person(10, 13, "Ada")]],
);
assert_eq!(merged.len(), 1, "same span from two windows must collapse: {merged:?}");
assert_eq!((merged[0].start, merged[0].end), (100, 103));
}
#[test]
fn entities_for_every_occurrence_finds_all_three_occurrences() {
let text = "Ada paid Ada then Ada left";
let found = entities_for_every_occurrence(text, "Ada", EntityCategory::Person, Some(0.9));
assert_eq!(found.len(), 3);
assert_eq!(
found.iter().map(|e| (e.start, e.end)).collect::<Vec<_>>(),
vec![(0, 3), (9, 12), (18, 21)]
);
for entity in &found {
assert_eq!(&text[entity.start as usize..entity.end as usize], "Ada");
}
}
#[test]
fn entities_for_every_occurrence_returns_nothing_for_empty_mention() {
assert!(entities_for_every_occurrence("Ada", "", EntityCategory::Person, None).is_empty());
}
}