#!/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()