from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Tuple
import unicodedata
import re

from markdown_it import MarkdownIt
from app.core.logger import logger

# (section_title, heading_level, body_content)
ParsedSection = Tuple[str, int, str]


@dataclass
class MarkdownParseResult:
    sections: List[ParsedSection]
    section_metadata: List[Dict[str, Any]]
    section_segments: List[Dict[str, Any]]
    tokens: List[Dict[str, Any]]
    normalized_content: str
    links: List[Dict[str, Any]]
    code_blocks: List[Dict[str, Any]]
    inline_code: List[Dict[str, Any]]
    tables: List[Dict[str, Any]]


class MarkdownParser:
    """Parses markdown using markdown-it-py and exposes traversable section/token output."""

    def __init__(self) -> None:
        # Keep parsing safe for repository ingestion by disabling raw HTML rendering.
        self._md = MarkdownIt(
            "commonmark",
            {
                "html": False,
                "linkify": False,
                "typographer": False,
            },
        )
        self._md.enable("table")
        self._md.enable("strikethrough")
        self._http_method_re = re.compile(r"^(GET|POST|PUT|PATCH|DELETE|HEAD|OPTIONS)\b")

    def parse(self, content: str) -> MarkdownParseResult:
        normalized = self._normalize_input(content)
        if not normalized.strip():
            return MarkdownParseResult(
                sections=[],
                section_metadata=[],
                section_segments=[],
                tokens=[],
                normalized_content=normalized,
                links=[],
                code_blocks=[],
                inline_code=[],
                tables=[],
            )

        tokens = self._md.parse(normalized)
        token_dicts = [
            {
                "type": token.type,
                "tag": token.tag,
                "content": token.content,
                "map": token.map,
                "nesting": token.nesting,
                "level": token.level,
            }
            for token in tokens
        ]

        sections, section_metadata = self._extract_sections(tokens, normalized)
        section_segments = self._build_section_segments(sections, section_metadata)
        structured = self._extract_structured_elements(tokens)
        return MarkdownParseResult(
            sections=sections,
            section_metadata=section_metadata,
            section_segments=section_segments,
            tokens=token_dicts,
            normalized_content=normalized,
            links=structured["links"],
            code_blocks=structured["code_blocks"],
            inline_code=structured["inline_code"],
            tables=structured["tables"],
        )

    def parse_sections(self, content: str) -> List[ParsedSection]:
        """Returns backward-compatible (title, level, body) tuples."""
        return self.parse(content).sections

    def _normalize_input(self, content: str) -> str:
        if not isinstance(content, str):
            raise ValueError("Markdown content must be a UTF-8 text string")

        normalized = unicodedata.normalize("NFC", content)
        if normalized.startswith("\ufeff"):
            normalized = normalized[1:]

        normalized = normalized.replace("\r\n", "\n").replace("\r", "\n")
        return self._normalize_markdown_lines(normalized)

    def _normalize_markdown_lines(self, content: str) -> str:
        lines = content.split("\n")
        normalized_lines: List[str] = []
        in_code_fence = False
        blank_run = 0
        previous_separator = False

        for line in lines:
            cleaned = self._strip_unsupported_chars(line)
            stripped = cleaned.strip()

            if self._is_fence_boundary(cleaned):
                in_code_fence = not in_code_fence
                normalized_lines.append(cleaned.rstrip())
                blank_run = 0
                previous_separator = False
                continue

            if in_code_fence:
                normalized_lines.append(cleaned)
                continue

            cleaned = cleaned.replace("\t", "    ").rstrip()

            if self._is_separator_line(cleaned):
                if previous_separator:
                    continue
                normalized_lines.append("---")
                blank_run = 0
                previous_separator = True
                continue

            previous_separator = False

            if not stripped:
                blank_run += 1
                if blank_run > 1:
                    continue
                normalized_lines.append("")
                continue

            blank_run = 0
            normalized_lines.append(cleaned)

        return "\n".join(normalized_lines)

    def _is_fence_boundary(self, line: str) -> bool:
        return bool(re.match(r"^\s*(```|~~~)", line))

    def _is_separator_line(self, line: str) -> bool:
        return bool(re.match(r"^\s*([-_*])\1{3,}\s*$", line))

    def _strip_unsupported_chars(self, text: str) -> str:
        chars: List[str] = []
        for ch in text:
            if ch in {"\t", "\n"}:
                chars.append(ch)
                continue
            category = unicodedata.category(ch)
            if category.startswith("C"):
                continue
            chars.append(ch)
        return "".join(chars)

    def _extract_sections(self, tokens: List[Any], content: str) -> Tuple[List[ParsedSection], List[Dict[str, Any]]]:
        lines = content.split("\n")
        headings: List[Dict[str, Any]] = []

        for idx, token in enumerate(tokens):
            if token.type != "heading_open":
                continue

            level = int(token.tag[1]) if token.tag.startswith("h") else 0
            start_line = token.map[0] if token.map else None
            end_line = token.map[1] if token.map else None
            inline = tokens[idx + 1] if idx + 1 < len(tokens) else None
            title = (
                inline.content.strip()
                if inline is not None and getattr(inline, "type", "") == "inline"
                else ""
            )
            headings.append(
                {
                    "title": title or "Untitled",
                    "level": level,
                    "start_line": start_line,
                    "end_line": end_line,
                }
            )

        if not headings:
            fallback_sections, fallback_metadata = self._extract_plaintext_sections(lines)
            if fallback_sections:
                return fallback_sections, fallback_metadata

            body = content.strip()
            if body:
                return [(
                    "Introduction",
                    0,
                    body,
                )], [
                    {
                        "section_title": "Introduction",
                        "section_level": 0,
                        "parent_section": None,
                        "heading_path": "Introduction",
                        "start_line": 1,
                        "end_line": len(lines),
                        "content_start_line": 1,
                        "content_end_line": len(lines),
                        "content_types": ["paragraph"],
                    }
                ]
            return [], []

        sections: List[ParsedSection] = []
        metadata: List[Dict[str, Any]] = []
        first_start = headings[0]["start_line"] if headings[0]["start_line"] is not None else 0
        intro = "\n".join(lines[:first_start]).strip()
        if intro:
            sections.append(("Introduction", 0, intro))
            metadata.append(
                {
                    "section_title": "Introduction",
                    "section_level": 0,
                    "parent_section": None,
                    "heading_path": "Introduction",
                    "start_line": 1,
                    "end_line": first_start,
                    "content_start_line": 1,
                    "content_end_line": self._trimmed_content_end_line(lines, 1, first_start),
                    "content_types": ["paragraph"],
                }
            )

        hierarchy_stack: List[Dict[str, Any]] = []

        for idx, heading in enumerate(headings):
            next_start = (
                headings[idx + 1]["start_line"]
                if idx + 1 < len(headings) and headings[idx + 1]["start_line"] is not None
                else len(lines)
            )
            start = heading["end_line"] if heading["end_line"] is not None else next_start
            body = "\n".join(lines[start:next_start]).strip()
            sections.append((heading["title"], heading["level"], body))

            while hierarchy_stack and hierarchy_stack[-1]["section_level"] >= heading["level"]:
                hierarchy_stack.pop()

            parent_section = (
                hierarchy_stack[-1]["section_title"] if hierarchy_stack else None
            )
            heading_path = (
                f"{hierarchy_stack[-1]['heading_path']} / {heading['title']}"
                if hierarchy_stack
                else heading["title"]
            )
            start_line = (
                (heading["start_line"] + 1)
                if heading["start_line"] is not None
                else 1
            )
            end_line = next_start if next_start >= start_line else start_line
            content_start_line = (
                (heading["end_line"] + 1)
                if heading["end_line"] is not None
                else start_line
            )
            if end_line >= content_start_line:
                content_end_line = self._trimmed_content_end_line(
                    lines, content_start_line, end_line
                )
            else:
                content_end_line = content_start_line

            current_meta = {
                "section_title": heading["title"],
                "section_level": heading["level"],
                "parent_section": parent_section,
                "heading_path": heading_path,
                "start_line": start_line,
                "end_line": end_line,
                "content_start_line": content_start_line,
                "content_end_line": content_end_line,
                "content_types": [],
            }
            metadata.append(current_meta)
            hierarchy_stack.append(current_meta)

        self._associate_content_types(tokens, metadata)
        return sections, metadata

    def _extract_plaintext_sections(self, lines: List[str]) -> Tuple[List[ParsedSection], List[Dict[str, Any]]]:
        heading_indexes = [
            idx for idx, _ in enumerate(lines) if self._is_plaintext_heading(lines, idx)
        ]
        if not heading_indexes:
            return [], []

        sections: List[ParsedSection] = []
        metadata: List[Dict[str, Any]] = []
        first_idx = heading_indexes[0]
        intro = "\n".join(lines[:first_idx]).strip()
        if intro:
            sections.append(("Introduction", 0, intro))
            metadata.append(
                {
                    "section_title": "Introduction",
                    "section_level": 0,
                    "parent_section": None,
                    "heading_path": "Introduction",
                    "start_line": 1,
                    "end_line": first_idx,
                    "content_start_line": 1,
                    "content_end_line": self._trimmed_content_end_line(lines, 1, first_idx),
                    "content_types": ["paragraph"],
                }
            )

        for pos, idx in enumerate(heading_indexes):
            title = lines[idx].strip()
            next_idx = heading_indexes[pos + 1] if pos + 1 < len(heading_indexes) else len(lines)
            body = "\n".join(lines[idx + 1:next_idx]).strip()
            sections.append((title, 1, body))
            metadata.append(
                {
                    "section_title": title,
                    "section_level": 1,
                    "parent_section": None,
                    "heading_path": title,
                    "start_line": idx + 1,
                    "end_line": next_idx,
                    "content_start_line": idx + 2,
                    "content_end_line": self._trimmed_content_end_line(
                        lines, idx + 2, next_idx
                    ),
                    "content_types": ["paragraph"],
                }
            )

        return sections, metadata

    def _extract_structured_elements(self, tokens: List[Any]) -> Dict[str, List[Dict[str, Any]]]:
        try:
            links = self._extract_links(tokens)
        except Exception:
            logger.exception("Failed to extract markdown links")
            links = []

        try:
            code_blocks = self._extract_code_blocks(tokens)
        except Exception:
            logger.exception("Failed to extract markdown code blocks")
            code_blocks = []

        try:
            inline_code = self._extract_inline_code(tokens)
        except Exception:
            logger.exception("Failed to extract markdown inline code")
            inline_code = []

        try:
            tables = self._extract_tables(tokens)
        except Exception:
            logger.exception("Failed to extract markdown tables")
            tables = []

        return {
            "links": links,
            "code_blocks": code_blocks,
            "inline_code": inline_code,
            "tables": tables,
        }

    def _extract_links(self, tokens: List[Any]) -> List[Dict[str, Any]]:
        links: List[Dict[str, Any]] = []
        for token in tokens:
            if token.type != "inline" or not getattr(token, "children", None):
                continue

            line_number = token.map[0] + 1 if token.map else None

            children = token.children
            idx = 0
            while idx < len(children):
                child = children[idx]
                if child.type != "link_open":
                    idx += 1
                    continue

                href = child.attrGet("href") or ""
                text_parts: List[str] = []
                cursor = idx + 1
                while cursor < len(children) and children[cursor].type != "link_close":
                    current = children[cursor]
                    if current.type in {"text", "code_inline"}:
                        text_parts.append(current.content)
                    cursor += 1

                links.append(
                    {
                        "text": "".join(text_parts).strip(),
                        "target": href,
                        "link_type": self._classify_link_type(href),
                        "line": line_number,
                    }
                )
                idx = cursor + 1

        return links

    def _extract_code_blocks(self, tokens: List[Any]) -> List[Dict[str, Any]]:
        code_blocks: List[Dict[str, Any]] = []
        for token in tokens:
            if token.type not in {"fence", "code_block"}:
                continue

            info = (getattr(token, "info", "") or "").strip()
            language = info.split()[0] if info else None
            start_line = token.map[0] + 1 if token.map else None
            end_line = token.map[1] if token.map else None

            code_blocks.append(
                {
                    "language": language,
                    "info": info or None,
                    "content": token.content,
                    "start_line": start_line,
                    "end_line": end_line,
                }
            )

        return code_blocks

    def _extract_inline_code(self, tokens: List[Any]) -> List[Dict[str, Any]]:
        snippets: List[Dict[str, Any]] = []
        for token in tokens:
            if token.type != "inline" or not getattr(token, "children", None):
                continue

            line = token.map[0] + 1 if token.map else None
            for child in token.children:
                if child.type != "code_inline":
                    continue
                snippets.append(
                    {
                        "content": child.content,
                        "line": line,
                    }
                )
        return snippets

    def _extract_tables(self, tokens: List[Any]) -> List[Dict[str, Any]]:
        tables: List[Dict[str, Any]] = []
        idx = 0
        while idx < len(tokens):
            token = tokens[idx]
            if token.type != "table_open":
                idx += 1
                continue

            depth = 1
            cursor = idx + 1
            while cursor < len(tokens) and depth > 0:
                current = tokens[cursor]
                if current.type == "table_open":
                    depth += 1
                elif current.type == "table_close":
                    depth -= 1
                cursor += 1

            table_tokens = tokens[idx:cursor]
            headers: List[str] = []
            rows: List[List[str]] = []
            current_row: List[str] = []
            in_header = False

            inner = 0
            while inner < len(table_tokens):
                current = table_tokens[inner]
                if current.type == "thead_open":
                    in_header = True
                elif current.type == "thead_close":
                    in_header = False
                elif current.type == "tr_open":
                    current_row = []
                elif current.type in {"th_open", "td_open"}:
                    close_type = "th_close" if current.type == "th_open" else "td_close"
                    content_parts: List[str] = []
                    cell_cursor = inner + 1
                    while cell_cursor < len(table_tokens) and table_tokens[cell_cursor].type != close_type:
                        cell_token = table_tokens[cell_cursor]
                        if cell_token.type == "inline":
                            content_parts.append(cell_token.content.strip())
                        cell_cursor += 1
                    current_row.append(" ".join([part for part in content_parts if part]).strip())
                    inner = cell_cursor
                elif current.type == "tr_close":
                    if current_row:
                        if in_header and not headers:
                            headers = list(current_row)
                        else:
                            rows.append(list(current_row))

                inner += 1

            table_start = token.map[0] + 1 if token.map else None
            table_end = token.map[1] if token.map else None
            column_count = len(headers) if headers else (max((len(row) for row in rows), default=0))

            tables.append(
                {
                    "headers": headers,
                    "rows": rows,
                    "row_count": len(rows),
                    "column_count": column_count,
                    "normalized_text": self._normalize_table_text(headers, rows),
                    "start_line": table_start,
                    "end_line": table_end,
                }
            )
            idx = cursor

        return tables

    def _normalize_table_text(self, headers: List[str], rows: List[List[str]]) -> str:
        lines: List[str] = []
        if headers:
            lines.append("Columns: " + " | ".join(headers))
        for row_idx, row in enumerate(rows, start=1):
            lines.append(f"Row {row_idx}: " + " | ".join(row))
        return "\n".join(lines)

    def _classify_link_type(self, target: str) -> str:
        value = (target or "").strip().lower()
        if value.startswith("#"):
            return "anchor"
        if value.startswith("http://") or value.startswith("https://") or value.startswith("mailto:"):
            return "external"
        return "internal"

    def _associate_content_types(self, tokens: List[Any], metadata: List[Dict[str, Any]]) -> None:
        if not metadata:
            return

        token_type_map = {
            "paragraph_open": "paragraph",
            "bullet_list_open": "list",
            "ordered_list_open": "list",
            "blockquote_open": "blockquote",
            "table_open": "table",
            "fence": "code_block",
            "code_block": "code_block",
        }

        for token in tokens:
            content_type = token_type_map.get(token.type)
            if not content_type or not token.map:
                continue

            token_line = token.map[0] + 1
            for section in reversed(metadata):
                start_line = section.get("content_start_line")
                end_line = section.get("content_end_line")
                if start_line is None or end_line is None:
                    continue
                if start_line <= token_line <= end_line:
                    if content_type not in section["content_types"]:
                        section["content_types"].append(content_type)
                    break

    def _build_section_segments(
        self,
        sections: List[ParsedSection],
        metadata: List[Dict[str, Any]],
    ) -> List[Dict[str, Any]]:
        segments: List[Dict[str, Any]] = []
        for idx, section in enumerate(sections):
            if idx >= len(metadata):
                break

            title, level, body = section
            segment = {
                "section_title": title,
                "section_level": level,
                "parent_section": metadata[idx].get("parent_section"),
                "heading_path": metadata[idx].get("heading_path"),
                "start_line": metadata[idx].get("start_line"),
                "end_line": metadata[idx].get("end_line"),
                "content_start_line": metadata[idx].get("content_start_line"),
                "content_end_line": metadata[idx].get("content_end_line"),
                "content_types": metadata[idx].get("content_types", []),
                "content": body,
            }
            segments.append(segment)

        return segments

    def _trimmed_content_end_line(self, lines: List[str], start_line: int, end_line: int) -> int:
        if end_line < start_line:
            return start_line

        for line_no in range(end_line, start_line - 1, -1):
            if 1 <= line_no <= len(lines) and lines[line_no - 1].strip():
                return line_no
        return start_line

    def _is_plaintext_heading(self, lines: List[str], idx: int) -> bool:
        line = lines[idx].strip()
        if not line:
            return False
        if len(line) > 80 or "\t" in lines[idx] or "|" in line or "`" in line:
            return False
        if line.startswith(("-", "*", "+", ">")):
            return False
        if re.match(r"^\d+\.\s+", line):
            return False
        if line.endswith((".", ":", ";")):
            return False

        words = line.split()
        if not (1 <= len(words) <= 6):
            return False
        if not any(any(ch.isalpha() for ch in word) for word in words):
            return False

        next_non_empty: Optional[str] = None
        next_non_empty_raw: Optional[str] = None
        for next_line in lines[idx + 1:]:
            stripped = next_line.strip()
            if stripped:
                next_non_empty = stripped
                next_non_empty_raw = next_line
                break
        if not next_non_empty:
            return False

        if self._http_method_re.match(next_non_empty):
            return False

        if len(words) >= 3 and next_non_empty_raw and "\t" in next_non_empty_raw:
            return False

        if not line[0].isalpha() or not line[0].isupper():
            return False

        return True
