from __future__ import annotations

from dataclasses import dataclass

from app.core.config import settings
from app.models.context import BuiltContext, Citation, ContextChunk
from app.models.intent import IntentResult, QueryIntent
from app.models.internal_retrieve import (
    CodeSnippet,
    DocExcerpt,
    RelatedCommit,
    RetrieveResponse,
    SourceType,
)
from app.utils.tokenizer import get_tokenizer


@dataclass(frozen=True)
class _ContextItem:
    record_id: str
    source_type: SourceType
    score: float
    text: str
    label: str
    file_path: str | None = None
    doc_path: str | None = None
    commit_sha: str | None = None
    symbol_name: str | None = None
    section_title: str | None = None


class ContextBuilder:
    """Assemble, dedupe, rank, and trim hydrated retrieval context for LLM prompts."""

    def __init__(self, token_budget: int | None = None):
        self._token_budget = token_budget or settings.LLM_CONTEXT_TOKEN_BUDGET
        self._tokenizer = get_tokenizer()

    def build(
        self,
        retrieve_response: RetrieveResponse,
        intent: IntentResult,
        *,
        token_budget: int | None = None,
    ) -> BuiltContext:
        budget = token_budget or self._token_budget
        items = self._flatten(retrieve_response)
        items = self._dedupe(items)
        items = self._rank(items, intent)
        return self._pack(items, budget)

    def map_citations(
        self,
        used_indices: list[int],
        built: BuiltContext,
    ) -> list[Citation]:
        by_index = {chunk.index: chunk for chunk in built.chunks}
        citations: list[Citation] = []
        for index in used_indices:
            chunk = by_index.get(index)
            if not chunk:
                continue
            citations.append(
                Citation(
                    index=chunk.index,
                    record_id=chunk.record_id,
                    source_type=chunk.source_type,
                    score=chunk.score,
                    file_path=chunk.file_path,
                    doc_path=chunk.doc_path,
                    commit_sha=chunk.commit_sha,
                    symbol_name=chunk.symbol_name,
                    section_title=chunk.section_title,
                )
            )
        return citations

    def _flatten(self, response: RetrieveResponse) -> list[_ContextItem]:
        items: list[_ContextItem] = []
        for snippet in response.code_snippets:
            items.append(self._from_code(snippet))
        for excerpt in response.doc_excerpts:
            items.append(self._from_doc(excerpt))
        for commit in response.related_commits:
            items.append(self._from_commit(commit))
        return items

    def _from_code(self, snippet: CodeSnippet) -> _ContextItem:
        symbol = snippet.symbol_name or "code"
        return _ContextItem(
            record_id=snippet.record_id,
            source_type=SourceType.CODE,
            score=snippet.score,
            text=snippet.text,
            label=f"{snippet.file_path} ({symbol})",
            file_path=snippet.file_path,
            symbol_name=snippet.symbol_name,
        )

    def _from_doc(self, excerpt: DocExcerpt) -> _ContextItem:
        title = excerpt.section_title or excerpt.doc_path
        return _ContextItem(
            record_id=excerpt.record_id,
            source_type=SourceType.DOCS,
            score=excerpt.score,
            text=excerpt.text,
            label=f"{excerpt.doc_path} — {title}",
            doc_path=excerpt.doc_path,
            section_title=excerpt.section_title,
        )

    def _from_commit(self, commit: RelatedCommit) -> _ContextItem:
        parts: list[str] = []

        author = self._format_person(commit.author_name, commit.author_email)
        if author:
            when = (commit.authored_at or "").strip() or "unknown"
            parts.append(f"Author: {author} ({when})")

        committer = self._format_person(commit.committer_name, commit.committer_email)
        if committer and committer != author:
            when = (commit.committed_at or "").strip() or "unknown"
            parts.append(f"Committer: {committer} ({when})")
        elif commit.committed_at and commit.committed_at.strip():
            parts.append(f"Committed at: {commit.committed_at.strip()}")

        if commit.commit_message and commit.commit_message.strip():
            parts.append("Message:")
            parts.append(commit.commit_message.strip())

        if commit.summary.strip():
            parts.append(f"Summary: {commit.summary.strip()}")

        if commit.changed_files:
            parts.append(
                "Changed files: " + ", ".join(commit.changed_files),
            )
        if commit.impacted_symbols:
            preview = commit.impacted_symbols[:20]
            suffix = "..." if len(commit.impacted_symbols) > len(preview) else ""
            parts.append(
                "Impacted symbols: " + ", ".join(preview) + suffix,
            )
        return _ContextItem(
            record_id=commit.record_id,
            source_type=SourceType.COMMIT,
            score=commit.score,
            text="\n".join(parts),
            label=f"commit {commit.commit_sha}",
            commit_sha=commit.commit_sha,
        )

    @staticmethod
    def _format_person(name: str | None, email: str | None) -> str:
        name = (name or "").strip()
        email = (email or "").strip()
        if name and email:
            return f"{name} <{email}>"
        return name or email

    def _dedupe(self, items: list[_ContextItem]) -> list[_ContextItem]:
        best: dict[str, _ContextItem] = {}
        for item in items:
            existing = best.get(item.record_id)
            if existing is None or item.score > existing.score:
                best[item.record_id] = item
        return list(best.values())

    def _rank(self, items: list[_ContextItem], intent: IntentResult) -> list[_ContextItem]:
        preferred = set(intent.source_types_used)

        def sort_key(item: _ContextItem) -> tuple[float, float]:
            boost = 0.1 if item.source_type in preferred else 0.0
            if intent.type == QueryIntent.ARCHITECTURE and item.source_type == SourceType.DOCS:
                boost += 0.15
                label_lower = item.label.lower()
                if "repo_overview" in label_lower or "repository overview" in label_lower:
                    boost += 0.2
            return (item.score + boost, item.score)

        return sorted(items, key=sort_key, reverse=True)

    def _pack(self, items: list[_ContextItem], budget: int) -> BuiltContext:
        chunks: list[ContextChunk] = []
        citations: list[Citation] = []
        parts: list[str] = []
        used_tokens = 0
        trimmed_count = 0
        index = 1

        for item in items:
            header = f"[{index}] {item.label}\n"
            body = item.text.strip()
            block = f"{header}{body}\n"
            block_tokens = self._tokenizer.count_tokens(block)

            if used_tokens + block_tokens > budget:
                remaining = budget - used_tokens - self._tokenizer.count_tokens(header)
                if remaining <= 0:
                    trimmed_count += len(items) - len(chunks)
                    break
                body = self._tokenizer.truncate_text(body, remaining)
                block = f"{header}{body}\n"
                block_tokens = self._tokenizer.count_tokens(block)
                if used_tokens + block_tokens > budget:
                    trimmed_count += len(items) - len(chunks)
                    break

            chunk = ContextChunk(
                index=index,
                source_type=item.source_type,
                record_id=item.record_id,
                label=item.label,
                text=body,
                score=item.score,
                file_path=item.file_path,
                doc_path=item.doc_path,
                commit_sha=item.commit_sha,
                symbol_name=item.symbol_name,
                section_title=item.section_title,
            )
            chunks.append(chunk)
            citations.append(
                Citation(
                    index=index,
                    record_id=item.record_id,
                    source_type=item.source_type,
                    score=item.score,
                    file_path=item.file_path,
                    doc_path=item.doc_path,
                    commit_sha=item.commit_sha,
                    symbol_name=item.symbol_name,
                    section_title=item.section_title,
                )
            )
            parts.append(block)
            used_tokens += block_tokens
            index += 1

        if len(chunks) < len(items):
            trimmed_count = max(trimmed_count, len(items) - len(chunks))

        return BuiltContext(
            chunks=chunks,
            prompt_text="".join(parts),
            citations=citations,
            total_tokens=used_tokens,
            trimmed_count=trimmed_count,
        )
