use std::ops::Range;
use itertools::Itertools;
use memchr::memchr2_iter;
use super::super::{
ChunkConfig, ChunkSizer,
splitter::{SemanticLevel, Splitter},
};
use super::fallback::GRAPHEME_SEGMENTER;
#[derive(Debug)]
pub struct TextSplitter<Sizer>
where
Sizer: ChunkSizer,
{
chunk_config: ChunkConfig<Sizer>,
}
impl<Sizer> TextSplitter<Sizer>
where
Sizer: ChunkSizer,
{
#[must_use]
pub fn new(chunk_config: impl Into<ChunkConfig<Sizer>>) -> Self {
Self {
chunk_config: chunk_config.into(),
}
}
pub fn chunks<'splitter, 'text: 'splitter>(
&'splitter self,
text: &'text str,
) -> impl Iterator<Item = &'text str> + 'splitter {
Splitter::<_>::chunks(self, text)
}
}
impl<Sizer> Splitter<Sizer> for TextSplitter<Sizer>
where
Sizer: ChunkSizer,
{
type Level = LineBreaks;
fn chunk_config(&self) -> &ChunkConfig<Sizer> {
&self.chunk_config
}
fn parse(&self, text: &str) -> Vec<(Self::Level, Range<usize>)> {
memchr2_iter(b'\n', b'\r', text.as_bytes())
.map(|i| i..i + 1)
.coalesce(|a, b| {
if a.end == b.start {
Ok(a.start..b.end)
} else {
Err((a, b))
}
})
.map(|range| {
let level = GRAPHEME_SEGMENTER
.segment_str(text.get(range.start..range.end).unwrap())
.tuple_windows::<(usize, usize)>()
.count();
(
match level {
0 => unreachable!("regex should always match at least one newline"),
n => LineBreaks(n),
},
range,
)
})
.collect()
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd)]
pub struct LineBreaks(usize);
impl SemanticLevel for LineBreaks {}