#!/usr/bin/env python3
"""
Parent-Child Chunking Script using Docling for document conversion.

Converts documents (PDF, XLSX, DOCX, PPTX) to markdown using Docling,
applies parent-child chunking strategy, and outputs to JSON.

Parent-child chunking creates two levels of chunks:
- Parent chunks: Larger segments (800-1500 tokens) for context
- Child chunks: Smaller segments (200-400 tokens) for precise retrieval

Usage:
    python scripts/parent_child_chunker.py input.pdf -o output.json
    python scripts/parent_child_chunker.py input.docx --parent-size 1500 --child-size 400
"""

import argparse
import json
import sys
import uuid
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any

# Add project root to path
sys.path.insert(0, str(Path(__file__).parent.parent))

import tiktoken
from docling.document_converter import DocumentConverter
from langchain_text_splitters import (
    MarkdownHeaderTextSplitter,
    RecursiveCharacterTextSplitter,
)


@dataclass
class ChunkConfig:
    """Configuration for parent-child chunking."""

    parent_chunk_size: int = 1500  # tokens
    parent_chunk_overlap: int = 150
    child_chunk_size: int = 400  # tokens
    child_chunk_overlap: int = 50
    # Approximate chars per token (for RecursiveCharacterTextSplitter)
    chars_per_token: float = 4.0


@dataclass
class ChildChunk:
    """A child chunk with reference to its parent."""

    id: str
    parent_id: str
    text: str
    token_count: int
    position_in_parent: int
    metadata: dict = field(default_factory=dict)


@dataclass
class ParentChunk:
    """A parent chunk containing child chunks."""

    id: str
    text: str
    token_count: int
    children: list[ChildChunk] = field(default_factory=list)
    metadata: dict = field(default_factory=dict)


