use async_trait::async_trait;
use crate::phoneme::TextEmbedder;
use crate::processor::{ProcessError, ProcessResult, TextProcessor};
use crate::types::{ContextSnapshot, FieldType};
pub fn split_sentences(text: &str) -> Vec<String> {
let mut sentences = Vec::new();
let mut current = String::new();
for c in text.chars() {
current.push(c);
if matches!(c, '.' | '!' | '?' | '。' | '!' | '?') {
let trimmed = current.trim().to_string();
if !trimmed.is_empty() {
sentences.push(trimmed);
}
current.clear();
}
}
let trimmed = current.trim().to_string();
if !trimmed.is_empty() {
sentences.push(trimmed);
}
sentences
}
pub struct ParagraphSplitter {
embedder: Option<Box<dyn TextEmbedder>>,
pub depth_ratio: f32,
pub min_similarity_range: f32,
pub center_embeddings: bool,
pub max_sentences: usize,
pub separator: String,
}
impl Default for ParagraphSplitter {
fn default() -> Self {
Self::new()
}
}
impl ParagraphSplitter {
pub fn new() -> Self {
Self {
embedder: None,
depth_ratio: 0.5,
min_similarity_range: 0.05,
center_embeddings: false,
max_sentences: 8,
separator: "\n\n".to_string(),
}
}
pub fn with_embedder(mut self, embedder: impl TextEmbedder + 'static) -> Self {
self.embedder = Some(Box::new(embedder));
self
}
pub fn with_depth_ratio(mut self, ratio: f32) -> Self {
self.depth_ratio = ratio;
self
}
pub fn with_min_similarity_range(mut self, range: f32) -> Self {
self.min_similarity_range = range;
self
}
pub fn with_center_embeddings(mut self, center: bool) -> Self {
self.center_embeddings = center;
self
}
pub fn breaks_for_sentences(&self, sentences: &[String]) -> Vec<usize> {
let embeddings = self.embed_sentences(sentences);
let similarities = self.inter_sentence_similarities(&embeddings);
self.find_breaks(sentences.len(), &similarities)
}
pub fn with_max_sentences(mut self, max: usize) -> Self {
self.max_sentences = max;
self
}
fn should_split(field_type: &Option<FieldType>) -> bool {
match field_type {
None => true, Some(FieldType::Document) | Some(FieldType::EmailCompose) => true,
Some(FieldType::ChatMessage)
| Some(FieldType::Terminal)
| Some(FieldType::SearchBar)
| Some(FieldType::CodeEditor) => false,
Some(FieldType::Generic) => true,
}
}
fn embed_sentences(&self, sentences: &[String]) -> Vec<Option<Vec<f32>>> {
let embedder = match &self.embedder {
Some(e) => e,
None => return vec![None; sentences.len()],
};
sentences.iter().map(|s| embedder.embed(s).ok()).collect()
}
fn cosine_sim(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return 0.0;
}
a.iter()
.zip(b.iter())
.map(|(x, y)| x * y)
.sum::<f32>()
.max(0.0)
}
fn center(embeddings: &[Option<Vec<f32>>]) -> Vec<Option<Vec<f32>>> {
let present: Vec<&Vec<f32>> = embeddings.iter().flatten().collect();
if present.len() < 2 {
return embeddings.to_vec();
}
let dim = present[0].len();
if present.iter().any(|v| v.len() != dim) {
return embeddings.to_vec();
}
let mut mean = vec![0.0f32; dim];
for v in &present {
for (m, x) in mean.iter_mut().zip(v.iter()) {
*m += x;
}
}
for m in mean.iter_mut() {
*m /= present.len() as f32;
}
embeddings
.iter()
.map(|e| {
e.as_ref().map(|v| {
let mut c: Vec<f32> = v.iter().zip(&mean).map(|(x, m)| x - m).collect();
let norm: f32 = c.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in c.iter_mut() {
*x /= norm;
}
}
c
})
})
.collect()
}
fn inter_sentence_similarities(&self, embeddings: &[Option<Vec<f32>>]) -> Vec<Option<f32>> {
if embeddings.len() < 2 {
return vec![];
}
let owned;
let embeddings = if self.center_embeddings {
owned = Self::center(embeddings);
&owned[..]
} else {
embeddings
};
(0..embeddings.len() - 1)
.map(|i| match (&embeddings[i], &embeddings[i + 1]) {
(Some(a), Some(b)) => Some(Self::cosine_sim(a, b)),
_ => None,
})
.collect()
}
fn find_breaks(&self, n_sentences: usize, similarities: &[Option<f32>]) -> Vec<usize> {
if n_sentences <= 1 {
return vec![];
}
let breaks = self.semantic_breaks(similarities);
let mut final_breaks = Vec::new();
let mut prev_break = 0;
for &br in &breaks {
self.split_long_segment(prev_break, br, similarities, &mut final_breaks);
final_breaks.push(br);
prev_break = br;
}
self.split_long_segment(prev_break, n_sentences, similarities, &mut final_breaks);
final_breaks.sort();
final_breaks.dedup();
final_breaks
}
fn valley_depth(sims: &[f32], i: usize) -> f32 {
let mut left = sims[i];
let mut j = i;
while j > 0 && sims[j - 1] >= left {
left = sims[j - 1];
j -= 1;
}
let mut right = sims[i];
let mut j = i;
while j + 1 < sims.len() && sims[j + 1] >= right {
right = sims[j + 1];
j += 1;
}
(left - sims[i]) + (right - sims[i])
}
fn is_local_min(sims: &[f32], i: usize) -> bool {
let left_ok = i == 0 || sims[i - 1] > sims[i];
let right_ok = i + 1 == sims.len() || sims[i + 1] > sims[i];
left_ok && right_ok
}
fn semantic_breaks(&self, similarities: &[Option<f32>]) -> Vec<usize> {
if similarities.len() < 2 {
return Vec::new();
}
let sims: Vec<f32> = match similarities.iter().copied().collect::<Option<Vec<f32>>>() {
Some(v) => v,
None => return Vec::new(),
};
let max = sims.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let min = sims.iter().copied().fold(f32::INFINITY, f32::min);
if max - min < self.min_similarity_range {
return Vec::new();
}
let candidates: Vec<(usize, f32)> = (0..sims.len())
.filter(|&i| Self::is_local_min(&sims, i))
.map(|i| (i, Self::valley_depth(&sims, i)))
.collect();
let deepest = candidates
.iter()
.map(|(_, d)| *d)
.fold(f32::NEG_INFINITY, f32::max);
if !deepest.is_finite() || deepest <= 0.0 {
return Vec::new();
}
candidates
.iter()
.filter(|(_, d)| *d >= self.depth_ratio * deepest)
.map(|(i, _)| i + 1) .collect()
}
fn split_long_segment(
&self,
start: usize,
end: usize,
similarities: &[Option<f32>],
breaks: &mut Vec<usize>,
) {
let len = end - start;
if len <= self.max_sentences {
return;
}
let mut min_sim = f32::INFINITY;
let mut min_idx = start + len / 2;
for i in start..end.saturating_sub(1) {
if i < similarities.len() {
if let Some(s) = similarities[i] {
if s < min_sim {
min_sim = s;
min_idx = i + 1;
}
}
}
}
breaks.push(min_idx);
self.split_long_segment(start, min_idx, similarities, breaks);
self.split_long_segment(min_idx, end, similarities, breaks);
}
}
#[async_trait]
impl TextProcessor for ParagraphSplitter {
async fn process(
&self,
text: &str,
ctx: &ContextSnapshot,
) -> Result<ProcessResult, ProcessError> {
if !Self::should_split(&ctx.field_type) {
return Ok(ProcessResult {
text: text.to_string(),
corrections: vec![],
});
}
let sentences = split_sentences(text);
if sentences.len() <= 1 {
return Ok(ProcessResult {
text: text.to_string(),
corrections: vec![],
});
}
let embeddings = self.embed_sentences(&sentences);
let similarities = self.inter_sentence_similarities(&embeddings);
let breaks = self.find_breaks(sentences.len(), &similarities);
if breaks.is_empty() {
return Ok(ProcessResult {
text: text.to_string(),
corrections: vec![],
});
}
let mut paragraphs: Vec<String> = Vec::new();
let mut current = Vec::new();
for (i, sentence) in sentences.iter().enumerate() {
if breaks.contains(&i) && !current.is_empty() {
paragraphs.push(current.join(" "));
current.clear();
}
current.push(sentence.as_str());
}
if !current.is_empty() {
paragraphs.push(current.join(" "));
}
let result = paragraphs.join(&self.separator);
tracing::debug!(
n_sentences = sentences.len(),
n_paragraphs = paragraphs.len(),
breaks = ?breaks,
"paragraph splitting applied"
);
Ok(ProcessResult {
text: result,
corrections: vec![],
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_split_sentences_basic() {
let s = split_sentences("Hello world. How are you? I am fine!");
assert_eq!(s, vec!["Hello world.", "How are you?", "I am fine!"]);
}
#[test]
fn test_split_sentences_no_terminal() {
let s = split_sentences("Hello world");
assert_eq!(s, vec!["Hello world"]);
}
#[test]
fn test_split_sentences_japanese() {
let s = split_sentences("今日は天気がいい。明日は雨だ。");
assert_eq!(s, vec!["今日は天気がいい。", "明日は雨だ。"]);
}
#[test]
fn test_split_sentences_empty() {
let s = split_sentences("");
assert!(s.is_empty());
}
fn depths(sims: &[f32]) -> Vec<f32> {
(0..sims.len())
.map(|i| ParagraphSplitter::valley_depth(sims, i))
.collect()
}
#[test]
fn valley_depth_is_invariant_under_a_uniform_shift() {
let bge = [0.68, 0.66, 0.51, 0.67, 0.69];
let granite: Vec<f32> = bge.iter().map(|s| s + 0.20).collect();
for (a, b) in depths(&bge).iter().zip(depths(&granite).iter()) {
assert!((a - b).abs() < 1e-5, "depth moved under shift: {a} vs {b}");
}
}
#[test]
fn valley_depth_measures_the_drop_from_both_sides() {
let sims = [0.88, 0.86, 0.71, 0.87, 0.89];
let d = ParagraphSplitter::valley_depth(&sims, 2);
assert!((d - 0.35).abs() < 1e-5, "got {d}");
}
#[test]
fn a_monotone_decline_has_no_interior_valley() {
let sims = [0.9, 0.8, 0.7, 0.6];
for i in 0..3 {
assert!(!ParagraphSplitter::is_local_min(&sims, i), "index {i}");
}
assert!(ParagraphSplitter::is_local_min(&sims, 3));
}
#[test]
fn ends_count_as_minima_so_edge_shifts_are_not_missed() {
assert!(ParagraphSplitter::is_local_min(&[0.2, 0.9, 0.9], 0));
assert!(ParagraphSplitter::is_local_min(&[0.9, 0.9, 0.2], 2));
}
#[test]
fn same_breaks_on_both_backends_for_the_same_text() {
let splitter = ParagraphSplitter::new();
let bge: Vec<Option<f32>> = [0.68, 0.66, 0.51, 0.67, 0.69]
.iter()
.map(|s| Some(*s))
.collect();
let granite: Vec<Option<f32>> = bge.iter().map(|s| Some(s.unwrap() + 0.20)).collect();
let a = splitter.semantic_breaks(&bge);
let b = splitter.semantic_breaks(&granite);
assert_eq!(a, b, "backend shift changed the break points");
assert_eq!(a, vec![3], "expected the single valley at index 2");
}
#[test]
fn the_old_absolute_rule_would_have_disagreed() {
const OLD_THRESHOLD: f32 = 0.5;
let bge = [0.68, 0.66, 0.51, 0.67, 0.69];
let granite: Vec<f32> = bge.iter().map(|s| s + 0.20).collect();
assert_eq!(bge.iter().filter(|s| **s < OLD_THRESHOLD).count(), 0);
assert_eq!(granite.iter().filter(|s| **s < OLD_THRESHOLD).count(), 0);
const BGE_TUNED: f32 = 0.6;
assert_eq!(bge.iter().filter(|s| **s < BGE_TUNED).count(), 1);
assert_eq!(granite.iter().filter(|s| **s < BGE_TUNED).count(), 0);
}
#[test]
fn flat_similarity_does_not_split() {
let splitter = ParagraphSplitter::new();
let flat: Vec<Option<f32>> = [0.90, 0.90, 0.89, 0.90].iter().map(|s| Some(*s)).collect();
assert!(splitter.semantic_breaks(&flat).is_empty());
}
#[test]
fn a_failed_embedding_disables_semantic_splitting() {
let splitter = ParagraphSplitter::new();
let holed = vec![Some(0.9), None, Some(0.2), Some(0.9)];
assert!(splitter.semantic_breaks(&holed).is_empty());
}
#[test]
fn a_single_boundary_is_not_enough_to_judge() {
let splitter = ParagraphSplitter::new();
assert!(splitter.semantic_breaks(&[Some(0.1)]).is_empty());
assert!(splitter.semantic_breaks(&[]).is_empty());
}
#[test]
fn shallow_valleys_are_dropped_relative_to_the_deepest() {
let splitter = ParagraphSplitter::new();
let sims: Vec<Option<f32>> = [0.90, 0.80, 0.90, 0.30, 0.90]
.iter()
.map(|s| Some(*s))
.collect();
assert_eq!(splitter.semantic_breaks(&sims), vec![4]);
let permissive = ParagraphSplitter::new().with_depth_ratio(0.1);
assert_eq!(permissive.semantic_breaks(&sims), vec![2, 4]);
}
#[tokio::test]
async fn test_splitter_single_sentence() {
let splitter = ParagraphSplitter::new();
let ctx = ContextSnapshot::default();
let r = splitter.process("Hello world.", &ctx).await.unwrap();
assert_eq!(r.text, "Hello world.");
}
#[tokio::test]
async fn test_splitter_skip_chat() {
let splitter = ParagraphSplitter::new();
let ctx = ContextSnapshot {
field_type: Some(FieldType::ChatMessage),
..Default::default()
};
let text = "First sentence. Second sentence. Third sentence. Fourth. Fifth. Sixth. Seventh. Eighth. Ninth. Tenth.";
let r = splitter.process(text, &ctx).await.unwrap();
assert_eq!(r.text, text); }
#[tokio::test]
async fn test_splitter_max_sentences_no_embedder() {
let splitter = ParagraphSplitter::new().with_max_sentences(3);
let ctx = ContextSnapshot::default();
let text = "One. Two. Three. Four. Five. Six.";
let r = splitter.process(text, &ctx).await.unwrap();
assert!(
r.text.contains("\n\n"),
"Expected paragraph break in: {}",
r.text
);
}
#[tokio::test]
async fn test_splitter_under_max_no_change() {
let splitter = ParagraphSplitter::new().with_max_sentences(10);
let ctx = ContextSnapshot::default();
let text = "One. Two. Three.";
let r = splitter.process(text, &ctx).await.unwrap();
assert!(!r.text.contains("\n\n"));
}
#[tokio::test]
async fn test_splitter_semantic_break() {
let emb_a = vec![1.0, 0.0, 0.0]; let emb_b = vec![0.0, 1.0, 0.0];
struct OrderedEmbedder {
embeddings: std::sync::Mutex<std::collections::VecDeque<Vec<f32>>>,
}
impl TextEmbedder for OrderedEmbedder {
fn embed(&self, _text: &str) -> Result<Vec<f32>, ProcessError> {
let mut q = self.embeddings.lock().unwrap();
Ok(q.pop_front().unwrap_or_else(|| vec![0.0; 3]))
}
}
let embedder = OrderedEmbedder {
embeddings: std::sync::Mutex::new(
vec![emb_a.clone(), emb_a.clone(), emb_b.clone()].into(),
),
};
let splitter = ParagraphSplitter::new()
.with_embedder(embedder)
.with_depth_ratio(0.5);
let ctx = ContextSnapshot::default();
let text = "Dogs are great pets. Cats are also wonderful. The stock market crashed today.";
let r = splitter.process(text, &ctx).await.unwrap();
assert!(
r.text.contains("\n\n"),
"Expected paragraph break in: {}",
r.text
);
let parts: Vec<&str> = r.text.split("\n\n").collect();
assert_eq!(parts.len(), 2);
}
}