from __future__ import annotations

import html
import json
import re
from typing import Any
from urllib.parse import parse_qs, unquote, urlparse

import httpx

from app.core.config import settings
from app.core.logger import logger

_DEFAULT_MAX_RESULTS = 5
_MAX_RESULTS_CAP = 10
_MAX_OUTPUT_BYTES = 24_384
_HTML_TAG_RE = re.compile(r"(?s)<[^>]*>")
_DDG_LINK_RE = re.compile(
    r'(?i)<a[^>]+rel="nofollow"[^>]+class="result__a"[^>]+href="([^"]+)"[^>]*>(.*?)</a>'
)
_DDG_SNIPPET_RE = re.compile(
    r'(?i)<a[^>]+class="result__snippet"[^>]*>(.*?)</a>|'
    r'<td[^>]+class="result-snippet"[^>]*>(.*?)</td>'
)


class WebSearchClient:
    """Cloud-side web search used by agent turns (not executed by the local daemon)."""

    def __init__(self, *, client: httpx.AsyncClient | None = None, timeout_sec: float | None = None):
        self._client = client
        self._timeout = timeout_sec if timeout_sec is not None else settings.WEB_SEARCH_TIMEOUT_SEC

    async def search(self, *, query: str, max_results: int = _DEFAULT_MAX_RESULTS) -> dict[str, Any]:
        cleaned = (query or "").strip()
        if not cleaned:
            return {
                "tool": "web_search",
                "success": False,
                "error": "query is required",
                "result_count": 0,
                "results": [],
            }

        limit = max(1, min(int(max_results or _DEFAULT_MAX_RESULTS), _MAX_RESULTS_CAP))
        try:
            hits, provider = await self._run_search(cleaned, limit)
        except Exception as exc:  # noqa: BLE001 — surface as tool failure to the model
            logger.warning("web_search failed query={query!r} error={error}", query=cleaned, error=str(exc))
            return {
                "tool": "web_search",
                "query": cleaned,
                "success": False,
                "error": str(exc),
                "result_count": 0,
                "results": [],
            }

        output = _format_output(cleaned, provider, hits)
        truncated = False
        if len(output) > _MAX_OUTPUT_BYTES:
            output = output[:_MAX_OUTPUT_BYTES] + "…"
            truncated = True

        observation: dict[str, Any] = {
            "tool": "web_search",
            "query": cleaned,
            "provider": provider,
            "success": True,
            "result_count": len(hits),
            "results": hits,
            "output": output,
        }
        if truncated:
            observation["truncated"] = True
        return observation

    async def _run_search(self, query: str, max_results: int) -> tuple[list[dict[str, str]], str]:
        api_key = (settings.BRAVE_SEARCH_API_KEY or "").strip()
        if api_key:
            try:
                hits = await self._brave_search(query, max_results, api_key)
                if hits:
                    return hits, "brave"
            except Exception as exc:  # noqa: BLE001
                logger.warning("brave web_search failed; falling back to duckduckgo error={error}", error=str(exc))

        hits = await self._duckduckgo_search(query, max_results)
        return hits, "duckduckgo"

    async def _brave_search(self, query: str, max_results: int, api_key: str) -> list[dict[str, str]]:
        params = {"q": query, "count": max_results}
        data = await self._get_json(
            "https://api.search.brave.com/res/v1/web/search",
            params=params,
            headers={
                "Accept": "application/json",
                "X-Subscription-Token": api_key,
                "User-Agent": "Lucos-RAG/1.0",
            },
        )
        results = (((data or {}).get("web") or {}).get("results")) or []
        hits: list[dict[str, str]] = []
        for item in results:
            title = str(item.get("title") or "").strip()
            url = str(item.get("url") or "").strip()
            snippet = str(item.get("description") or "").strip()
            if not title or not url:
                continue
            hits.append({"title": title, "url": url, "snippet": snippet})
            if len(hits) >= max_results:
                break
        return hits

    async def _duckduckgo_search(self, query: str, max_results: int) -> list[dict[str, str]]:
        hits = await self._duckduckgo_instant(query)
        if len(hits) >= max_results:
            return hits[:max_results]
        try:
            html_hits = await self._duckduckgo_html(query, max_results)
            hits = _append_unique(hits, html_hits, max_results)
        except Exception as exc:  # noqa: BLE001
            if hits:
                logger.warning(
                    "duckduckgo html enrichment failed; using instant answers error={error}",
                    error=str(exc),
                )
            else:
                raise
        return hits[:max_results]

    async def _duckduckgo_instant(self, query: str) -> list[dict[str, str]]:
        data = await self._get_json(
            "https://api.duckduckgo.com/",
            params={
                "q": query,
                "format": "json",
                "no_redirect": "1",
                "no_html": "1",
                "skip_disambig": "1",
            },
            headers={"Accept": "application/json", "User-Agent": "Lucos-RAG/1.0"},
        )
        hits: list[dict[str, str]] = []
        abstract = str(data.get("AbstractText") or "").strip()
        if not abstract:
            abstract = str(data.get("Answer") or "").strip()
        if not abstract:
            abstract = str(data.get("Definition") or "").strip()
        if abstract:
            title = str(data.get("Heading") or "").strip() or query
            url = str(data.get("AbstractURL") or "").strip() or f"https://duckduckgo.com/?q={query}"
            hits.append({"title": title, "url": url, "snippet": abstract})

        for topic in data.get("RelatedTopics") or []:
            if not isinstance(topic, dict):
                continue
            text = str(topic.get("Text") or "").strip()
            url = str(topic.get("FirstURL") or "").strip()
            if not text or not url:
                continue
            title = text
            if " - " in text:
                title = text.split(" - ", 1)[0].strip()
            hits.append({"title": title, "url": url, "snippet": text})
            if len(hits) >= 4:
                break
        return hits

    async def _duckduckgo_html(self, query: str, max_results: int) -> list[dict[str, str]]:
        text = await self._get_text(
            "https://html.duckduckgo.com/html/",
            params={"q": query},
            headers={"Accept": "text/html", "User-Agent": "Lucos-RAG/1.0"},
        )
        link_matches = _DDG_LINK_RE.findall(text)
        snippet_matches = _DDG_SNIPPET_RE.findall(text)
        hits: list[dict[str, str]] = []
        for index, (raw_url, raw_title) in enumerate(link_matches):
            url = _decode_ddg_redirect(html.unescape(raw_url.strip()))
            title = _clean_html_text(raw_title)
            if not url or not title:
                continue
            snippet = ""
            if index < len(snippet_matches):
                first, second = snippet_matches[index]
                snippet = _clean_html_text(first or second)
            hits.append({"title": title, "url": url, "snippet": snippet})
            if len(hits) >= max_results:
                break
        return hits

    async def _get_json(
        self,
        url: str,
        *,
        params: dict[str, Any] | None = None,
        headers: dict[str, str] | None = None,
    ) -> dict[str, Any]:
        response = await self._request(url, params=params, headers=headers)
        try:
            data = response.json()
        except json.JSONDecodeError as exc:
            raise RuntimeError(f"invalid JSON from {url}") from exc
        if not isinstance(data, dict):
            raise RuntimeError(f"unexpected JSON payload from {url}")
        return data

    async def _get_text(
        self,
        url: str,
        *,
        params: dict[str, Any] | None = None,
        headers: dict[str, str] | None = None,
    ) -> str:
        response = await self._request(url, params=params, headers=headers)
        return response.text

    async def _request(
        self,
        url: str,
        *,
        params: dict[str, Any] | None = None,
        headers: dict[str, str] | None = None,
    ) -> httpx.Response:
        if self._client is not None:
            response = await self._client.get(url, params=params, headers=headers)
        else:
            async with httpx.AsyncClient(timeout=self._timeout) as client:
                response = await client.get(url, params=params, headers=headers)
        if response.status_code < 200 or response.status_code >= 300:
            raise RuntimeError(f"web search HTTP {response.status_code} from {url}")
        return response


