use crate::ast::KindData;
use crate::parser::{NoParserOptions, Parser, ParserExtension, ParserExtensionFn, TableAstTransformer, TableParagraphTransformer};
use crate::text::BasicReader;
use std::fs::File;
use memmap2::Mmap;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BlockType {
Heading,
Paragraph,
CodeBlock,
List,
Table,
Blockquote,
Diagram,
Other,
}
impl BlockType {
pub fn as_str(&self) -> &'static str {
match self {
BlockType::Heading => "Heading",
BlockType::Paragraph => "Paragraph",
BlockType::CodeBlock => "CodeBlock",
BlockType::List => "List",
BlockType::Table => "Table",
BlockType::Blockquote => "Blockquote",
BlockType::Diagram => "Diagram",
BlockType::Other => "Other",
}
}
}
fn parse_block_type(s: &str) -> BlockType {
match s {
"Heading" => BlockType::Heading,
"Paragraph" => BlockType::Paragraph,
"CodeBlock" => BlockType::CodeBlock,
"List" => BlockType::List,
"Table" => BlockType::Table,
"Blockquote" => BlockType::Blockquote,
"Diagram" => BlockType::Diagram,
_ => BlockType::Other,
}
}
#[derive(Clone, Copy, Debug)]
struct NodeInfo {
block_type: BlockType,
start: usize,
end: usize,
}
fn extract_nodes(source: &str) -> Vec<NodeInfo> {
let gfm_ext = ParserExtensionFn::new(|p: &mut Parser| {
p.add_ast_transformer(TableAstTransformer::new, NoParserOptions, 0);
p.add_paragraph_transformer(TableParagraphTransformer::new, NoParserOptions, 200);
});
let diagram_ext = crate::diagram::diagram_parser_extension(
crate::diagram::DiagramParserOptions::default(),
);
let parser = Parser::with_extensions(
crate::parser::Options::default(),
gfm_ext.and(diagram_ext),
);
let mut reader = BasicReader::new(source);
let (arena, doc_ref) = parser.parse(&mut reader);
let mut starts: Vec<(BlockType, usize)> = Vec::new();
let mut child = arena[doc_ref].first_child();
while let Some(cref) = child {
let node = &arena[cref];
if let Some(start) = node.pos() {
let block_type = match node.kind_data() {
KindData::Heading(_) => BlockType::Heading,
KindData::Paragraph(_) => BlockType::Paragraph,
KindData::CodeBlock(_) => BlockType::CodeBlock,
KindData::List(_) => BlockType::List,
KindData::Table(_) => BlockType::Table,
KindData::Blockquote(_) => BlockType::Blockquote,
KindData::Extension(ref d) => {
if (d.as_ref() as &dyn std::any::Any).is::<crate::diagram::Diagram>() {
BlockType::Diagram
} else {
BlockType::Other
}
}
_ => BlockType::Other,
};
starts.push((block_type, start));
}
child = arena[cref].next_sibling();
}
let n = starts.len();
let mut nodes = Vec::with_capacity(n);
for i in 0..n {
let (block_type, start) = starts[i];
let end = if i + 1 < n { starts[i + 1].1 } else { source.len() };
nodes.push(NodeInfo { block_type, start, end });
}
nodes
}
pub enum TextSource {
Owned(String),
Mapped { file: File, mmap: Mmap },
}
#[derive(Debug)]
pub enum ChunkerError {
Io(std::io::Error),
InvalidUtf8(String),
}
impl std::fmt::Display for ChunkerError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ChunkerError::Io(e) => write!(f, "{e}"),
ChunkerError::InvalidUtf8(msg) => write!(f, "{msg}"),
}
}
}
impl std::error::Error for ChunkerError {}
impl TextSource {
pub fn owned(text: String) -> Self {
TextSource::Owned(text)
}
pub fn from_mmap_path(path: &str) -> Result<Self, ChunkerError> {
let file = File::open(path).map_err(ChunkerError::Io)?;
let mmap = unsafe { Mmap::map(&file) }.map_err(ChunkerError::Io)?;
Ok(TextSource::Mapped { file, mmap })
}
pub fn str_checked(&self) -> Result<&str, ChunkerError> {
match self {
TextSource::Owned(s) => Ok(s.as_str()),
TextSource::Mapped { mmap, .. } => std::str::from_utf8(&mmap[..])
.map_err(|e| ChunkerError::InvalidUtf8(format!("Invalid UTF-8 in file: {e}"))),
}
}
}
#[derive(Clone, Debug)]
pub struct Chunk {
pub text: String,
pub block_type: BlockType,
pub start_offset: usize,
pub end_offset: usize,
}
fn source_str(source: &TextSource) -> &str {
match source {
TextSource::Owned(s) => s.as_str(),
TextSource::Mapped { mmap, .. } => unsafe {
std::str::from_utf8_unchecked(&mmap[..])
},
}
}
pub struct MarkdownChunker {
source: TextSource,
nodes: Vec<NodeInfo>,
index: usize,
current_header: Option<(usize, usize)>,
}
impl MarkdownChunker {
pub fn new(text: String) -> Self {
let nodes = extract_nodes(&text);
Self {
source: TextSource::Owned(text),
nodes,
index: 0,
current_header: None,
}
}
pub fn from_file(path: &str) -> Result<Self, ChunkerError> {
let bytes = std::fs::read(path).map_err(ChunkerError::Io)?;
let text = String::from_utf8(bytes).map_err(|e| {
ChunkerError::InvalidUtf8(format!("Invalid UTF-8 in file: {e}"))
})?;
Ok(Self::new(text))
}
pub fn from_text_source(source: TextSource) -> Result<Self, ChunkerError> {
source.str_checked()?;
let nodes = {
let text = source.str_checked()?;
extract_nodes(text)
};
Ok(Self {
source,
nodes,
index: 0,
current_header: None,
})
}
fn text(&self) -> &str {
source_str(&self.source)
}
pub fn next_chunk(&mut self) -> Option<String> {
let text = source_str(&self.source);
while self.index < self.nodes.len() {
let node = self.nodes[self.index]; self.index += 1;
let raw = text[node.start..node.end].trim_end();
match node.block_type {
BlockType::Heading => {
let end = node.start + raw.len();
self.current_header = Some((node.start, end));
}
BlockType::Paragraph
| BlockType::CodeBlock
| BlockType::List
| BlockType::Table
| BlockType::Blockquote
| BlockType::Diagram => {
return Some(raw.to_string());
}
BlockType::Other => {
}
}
}
None
}
pub fn current_header(&self) -> Option<String> {
let text = source_str(&self.source);
self.current_header.map(|(s, e)| text[s..e].to_string())
}
pub fn node_count(&self) -> usize {
self.nodes.len()
}
pub fn get_chunks(&mut self) -> Vec<Chunk> {
let text = source_str(&self.source);
let mut chunks = Vec::new();
for node in &self.nodes {
match node.block_type {
BlockType::Heading => {
let raw = text[node.start..node.end].trim_end();
let end = node.start + raw.len();
self.current_header = Some((node.start, end));
}
BlockType::Paragraph
| BlockType::CodeBlock
| BlockType::List
| BlockType::Table
| BlockType::Blockquote
| BlockType::Diagram => {
let raw = text[node.start..node.end].trim_end();
chunks.push(Chunk {
text: raw.to_string(),
block_type: node.block_type,
start_offset: node.start,
end_offset: node.start + raw.len(),
});
}
BlockType::Other => {
}
}
}
chunks
}
pub fn get_chunks_with_context(&mut self) -> Vec<Chunk> {
let text = source_str(&self.source);
let mut chunks = Vec::new();
for node in &self.nodes {
match node.block_type {
BlockType::Heading => {
let raw = text[node.start..node.end].trim_end();
let end = node.start + raw.len();
self.current_header = Some((node.start, end));
}
BlockType::Paragraph
| BlockType::CodeBlock
| BlockType::List
| BlockType::Table
| BlockType::Blockquote
| BlockType::Diagram => {
let raw = text[node.start..node.end].trim_end();
let text_with_context: String = match self.current_header {
Some((h_start, h_end)) => {
format!("{}\n\n{}", &text[h_start..h_end], raw)
}
None => raw.to_string(),
};
chunks.push(Chunk {
text: text_with_context,
block_type: node.block_type,
start_offset: node.start,
end_offset: node.start + raw.len(),
});
}
BlockType::Other => {
}
}
}
chunks
}
pub fn get_all_chunks(&mut self) -> Vec<Chunk> {
let text = source_str(&self.source);
let mut chunks = Vec::new();
for node in &self.nodes {
match node.block_type {
BlockType::Heading => {
let raw = text[node.start..node.end].trim_end();
let end = node.start + raw.len();
self.current_header = Some((node.start, end));
chunks.push(Chunk {
text: raw.to_string(),
block_type: BlockType::Heading,
start_offset: node.start,
end_offset: node.start + raw.len(),
});
}
BlockType::Paragraph
| BlockType::CodeBlock
| BlockType::List
| BlockType::Table
| BlockType::Blockquote
| BlockType::Diagram => {
let raw = text[node.start..node.end].trim_end();
chunks.push(Chunk {
text: raw.to_string(),
block_type: node.block_type,
start_offset: node.start,
end_offset: node.start + raw.len(),
});
}
BlockType::Other => {
}
}
}
chunks
}
pub fn get_bare_chunks(&mut self) -> Vec<String> {
let text = source_str(&self.source);
let mut chunks = Vec::new();
for node in &self.nodes {
match node.block_type {
BlockType::Heading => {
let raw = text[node.start..node.end].trim_end();
let end = node.start + raw.len();
self.current_header = Some((node.start, end));
}
BlockType::Paragraph
| BlockType::CodeBlock
| BlockType::List
| BlockType::Table
| BlockType::Blockquote
| BlockType::Diagram => {
let raw = text[node.start..node.end].trim_end();
chunks.push(raw.to_string());
}
BlockType::Other => {
}
}
}
chunks
}
pub fn compute_overlap_payloads(&mut self, overlap_words: usize) -> Vec<(String, String)> {
let text = source_str(&self.source);
let mut payloads = Vec::new();
let mut prev_tail: String = String::new();
let mut chunk_index = 0usize;
for node in &self.nodes {
match node.block_type {
BlockType::Heading => {
let raw = text[node.start..node.end].trim_end();
let end = node.start + raw.len();
self.current_header = Some((node.start, end));
}
BlockType::Paragraph
| BlockType::CodeBlock
| BlockType::List
| BlockType::Table
| BlockType::Blockquote
| BlockType::Diagram => {
let raw = text[node.start..node.end].trim_end();
let embed_text: String = if !prev_tail.is_empty() {
format!("{}\n\n{}", prev_tail, raw)
} else {
raw.to_string()
};
payloads.push((format!("chunk:{}", chunk_index), embed_text));
let words: Vec<&str> = raw.split_whitespace().collect();
if overlap_words > 0 && !words.is_empty() {
let tail_start = if words.len() > overlap_words {
words.len() - overlap_words
} else {
0
};
prev_tail = words[tail_start..].join(" ");
} else {
prev_tail = String::new();
}
chunk_index += 1;
}
BlockType::Other => {
}
}
}
payloads
}
pub fn get_delimiter(prev: &str, curr: &str) -> String {
if prev == "List" && curr == "List" {
"\n".to_string()
} else if prev == "Blockquote" && curr == "Blockquote" {
"\n> ".to_string()
} else {
"\n\n".to_string()
}
}
}
pub fn block_type_from_str(s: &str) -> BlockType {
parse_block_type(s)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn basic_chunking() {
let src = "# Title\n\nPara one.\n\n```rust\ncode\n```\n\n- a\n- b\n";
let mut c = MarkdownChunker::new(src.to_string());
assert!(c.node_count() >= 4);
let chunks = c.get_chunks();
let types: Vec<&str> = chunks.iter().map(|k| k.block_type.as_str()).collect();
assert!(!types.contains(&"Heading"));
assert!(types.contains(&"Paragraph"));
assert!(types.contains(&"CodeBlock"));
assert!(types.contains(&"List"));
}
#[test]
fn get_all_chunks_includes_heading() {
let src = "# Title\n\nBody.\n";
let mut c = MarkdownChunker::new(src.to_string());
let all = c.get_all_chunks();
assert_eq!(all[0].block_type.as_str(), "Heading");
assert_eq!(all[0].text, "# Title");
}
#[test]
fn context_prefix() {
let src = "# H\n\nBody text.\n";
let mut c = MarkdownChunker::new(src.to_string());
let chunks = c.get_chunks_with_context();
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0].text, "# H\n\nBody text.");
}
#[test]
fn current_header_tracks_headings() {
let src = "# First\n\npara\n\n## Second\n\npara2\n";
let mut c = MarkdownChunker::new(src.to_string());
let _ = c.next_chunk(); assert_eq!(c.current_header().as_deref(), Some("# First"));
let _ = c.next_chunk(); assert_eq!(c.current_header().as_deref(), Some("## Second"));
}
#[test]
fn bare_iteration_matches_get_bare_chunks() {
let src = "Intro\n\n# H\n\nA\n\nB\n";
let mut c1 = MarkdownChunker::new(src.to_string());
let mut iterated = Vec::new();
while let Some(s) = c1.next_chunk() {
iterated.push(s);
}
let mut c2 = MarkdownChunker::new(src.to_string());
assert_eq!(iterated, c2.get_bare_chunks());
assert_eq!(iterated, vec!["Intro".to_string(), "A".to_string(), "B".to_string()]);
}
#[test]
fn delimiters() {
assert_eq!(MarkdownChunker::get_delimiter("List", "List"), "\n");
assert_eq!(MarkdownChunker::get_delimiter("Blockquote", "Blockquote"), "\n> ");
assert_eq!(MarkdownChunker::get_delimiter("Paragraph", "List"), "\n\n");
}
#[test]
fn overlap_payloads() {
let src = "one two three four\n\nfive six\n";
let mut c = MarkdownChunker::new(src.to_string());
let payloads = c.compute_overlap_payloads(2);
assert_eq!(payloads.len(), 2);
assert_eq!(payloads[0].0, "chunk:0");
assert_eq!(payloads[0].1, "one two three four");
assert!(payloads[1].1.starts_with("three four\n\nfive six"));
}
#[test]
fn block_type_roundtrip() {
assert_eq!(block_type_from_str("Heading"), BlockType::Heading);
assert_eq!(block_type_from_str("Diagram"), BlockType::Diagram);
assert_eq!(block_type_from_str("whatever"), BlockType::Other);
}
#[test]
fn offsets_index_original_source() {
let src = "# T\n\nHello world.\n";
let mut c = MarkdownChunker::new(src.to_string());
let chunks = c.get_chunks();
let ch = &chunks[0];
assert_eq!(&src[ch.start_offset..ch.end_offset], "Hello world.");
}
}