Cognitive-rag / 02_chunking / chunk_documents.py
chunk_documents.py
Raw
#!/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()