class ParentChildChunker:
    """
    Parent-child chunking strategy for RAG applications.

    Creates hierarchical chunks where:
    - Parent chunks preserve broader context (for LLM generation)
    - Child chunks enable precise retrieval (for embedding search)
    """

    SUPPORTED_EXTENSIONS = {".pdf", ".xlsx", ".xls", ".docx", ".pptx"}

    def __init__(self, config: ChunkConfig | None = None):
        """
        Initialize the chunker with configuration.

        Args:
            config: Chunking configuration. Uses defaults if not provided.
        """
        self.config = config or ChunkConfig()

        # Initialize tokenizer for accurate token counting
        self.encoding = tiktoken.get_encoding("cl100k_base")

        # Initialize document converter
        self.converter = DocumentConverter()

        # Convert token sizes to character sizes for text splitters
        parent_char_size = int(self.config.parent_chunk_size * self.config.chars_per_token)
        parent_char_overlap = int(self.config.parent_chunk_overlap * self.config.chars_per_token)
        child_char_size = int(self.config.child_chunk_size * self.config.chars_per_token)
        child_char_overlap = int(self.config.child_chunk_overlap * self.config.chars_per_token)

        # Markdown header splitter for initial structure-aware splitting
        self.header_splitter = MarkdownHeaderTextSplitter(
            headers_to_split_on=[
                ("#", "h1"),
                ("##", "h2"),
                ("###", "h3"),
                ("####", "h4"),
            ],
            strip_headers=False,
        )

        # Parent chunker (larger chunks for context)
        self.parent_splitter = RecursiveCharacterTextSplitter(
            chunk_size=parent_char_size,
            chunk_overlap=parent_char_overlap,
            separators=["\n\n", "\n", ". ", " ", ""],
            length_function=len,
        )

        # Child chunker (smaller chunks for retrieval)
        self.child_splitter = RecursiveCharacterTextSplitter(
            chunk_size=child_char_size,
            chunk_overlap=child_char_overlap,
            separators=["\n\n", "\n", ". ", " ", ""],
            length_function=len,
        )

    def count_tokens(self, text: str) -> int:
        """Count tokens in text using tiktoken."""
        return len(self.encoding.encode(text))

    def convert_to_markdown(self, file_path: Path) -> str:
        """
        Convert document to markdown using Docling.

        Args:
            file_path: Path to the input document

        Returns:
            Markdown string representation of the document

        Raises:
            ValueError: If file type is not supported
        """
        suffix = file_path.suffix.lower()
        if suffix not in self.SUPPORTED_EXTENSIONS:
            raise ValueError(
                f"Unsupported file type: {suffix}. "
                f"Supported: {', '.join(self.SUPPORTED_EXTENSIONS)}"
            )

        print(f"Converting {file_path.name} to markdown...")
        result = self.converter.convert(file_path)
        markdown = result.document.export_to_markdown()
        print(f"Converted to {len(markdown)} characters of markdown")

        return markdown

    def create_parent_chunks(self, markdown: str, source_filename: str) -> list[ParentChunk]:
        """
        Create parent chunks from markdown, respecting document structure.

        Args:
            markdown: Markdown content
            source_filename: Original filename for metadata

        Returns:
            List of ParentChunk objects
        """
        parent_chunks = []

        # First, split by headers to respect document structure
        header_splits = self.header_splitter.split_text(markdown)

        for header_doc in header_splits:
            content = header_doc.page_content
            header_metadata = header_doc.metadata

            # If content is too large, split further
            if len(content) > self.config.parent_chunk_size * self.config.chars_per_token:
                sub_splits = self.parent_splitter.split_text(content)
                for i, sub_content in enumerate(sub_splits):
                    parent_id = str(uuid.uuid4())
                    parent_chunks.append(
                        ParentChunk(
                            id=parent_id,
                            text=sub_content,
                            token_count=self.count_tokens(sub_content),
                            metadata={
                                "source_filename": source_filename,
                                "headers": header_metadata,
                                "split_index": i,
                            },
                        )
                    )
            else:
                parent_id = str(uuid.uuid4())
                parent_chunks.append(
                    ParentChunk(
                        id=parent_id,
                        text=content,
                        token_count=self.count_tokens(content),
                        metadata={
                            "source_filename": source_filename,
                            "headers": header_metadata,
                        },
                    )
                )

        return parent_chunks

    def create_child_chunks(self, parent: ParentChunk) -> list[ChildChunk]:
        """
        Create child chunks from a parent chunk.

        Args:
            parent: Parent chunk to split

        Returns:
            List of ChildChunk objects with references to parent
        """
        child_chunks = []

        # Split parent text into smaller child chunks
        child_texts = self.child_splitter.split_text(parent.text)

        for i, child_text in enumerate(child_texts):
            child_id = str(uuid.uuid4())
            child_chunks.append(
                ChildChunk(
                    id=child_id,
                    parent_id=parent.id,
                    text=child_text,
                    token_count=self.count_tokens(child_text),
                    position_in_parent=i,
                    metadata={
                        **parent.metadata,
                        "child_index": i,
                        "total_children": len(child_texts),
                    },
                )
            )

        return child_chunks

    def chunk_document(self, file_path: Path) -> dict[str, Any]:
        """
        Process a document and create parent-child chunks.

        Args:
            file_path: Path to the input document

        Returns:
            Dictionary with markdown, parent chunks, child chunks, and metadata
        """
        file_path = Path(file_path)

        # Convert to markdown
        markdown = self.convert_to_markdown(file_path)

        # Create parent chunks
        print("Creating parent chunks...")
        parent_chunks = self.create_parent_chunks(markdown, file_path.name)
        print(f"Created {len(parent_chunks)} parent chunks")

        # Create child chunks for each parent
        print("Creating child chunks...")
        all_children = []
        for parent in parent_chunks:
            children = self.create_child_chunks(parent)
            parent.children = children
            all_children.extend(children)
        print(f"Created {len(all_children)} child chunks")

        # Calculate statistics
        parent_tokens = sum(p.token_count for p in parent_chunks)
        child_tokens = sum(c.token_count for c in all_children)

        result = {
            "source_file": str(file_path),
            "source_filename": file_path.name,
            "markdown": markdown,
            "config": {
                "parent_chunk_size": self.config.parent_chunk_size,
                "parent_chunk_overlap": self.config.parent_chunk_overlap,
                "child_chunk_size": self.config.child_chunk_size,
                "child_chunk_overlap": self.config.child_chunk_overlap,
            },
            "statistics": {
                "markdown_chars": len(markdown),
                "markdown_tokens": self.count_tokens(markdown),
                "parent_chunk_count": len(parent_chunks),
                "child_chunk_count": len(all_children),
                "total_parent_tokens": parent_tokens,
                "total_child_tokens": child_tokens,
                "avg_parent_tokens": parent_tokens / len(parent_chunks) if parent_chunks else 0,
                "avg_child_tokens": child_tokens / len(all_children) if all_children else 0,
            },
            "parent_chunks": [
                {
                    "id": p.id,
                    "text": p.text,
                    "token_count": p.token_count,
                    "metadata": p.metadata,
                    "child_ids": [c.id for c in p.children],
                }
                for p in parent_chunks
            ],
            "child_chunks": [
                {
                    "id": c.id,
                    "parent_id": c.parent_id,
                    "text": c.text,
                    "token_count": c.token_count,
                    "position_in_parent": c.position_in_parent,
                    "metadata": c.metadata,
                }
                for c in all_children
            ],
        }

        return result


