from __future__ import annotations

from functools import lru_cache
from typing import Any

from app.core.config import settings
from app.generation.llm_client import LLMClientError
from app.llm import LLMRouter
from app.models.intent import INTENT_CONFIGS, IntentResult, QueryIntent
from app.models.internal_retrieve import SourceType


class Tier2Classifier:
    """LLM-based intent refinement for low-confidence Tier-1 results."""

    SYSTEM_PROMPT = """You classify developer questions about a codebase.
Return JSON with keys:
- intent: one of code_lookup, documentation, historical, architecture, general, mixed
- confidence: float 0-1
- source_types: array subset of code, docs, commit
- rewritten_query: search-optimized query string
- needs_graph_expansion: boolean
"""

    def __init__(
        self,
        llm_client: LLMRouter | None = None,
        *,
        router: LLMRouter | None = None,
    ):
        self._llm = llm_client or router or LLMRouter()

    async def classify(self, question: str) -> IntentResult | None:
        messages = [
            {"role": "system", "content": self.SYSTEM_PROMPT},
            {"role": "user", "content": question},
        ]
        try:
            data = await self._llm.complete_json(
                messages,
                model=settings.LLM_INTENT_MODEL,
            )
        except LLMClientError:
            return None
        return self._parse(data, question)

    def _parse(self, data: dict[str, Any], question: str) -> IntentResult | None:
        try:
            intent_raw = str(data.get("intent", "")).strip()
            intent = QueryIntent(intent_raw)
            confidence = float(data.get("confidence", 0.0))
            confidence = max(0.0, min(1.0, confidence))
            source_types = self._parse_source_types(data.get("source_types", []))
            rewritten = data.get("rewritten_query")
            rewritten_query = str(rewritten).strip() if rewritten else None
            needs_graph = data.get("needs_graph_expansion")
            needs_graph_expansion = bool(needs_graph) if needs_graph is not None else None
            if not source_types:
                base = INTENT_CONFIGS[intent]
                source_types = list(base.source_types or [])
            return IntentResult(
                type=intent,
                confidence=confidence,
                source_types_used=source_types,
                rewritten_query=rewritten_query or question,
                tier="tier2",
                needs_graph_expansion=needs_graph_expansion,
            )
        except (ValueError, TypeError):
            return None

    def _parse_source_types(self, raw: Any) -> list[SourceType]:
        if not isinstance(raw, list):
            return []
        result: list[SourceType] = []
        for item in raw:
            try:
                result.append(SourceType(str(item)))
            except ValueError:
                continue
        return result


@lru_cache(maxsize=256)
def _cache_key(question: str) -> str:
    return question.strip().lower()


_tier2_cache: dict[str, IntentResult] = {}


def get_cached_tier2(question: str) -> IntentResult | None:
    return _tier2_cache.get(_cache_key(question))


def set_cached_tier2(question: str, result: IntentResult) -> None:
    if len(_tier2_cache) >= 256:
        _tier2_cache.pop(next(iter(_tier2_cache)))
    _tier2_cache[_cache_key(question)] = result


def clear_tier2_cache() -> None:
    _tier2_cache.clear()