def _clean_html_text(value: str) -> str:
    text = _HTML_TAG_RE.sub("", value or "")
    text = html.unescape(text)
    return " ".join(text.split()).strip()


def _decode_ddg_redirect(raw: str) -> str:
    value = (raw or "").strip()
    if not value:
        return ""
    if value.startswith("//"):
        value = "https:" + value
    parsed = urlparse(value)
    if "duckduckgo.com" in (parsed.netloc or ""):
        uddg = parse_qs(parsed.query).get("uddg", [""])[0]
        if uddg:
            return unquote(uddg)
    return value


def _append_unique(
    existing: list[dict[str, str]],
    extra: list[dict[str, str]],
    max_results: int,
) -> list[dict[str, str]]:
    seen = {str(item.get("url") or "").lower() for item in existing}
    out = list(existing)
    for item in extra:
        key = str(item.get("url") or "").lower()
        if not key or key in seen:
            continue
        seen.add(key)
        out.append(item)
        if len(out) >= max_results:
            break
    return out


def _format_output(query: str, provider: str, hits: list[dict[str, str]]) -> str:
    lines = [f"Web search query: {query}", f"Provider: {provider}"]
    if not hits:
        lines.append("No results found.")
        return "\n".join(lines)
    for index, hit in enumerate(hits, start=1):
        lines.append("")
        lines.append(f"{index}. {hit.get('title') or ''}")
        lines.append(f"   {hit.get('url') or ''}")
        snippet = str(hit.get("snippet") or "").strip()
        if snippet:
            lines.append(f"   {snippet}")
    return "\n".join(lines)
