import hashlib
import re
from typing import Any, Dict, List, Optional

from app.core.logger import logger
from app.models.chunk import DocChunk
from app.models.document import RepoDocumentSection


class TextChunker:
    """
    Splits sections into deterministic embedding-ready chunks.
    Uses paragraph-aware boundaries and token/character thresholds.
    """

    def __init__(
        self,
        chunk_size: int = 1000,
        chunk_overlap: int = 200,
        max_tokens: int = 220,
    ):
        self.chunk_size = chunk_size
        self.chunk_overlap = chunk_overlap
        self.max_tokens = max_tokens

    def chunk_sections(
        self,
        sections: List[RepoDocumentSection],
        repo_id: str,
        doc_path: str,
        snapshot_metadata: Optional[Dict[str, Any]] = None,
        section_entities: Optional[Dict[str, Dict[str, int]]] = None,
    ) -> List[DocChunk]:
        chunks: List[DocChunk] = []
        chunk_order = 0
        snapshot_metadata = snapshot_metadata or {}
        section_entities = section_entities or {}
        for section_idx, section in enumerate(sections):
            if not section.content.strip():
                continue
            try:
                section_chunks = self._split_text(section.content)
            except Exception as exc:
                logger.bind(repo_id=repo_id, doc_path=doc_path, section=section.section_title).warning(
                    f"Failed to split section into chunks: {exc}"
                )
                continue

            for idx, chunk_text in enumerate(section_chunks):
                cleaned = chunk_text.strip()
                if not cleaned:
                    continue

                try:
                    chunk_type = self._infer_chunk_type(section.content_types)
                    token_count = self._token_count(cleaned)
                    chunk_hash = self._build_chunk_hash(
                        repo_id,
                        doc_path,
                        section.heading_path,
                        section.section_title,
                        idx,
                        cleaned,
                    )
                    chunk_id = self._build_chunk_id(repo_id, doc_path, chunk_hash)

                    chunks.append(
                        DocChunk(
                            chunk_id=chunk_id,
                            repo_id=repo_id,
                            snapshot_id=snapshot_metadata.get("snapshot_id"),
                            doc_path=doc_path,
                            section_title=section.section_title,
                            chunk_text=cleaned,
                            chunk_hash=chunk_hash,
                            source_type="repo_doc",
                            metadata={
                                "chunk_id": chunk_id,
                                "repo_id": repo_id,
                                "doc_path": doc_path,
                                "chunk_type": chunk_type,
                                "source_type": "repo_doc",
                                "section_title": section.section_title,
                                "section_level": section.section_level,
                                "heading": section.heading,
                                "parent_section": section.parent_section,
                                "heading_path": section.heading_path,
                                "start_line": section.start_line,
                                "end_line": section.end_line,
                                "content_start_line": section.content_start_line,
                                "content_end_line": section.content_end_line,
                                "content_types": section.content_types,
                                "chunk_index": idx,
                                "section_index": section_idx,
                                "chunk_order": chunk_order,
                                "char_count": len(cleaned),
                                "token_count": token_count,
                                "snapshot": {
                                    "commit_hash": snapshot_metadata.get("commit_hash"),
                                    "snapshot_id": snapshot_metadata.get("snapshot_id"),
                                },
                                "markdown_entities": section_entities.get(
                                    section.section_title,
                                    {"links": 0, "code": 0, "tables": 0},
                                ),
                            },
                        )
                    )
                except Exception as exc:
                    logger.bind(
                        repo_id=repo_id,
                        doc_path=doc_path,
                        section=section.section_title,
                        chunk_index=idx,
                    ).warning(f"Failed to build chunk metadata: {exc}")
                    continue
                chunk_order += 1
        return chunks

    def _build_chunk_hash(
        self,
        repo_id: str,
        doc_path: str,
        heading_path: Optional[str],
        section_title: str,
        chunk_index: int,
        chunk_text: str,
    ) -> str:
        canonical = "|".join(
            [
                repo_id,
                doc_path,
                heading_path or section_title,
                str(chunk_index),
                chunk_text,
            ]
        )
        return hashlib.sha256(canonical.encode()).hexdigest()

    def _build_chunk_id(self, repo_id: str, doc_path: str, chunk_hash: str) -> str:
        seed = f"{repo_id}|{doc_path}|{chunk_hash}"
        return f"chk_{hashlib.sha256(seed.encode()).hexdigest()[:24]}"

    def _infer_chunk_type(self, content_types: List[str]) -> str:
        if not content_types:
            return "text"
        if "table" in content_types:
            return "table"
        if "code_block" in content_types:
            return "code"
        return "text"

    def _split_text(self, text: str) -> List[str]:
        normalized = text.strip()
        if not normalized:
            return []

        paragraphs = [part.strip() for part in re.split(r"\n\s*\n", normalized) if part.strip()]
        if not paragraphs:
            return [normalized]

        chunks: List[str] = []
        current_parts: List[str] = []

        for paragraph in paragraphs:
            if self._fits(current_parts, paragraph):
                current_parts.append(paragraph)
                continue

            if current_parts:
                chunks.append("\n\n".join(current_parts).strip())
                current_parts = []

            if self._token_count(paragraph) <= self.max_tokens and len(paragraph) <= self.chunk_size:
                current_parts.append(paragraph)
                continue

            chunks.extend(self._split_large_paragraph(paragraph))

        if current_parts:
            chunks.append("\n\n".join(current_parts).strip())

        return chunks

    def _split_large_paragraph(self, paragraph: str) -> List[str]:
        sentences = re.split(r"(?<=[.!?])\s+", paragraph)
        sentences = [item.strip() for item in sentences if item.strip()]
        if not sentences:
            return self._split_by_char_window(paragraph)

        chunks: List[str] = []
        current: List[str] = []

        for sentence in sentences:
            if self._fits(current, sentence, separator=" "):
                current.append(sentence)
                continue

            if current:
                chunks.append(" ".join(current).strip())
                current = []

            if self._token_count(sentence) <= self.max_tokens and len(sentence) <= self.chunk_size:
                current.append(sentence)
                continue

            chunks.extend(self._split_by_char_window(sentence))

        if current:
            chunks.append(" ".join(current).strip())

        return chunks

    def _split_by_char_window(self, text: str) -> List[str]:
        value = text.strip()
        if not value:
            return []

        if len(value) <= self.chunk_size and self._token_count(value) <= self.max_tokens:
            return [value]

        chunks: List[str] = []
        start = 0
        overlap = max(0, min(self.chunk_overlap, self.chunk_size // 4))

        while start < len(value):
            end = min(start + self.chunk_size, len(value))
            window = value[start:end]

            while self._token_count(window) > self.max_tokens and len(window) > 1:
                end -= max(1, len(window) // 10)
                window = value[start:end]

            boundary = window.rfind(" ")
            if boundary > 20 and end < len(value):
                end = start + boundary
                window = value[start:end]

            chunk = window.strip()
            if chunk:
                chunks.append(chunk)

            if end >= len(value):
                break
            start = max(start + 1, end - overlap)

        return chunks

    def _token_count(self, text: str) -> int:
        return len(re.findall(r"\S+", text))

    def _fits(self, current_parts: List[str], candidate: str, separator: str = "\n\n") -> bool:
        if not current_parts:
            return self._token_count(candidate) <= self.max_tokens and len(candidate) <= self.chunk_size

        merged = separator.join(current_parts + [candidate])
        return self._token_count(merged) <= self.max_tokens and len(merged) <= self.chunk_size
