from __future__ import annotations

import uuid
from typing import Any

import httpx

from app.core.config import settings
from app.models.intent import IntentResult, QueryIntent
from app.models.search import SearchOptions
from app.models.internal_retrieve import (
    RetrieveFilters,
    RetrieveMetadata,
    RetrieveRequest,
    RetrieveResponse,
)
from app.retrieval.fixtures.retrieve_fixtures import build_mock_retrieve_response
from app.retrieval.query_rewriter import apply_query_rewrite


class RetrievalClientError(Exception):
    """Base error for retrieval client failures."""

    def __init__(self, message: str, status_code: int | None = None, request_id: str | None = None):
        super().__init__(message)
        self.status_code = status_code
        self.request_id = request_id


class RetrievalClient:
    """HTTP client for embedding-engine POST /internal/retrieve."""

    def __init__(
        self,
        base_url: str | None = None,
        timeout_sec: float | None = None,
        mock: bool | None = None,
        client: httpx.AsyncClient | None = None,
    ):
        self._base_url = (base_url or settings.EMBEDDING_RETRIEVAL_SERVICE_URL).rstrip("/")
        self._timeout = timeout_sec or settings.EMBEDDING_RETRIEVAL_TIMEOUT_SEC
        self._mock_override = mock
        self._client = client

    def _use_mock(self) -> bool:
        if self._mock_override is not None:
            return self._mock_override
        return settings.RETRIEVAL_CLIENT_MOCK

    async def retrieve(self, request: RetrieveRequest) -> RetrieveResponse:
        if self._use_mock():
            return self._mock_response(request)

        url = f"{self._base_url}/internal/retrieve"
        payload = request.model_dump(mode="json", exclude_none=True)

        if self._client is not None:
            response = await self._client.post(url, json=payload)
        else:
            async with httpx.AsyncClient(timeout=self._timeout) as client:
                response = await client.post(url, json=payload)

        if response.status_code >= 400:
            request_id = request.request_id
            detail = _extract_error_detail(response)
            raise RetrievalClientError(
                f"Retrieval failed ({response.status_code}): {detail}",
                status_code=response.status_code,
                request_id=request_id,
            )

        return RetrieveResponse.model_validate(response.json())

    def _mock_response(self, request: RetrieveRequest) -> RetrieveResponse:
        return build_mock_retrieve_response(request)


def resolve_retrieval_score_threshold(
    intent: IntentResult,
    options: SearchOptions | None = None,
) -> float:
    """Pick vector score cutoff: explicit UI override, then intent default, then global."""
    if options and options.score_threshold is not None:
        return options.score_threshold
    if intent.type == QueryIntent.HISTORICAL:
        return settings.QUERY_HISTORICAL_SCORE_THRESHOLD
    if intent.type == QueryIntent.ARCHITECTURE:
        return settings.QUERY_ARCHITECTURE_SCORE_THRESHOLD
    if intent.type == QueryIntent.MIXED:
        return settings.QUERY_MIXED_SCORE_THRESHOLD
    return settings.QUERY_DEFAULT_SCORE_THRESHOLD


def build_retrieve_request(
    *,
    repo_id: str,
    snapshot_id: str,
    question: str,
    intent: IntentResult,
    options: SearchOptions | None = None,
    request_id: str | None = None,
    source_types_override: list | None = None,
    top_k_override: int | None = None,
    score_threshold_override: float | None = None,
    graph_expand_override: bool | None = None,
) -> RetrieveRequest:
    """Map intent + public search options to internal retrieve request (contracts §7.1)."""
    config = intent.config
    source_types = source_types_override if source_types_override is not None else config.source_types
    top_k = top_k_override if top_k_override is not None else (
        options.top_k if options and options.top_k is not None else config.top_k
    )
    score_threshold = (
        score_threshold_override
        if score_threshold_override is not None
        else resolve_retrieval_score_threshold(intent, options)
    )
    graph_expand = (
        graph_expand_override
        if graph_expand_override is not None
        else config.needs_graph_expansion
    )

    filters = RetrieveFilters(source_types=source_types) if source_types else None

    return RetrieveRequest(
        repo_id=repo_id,
        snapshot_id=snapshot_id,
        query_text=apply_query_rewrite(intent, question),
        filters=filters,
        top_k=top_k,
        score_threshold=score_threshold,
        graph_expand=graph_expand,
        request_id=request_id or f"req_{uuid.uuid4().hex[:12]}",
    )


def _extract_error_detail(response: httpx.Response) -> str:
    try:
        body: dict[str, Any] = response.json()
        if "detail" in body:
            return str(body["detail"])
        return str(body)
    except Exception:
        return response.text or "unknown error"