def main():
    """Main entry point for the script."""
    parser = argparse.ArgumentParser(
        description="Convert documents to markdown and apply parent-child chunking.",
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog="""
Examples:
    python scripts/parent_child_chunker.py document.pdf
    python scripts/parent_child_chunker.py report.docx -o chunks.json
    python scripts/parent_child_chunker.py data.xlsx --parent-size 1000 --child-size 300
        """,
    )
    parser.add_argument("input_file", type=Path, help="Input document (PDF, XLSX, DOCX, PPTX)")
    parser.add_argument(
        "-o",
        "--output",
        type=Path,
        help="Output JSON file (default: <input_name>_chunks.json)",
    )
    parser.add_argument(
        "--parent-size",
        type=int,
        default=1500,
        help="Parent chunk size in tokens (default: 1500)",
    )
    parser.add_argument(
        "--parent-overlap",
        type=int,
        default=150,
        help="Parent chunk overlap in tokens (default: 150)",
    )
    parser.add_argument(
        "--child-size",
        type=int,
        default=400,
        help="Child chunk size in tokens (default: 400)",
    )
    parser.add_argument(
        "--child-overlap",
        type=int,
        default=50,
        help="Child chunk overlap in tokens (default: 50)",
    )
    parser.add_argument(
        "--markdown-only",
        action="store_true",
        help="Only output markdown, skip chunking",
    )

    args = parser.parse_args()

    # Validate input file
    if not args.input_file.exists():
        print(f"Error: Input file not found: {args.input_file}")
        sys.exit(1)

    # Set output path
    if args.output:
        output_path = args.output
    else:
        output_path = args.input_file.with_name(f"{args.input_file.stem}_chunks.json")

    # Create config
    config = ChunkConfig(
        parent_chunk_size=args.parent_size,
        parent_chunk_overlap=args.parent_overlap,
        child_chunk_size=args.child_size,
        child_chunk_overlap=args.child_overlap,
    )

    # Initialize chunker
    chunker = ParentChildChunker(config)

    if args.markdown_only:
        # Only convert to markdown
        markdown = chunker.convert_to_markdown(args.input_file)
        md_output_path = args.input_file.with_suffix(".md")
        md_output_path.write_text(markdown, encoding="utf-8")
        print(f"\nMarkdown saved to: {md_output_path}")
    else:
        # Full processing
        print(f"\nProcessing: {args.input_file}")
        print(f"Config: parent={config.parent_chunk_size}t, child={config.child_chunk_size}t")
        print("-" * 60)

        result = chunker.chunk_document(args.input_file)

        # Save to JSON
        with open(output_path, "w", encoding="utf-8") as f:
            json.dump(result, f, indent=2, ensure_ascii=False)

        # Print summary
        stats = result["statistics"]
        print("-" * 60)
        print("Summary:")
        print(f"  Markdown: {stats['markdown_chars']:,} chars, {stats['markdown_tokens']:,} tokens")
        print(f"  Parent chunks: {stats['parent_chunk_count']} (avg {stats['avg_parent_tokens']:.0f} tokens)")
        print(f"  Child chunks: {stats['child_chunk_count']} (avg {stats['avg_child_tokens']:.0f} tokens)")
        print(f"\nOutput saved to: {output_path}")


if __name__ == "__main__":
    main()
