#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TokenSpan {
pub start: usize,
pub end: usize,
}
impl TokenSpan {
pub fn new(start: usize, end: usize) -> Self {
Self { start, end }
}
pub fn intersects(&self, range_start: usize, range_end: usize) -> bool {
self.start < range_end && range_start < self.end
}
}
pub fn char_spans_to_byte_spans(text: &str, offsets: &[(usize, usize)]) -> Vec<TokenSpan> {
let char_to_byte: Vec<usize> = text.char_indices().map(|(b, _)| b).collect();
let text_len_chars = char_to_byte.len();
let mut out = Vec::with_capacity(offsets.len());
for &(start, end) in offsets {
if end <= start {
continue; }
let start = start.min(text_len_chars);
let end = end.min(text_len_chars);
if end <= start {
continue;
}
let byte_start = char_to_byte[start];
let byte_end = if end == text_len_chars {
text.len()
} else {
char_to_byte[end]
};
out.push(TokenSpan::new(byte_start, byte_end));
}
out
}
#[derive(Debug, Clone, PartialEq)]
pub struct TokenEmbedding {
pub span: TokenSpan,
pub vector: Vec<f32>,
}
#[trait_variant::make(TokenLevelEmbeddings: Send)]
pub trait LocalTokenLevelEmbeddings {
async fn embed_tokens(&self, text: &str) -> Result<Vec<TokenEmbedding>, crate::EmbeddingError>;
}
#[cfg(test)]
mod tests {
use super::*;
struct MockTokenEmbeddings;
impl TokenLevelEmbeddings for MockTokenEmbeddings {
async fn embed_tokens(
&self,
text: &str,
) -> Result<Vec<TokenEmbedding>, crate::EmbeddingError> {
if text.trim().is_empty() {
return Err(crate::EmbeddingError::EmptyInput);
}
let mut out = Vec::new();
let mut cursor = 0usize;
for word in text.split_whitespace() {
let start = text[cursor..]
.find(word)
.map(|p| cursor + p)
.unwrap_or(cursor);
let end = start + word.len();
cursor = end;
let vector = vec![word.bytes().map(|b| b as f32).sum::<f32>(), 1.0];
out.push(TokenEmbedding {
span: TokenSpan::new(start, end),
vector,
});
}
Ok(out)
}
}
async fn embed<E: TokenLevelEmbeddings>(
e: &E,
text: &str,
) -> Result<Vec<TokenEmbedding>, crate::EmbeddingError> {
e.embed_tokens(text).await
}
#[tokio::test]
async fn embed_tokens_returns_spans_and_vectors() {
let embedder = MockTokenEmbeddings;
let tokens = embed(&embedder, "hello late world").await.unwrap();
assert_eq!(tokens.len(), 3);
assert_eq!(tokens[0].span, TokenSpan::new(0, 5));
assert_eq!(tokens[1].span, TokenSpan::new(6, 10));
assert_eq!(tokens[2].span, TokenSpan::new(11, 16));
assert_eq!(
tokens[0].vector,
vec![b"hello".iter().map(|b| *b as f32).sum::<f32>(), 1.0]
);
}
async fn generic_len<E: TokenLevelEmbeddings>(e: &E, text: &str) -> usize {
e.embed_tokens(text).await.unwrap().len()
}
#[tokio::test]
async fn trait_variant_send_is_generically_usable() {
let embedder = MockTokenEmbeddings;
assert_eq!(generic_len(&embedder, "a b c").await, 3);
}
#[tokio::test]
async fn spans_are_byte_offsets() {
let embedder = MockTokenEmbeddings;
let tokens = embed(&embedder, "你好 world").await.unwrap();
assert_eq!(tokens.len(), 2);
assert_eq!(tokens[0].span, TokenSpan::new(0, 6));
assert_eq!(tokens[1].span, TokenSpan::new(7, 12));
assert_eq!(
&"你好 world"[tokens[1].span.start..tokens[1].span.end],
"world"
);
}
#[tokio::test]
async fn empty_input_rejected() {
let embedder = MockTokenEmbeddings;
let err = embed(&embedder, " ").await.unwrap_err();
assert!(matches!(err, crate::EmbeddingError::EmptyInput));
}
#[test]
fn span_intersects_boundaries() {
let span = TokenSpan::new(5, 10);
assert!(span.intersects(4, 6));
assert!(span.intersects(9, 11));
assert!(span.intersects(0, 100));
assert!(!span.intersects(0, 5), "end-exclusive");
assert!(!span.intersects(10, 20), "start-inclusive");
}
#[test]
fn char_spans_to_byte_spans_ascii() {
let spans = char_spans_to_byte_spans("hello world", &[(0, 5), (6, 11)]);
assert_eq!(spans, vec![TokenSpan::new(0, 5), TokenSpan::new(6, 11)]);
}
#[test]
fn char_spans_to_byte_spans_multibyte() {
let text = "你好 world";
let spans = char_spans_to_byte_spans(text, &[(3, 8)]);
assert_eq!(spans, vec![TokenSpan::new(7, 12)]);
assert_eq!(&text[spans[0].start..spans[0].end], "world");
}
#[test]
fn char_spans_to_byte_spans_out_of_range() {
let spans = char_spans_to_byte_spans("abc", &[(0, 10), (5, 2), (1, 1)]);
assert_eq!(spans.len(), 1);
assert_eq!(spans[0], TokenSpan::new(0, 3));
}
#[test]
fn char_spans_to_byte_spans_empty() {
assert!(char_spans_to_byte_spans("", &[]).is_empty());
assert!(char_spans_to_byte_spans("abc", &[]).is_empty());
}
}