#!/usr/bin/env python """Chunk ingested documents into semantically meaningful segments with relationship tracking. Uses LangChain's RecursiveCharacterTextSplitter for semantic-aware chunking. """ from __future__ import annotations import argparse import json from collections import defaultdict from datetime import datetime, timezone from pathlib import Path from typing import Dict, Iterator, List, Tuple from langchain_text_splitters import RecursiveCharacterTextSplitter def create_text_splitter(chunk_size: int, chunk_overlap: int) -> RecursiveCharacterTextSplitter: """Create a RecursiveCharacterTextSplitter with semantic separators.""" return RecursiveCharacterTextSplitter( chunk_size=chunk_size, chunk_overlap=chunk_overlap, length_function=len, is_separator_regex=False, separators=[ "\n\n", # Paragraph breaks "\n", # Line breaks ". ", # Sentence endings "? ", # Question endings "! ", # Exclamation endings "; ", # Semicolons ", ", # Commas " ", # Words "", # Characters (fallback) ], ) def chunk_text_with_offsets( text: str, splitter: RecursiveCharacterTextSplitter, ) -> List[Tuple[str, int, int]]: """Split text using LangChain splitter and compute character offsets. Returns list of (content, char_offset_start, char_offset_end) tuples. """ if not text.strip(): return [] chunks = splitter.split_text(text) if not chunks: return [] results = [] search_start = 0 for chunk_content in chunks: idx = text.find(chunk_content, search_start) if idx == -1: idx = text.find(chunk_content) if idx == -1: idx = search_start start = idx end = start + len(chunk_content) results.append((chunk_content, start, end)) search_start = start + 1 return results def process_record( record: Dict, splitter: RecursiveCharacterTextSplitter, chunk_size: int, global_chunk_index: int, ) -> Tuple[List[Dict], int]: """Convert an ingestion record into one or more chunks. Returns (list of chunk dicts, updated global_chunk_index). """ doc_id = record.get("doc_id", "unknown") content = record.get("content", "") source = record.get("source", {}) modality = source.get("modality", "text") if modality in ("audio", "video", "image") or len(content) <= chunk_size: chunk = { "chunk_id": f"{doc_id}_chunk_{global_chunk_index:06d}", "doc_id": doc_id, "document_name": source.get("filename", ""), "section": record.get("section"), "page": record.get("page"), "timestamp_start": source.get("timestamp_start"), "timestamp_end": source.get("timestamp_end"), "related_chunks": { "prev": None, "next": None, "same_section": [], "same_document": [], }, "content": content, "source": { "original_record_id": source.get("record_id"), "path": source.get("path"), "filetype": source.get("filetype"), "modality": modality, "char_offset_start": 0, "char_offset_end": len(content), }, } return [chunk], global_chunk_index + 1 text_chunks = chunk_text_with_offsets(content, splitter) chunks = [] for chunk_content, start, end in text_chunks: chunk = { "chunk_id": f"{doc_id}_chunk_{global_chunk_index:06d}", "doc_id": doc_id, "document_name": source.get("filename", ""), "section": record.get("section"), "page": record.get("page"), "timestamp_start": None, "timestamp_end": None, "related_chunks": { "prev": None, "next": None, "same_section": [], "same_document": [], }, "content": chunk_content, "source": { "original_record_id": source.get("record_id"), "path": source.get("path"), "filetype": source.get("filetype"), "modality": modality, "char_offset_start": start, "char_offset_end": end, }, } chunks.append(chunk) global_chunk_index += 1 return chunks, global_chunk_index def build_chunk_relationships(chunks: List[Dict], section_window: int = 5) -> None: """Add prev/next and same_section relationships to chunks in-place.""" for i, chunk in enumerate(chunks): if i > 0: chunk["related_chunks"]["prev"] = chunks[i - 1]["chunk_id"] if i < len(chunks) - 1: chunk["related_chunks"]["next"] = chunks[i + 1]["chunk_id"] section_groups: Dict[Tuple[str, str], List[int]] = defaultdict(list) for i, chunk in enumerate(chunks): key = (chunk["doc_id"], chunk["section"]) section_groups[key].append(i) for indices in section_groups.values(): for i, chunk_idx in enumerate(indices): same_section = [] for j in range(max(0, i - section_window), min(len(indices), i + section_window + 1)): if j != i: same_section.append(chunks[indices[j]]["chunk_id"]) chunks[chunk_idx]["related_chunks"]["same_section"] = same_section doc_groups: Dict[str, List[int]] = defaultdict(list) for i, chunk in enumerate(chunks): doc_groups[chunk["doc_id"]].append(i) for indices in doc_groups.values(): if len(indices) <= 10: for chunk_idx in indices: same_doc = [chunks[i]["chunk_id"] for i in indices if i != chunk_idx] chunks[chunk_idx]["related_chunks"]["same_document"] = same_doc else: for chunk_idx in indices: same_doc = [] if indices[0] != chunk_idx: same_doc.append(chunks[indices[0]]["chunk_id"]) if indices[-1] != chunk_idx: same_doc.append(chunks[indices[-1]]["chunk_id"]) mid = len(indices) // 2 if indices[mid] != chunk_idx: same_doc.append(chunks[indices[mid]]["chunk_id"]) chunks[chunk_idx]["related_chunks"]["same_document"] = same_doc def load_records(input_path: Path) -> Iterator[Dict]: """Stream records from a JSONL file.""" with open(input_path, "r", encoding="utf-8") as f: for line in f: line = line.strip() if line: yield json.loads(line) def chunk_file( input_path: Path, output_dir: Path, chunk_size: int, chunk_overlap: int, ) -> Dict: """Process a single JSONL file and produce chunked output.""" output_dir.mkdir(parents=True, exist_ok=True) output_path = output_dir / "chunked_documents.jsonl" splitter = create_text_splitter(chunk_size, chunk_overlap) all_chunks = [] global_chunk_index = 0 records_processed = 0 for record in load_records(input_path): record_chunks, global_chunk_index = process_record( record, splitter, chunk_size, global_chunk_index ) all_chunks.extend(record_chunks) records_processed += 1 build_chunk_relationships(all_chunks) with open(output_path, "w", encoding="utf-8") as f: for chunk in all_chunks: f.write(json.dumps(chunk, ensure_ascii=True) + "\n") stats = { "input_file": str(input_path), "output_file": str(output_path), "generated_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"), "config": { "chunk_size": chunk_size, "chunk_overlap": chunk_overlap, "splitter": "RecursiveCharacterTextSplitter", }, "counts": { "records_processed": records_processed, "chunks_produced": len(all_chunks), }, } stats_path = output_dir / "chunking_stats.json" with open(stats_path, "w", encoding="utf-8") as f: json.dump(stats, f, indent=2, ensure_ascii=True) return stats def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Chunk ingested documents using LangChain RecursiveCharacterTextSplitter." ) parser.add_argument( "--input", required=True, help="Path to input JSONL file from data ingestion.", ) parser.add_argument( "--output-dir", required=True, help="Output directory for chunked JSONL and stats.", ) parser.add_argument( "--chunk-size", type=int, default=2000, help="Maximum characters per chunk (default: 2000).", ) parser.add_argument( "--chunk-overlap", type=int, default=200, help="Overlap characters between chunks (default: 200).", ) return parser.parse_args() def main() -> None: args = parse_args() stats = chunk_file( input_path=Path(args.input), output_dir=Path(args.output_dir), chunk_size=args.chunk_size, chunk_overlap=args.chunk_overlap, ) print(f"Chunking complete: {stats['counts']['records_processed']} records -> " f"{stats['counts']['chunks_produced']} chunks") print(f"Output: {stats['output_file']}") if __name__ == "__main__": main()