from __future__ import annotations

import re
from typing import Any, Literal

from pydantic import BaseModel, Field

from app.services.notifications import resolve_client_id
from app.store import store

SearchCategory = Literal["clients", "documents", "suppliers", "questions", "activity"]


class SearchHit(BaseModel):
    id: str
    category: SearchCategory
    title: str
    subtitle: str | None = None
    meta: str | None = None
    status: str | None = None
    reference_type: str
    reference_id: str
    score: float = 0


class SearchResponse(BaseModel):
    query: str
    total: int
    items: list[SearchHit]
    categories: dict[str, list[SearchHit]] = Field(default_factory=dict)


def _norm(value: Any) -> str:
    return str(value or "").strip().lower()


def _score_text(query: str, *candidates: str) -> float:
    """Higher is better. Exact > prefix > contains."""
    q = _norm(query)
    if not q:
        return 0
    best = 0.0
    for raw in candidates:
        text = _norm(raw)
        if not text:
            continue
        if text == q:
            best = max(best, 100.0)
        elif text.startswith(q):
            best = max(best, 80.0)
        elif q in text:
            # shorter distance from start = slightly better
            idx = text.find(q)
            best = max(best, 55.0 - min(idx, 20) * 0.5)
        else:
            # token prefix match
            for token in re.split(r"[\s_\-./]+", text):
                if token.startswith(q) and len(q) >= 2:
                    best = max(best, 45.0)
    return best


def _amount_match(query: str, total: float) -> float:
    digits = re.sub(r"[^\d]", "", query)
    if len(digits) < 3:
        return 0
    amount_digits = re.sub(r"[^\d]", "", f"{total:.2f}")
    if digits == amount_digits or digits in amount_digits or amount_digits.startswith(digits):
        return 70.0
    # also match integer euros e.g. 1245 vs 1245.80
    int_part = re.sub(r"[^\d]", "", f"{int(total)}")
    if digits == int_part or int_part.startswith(digits) or digits in int_part:
        return 65.0
    return 0


def _client_facing_status(status: str) -> str:
    mapping = {
        "New": "Processing",
        "Processing": "Processing",
        "AI processed": "Processing",
        "Needs review": "Under review",
        "Auto-ready": "Under review",
        "Approved": "Complete",
        "Synced": "Complete",
        "Rejected": "Action needed",
        "Sync failed": "Action needed",
    }
    return mapping.get(status, status)


def _eur(total: float) -> str:
    return f"€{total:,.2f}"


def search(
    *,
    query: str,
    platform: Literal["firm", "client"],
    email: str,
    client_id: str | None = None,
    category: SearchCategory | Literal["all"] | None = "all",
    limit: int = 40,
    per_category: int = 5,
) -> SearchResponse:
    q = query.strip()
    if len(q) < 1:
        return SearchResponse(query=q, total=0, items=[], categories={})

    resolved_client = resolve_client_id(platform, email, client_id)
    state = store.snapshot()
    want = category or "all"

    hits: list[SearchHit] = []

    # --- Clients (firm only) ---
    if platform == "firm" and want in ("all", "clients"):
        for client in state.clients:
            score = _score_text(q, client.name, client.id, client.kvk)
            if score <= 0:
                continue
            hits.append(
                SearchHit(
                    id=f"client:{client.id}",
                    category="clients",
                    title=client.name,
                    subtitle="Client",
                    meta=f"KvK {client.kvk}" if client.kvk else None,
                    status=None,
                    reference_type="client",
                    reference_id=client.id,
                    score=score + 5,  # slight boost for entity type
                )
            )

    # --- Documents / invoices ---
    if want in ("all", "documents"):
        for inv in state.invoices:
            if platform == "client" and inv.clientId != resolved_client:
                continue
            client_name = next((c.name for c in state.clients if c.id == inv.clientId), inv.clientId)
            score = max(
                _score_text(
                    q,
                    inv.file,
                    inv.supplier,
                    inv.number,
                    inv.paymentRef,
                    inv.suggestion.relation if inv.suggestion else "",
                    *( [client_name] if platform == "firm" else [] ),
                ),
                _amount_match(q, inv.total),
            )
            if score <= 0:
                continue
            status_label = inv.status if platform == "firm" else _client_facing_status(inv.status)
            subtitle_bits = [
                "Invoice",
                _eur(inv.total),
            ]
            if platform == "firm":
                subtitle_bits.insert(1, client_name)
            hits.append(
                SearchHit(
                    id=f"invoice:{inv.id}",
                    category="documents",
                    title=inv.supplier or inv.file,
                    subtitle=" · ".join(subtitle_bits),
                    meta=inv.file,
                    status=status_label,
                    reference_type="invoice",
                    reference_id=inv.id,
                    score=score,
                )
            )

    # --- Suppliers (firm only — unique supplier names from invoices) ---
    if platform == "firm" and want in ("all", "suppliers"):
        seen: set[str] = set()
        for inv in state.invoices:
            key = _norm(inv.supplier)
            if not key or key in seen:
                continue
            score = _score_text(q, inv.supplier)
            if score <= 0:
                continue
            seen.add(key)
            hits.append(
                SearchHit(
                    id=f"supplier:{key}",
                    category="suppliers",
                    title=inv.supplier,
                    subtitle="Supplier",
                    meta=None,
                    status=None,
                    reference_type="supplier",
                    reference_id=inv.supplier,
                    score=score,
                )
            )

    # --- Questions ---
    if want in ("all", "questions"):
        for query_row in state.queries:
            if platform == "client" and query_row.clientId != resolved_client:
                continue
            inv = next((i for i in state.invoices if i.id == query_row.invoiceId), None)
            supplier = inv.supplier if inv else None
            client_name = next(
                (c.name for c in state.clients if c.id == query_row.clientId),
                query_row.clientId,
            )
            score = _score_text(
                q,
                query_row.question,
                query_row.title or "",
                supplier or "",
                client_name,
                query_row.id,
                *(m.message for m in (query_row.messages or [])),
            )
            if inv and inv.number:
                score = max(score, _score_text(q, inv.number))
            if inv and inv.file:
                score = max(score, _score_text(q, inv.file))
            if score <= 0:
                continue
            status_label = {
                "Requested": "Requested",
                "Delivered": "Waiting for client" if platform == "firm" else "Needs your answer",
                "Answered": "Answered",
                "Resolved": "Resolved",
                "Cancelled": "Cancelled",
                "Approved": "Resolved",
            }.get(query_row.status, query_row.status)
            hits.append(
                SearchHit(
                    id=f"query:{query_row.id}",
                    category="questions",
                    title=(query_row.question[:80] + ("…" if len(query_row.question) > 80 else "")),
                    subtitle=supplier or client_name,
                    meta=None,
                    status=status_label,
                    reference_type="query",
                    reference_id=query_row.id,
                    score=score,
                )
            )

    # --- Activity ---
    if want in ("all", "activity"):
        for event in state.activity:
            if platform == "client":
                if not event.clientId or event.clientId != resolved_client:
                    continue
                # Hide internal-only activity types from clients
                if event.type in ("edit", "match", "ai", "sync", "sync-failed"):
                    continue
            score = _score_text(q, event.text, event.actor, event.type)
            if score <= 0:
                continue
            hits.append(
                SearchHit(
                    id=f"activity:{event.id}",
                    category="activity",
                    title=event.text,
                    subtitle=event.actor,
                    meta=None,
                    status=None,
                    reference_type="activity",
                    reference_id=event.invoiceId or event.id,
                    score=score * 0.85,  # slightly lower rank
                )
            )

    hits.sort(key=lambda h: (-h.score, h.title.lower()))

    # Group + cap per category for response shape
    categories: dict[str, list[SearchHit]] = {}
    for hit in hits:
        bucket = categories.setdefault(hit.category, [])
        if len(bucket) < per_category or want != "all":
            # for filtered category allow more
            max_n = limit if want != "all" else per_category
            if len(bucket) < max_n:
                bucket.append(hit)

    if want == "all":
        # flatten capped groups preserving category order
        order: list[SearchCategory] = ["clients", "documents", "suppliers", "questions", "activity"]
        items: list[SearchHit] = []
        for key in order:
            items.extend(categories.get(key, []))
        items = items[:limit]
    else:
        items = categories.get(want, [])[:limit]

    return SearchResponse(
        query=q,
        total=len(hits),
        items=items,
        categories={k: v for k, v in categories.items() if v},
    )


def autocomplete(
    *,
    query: str,
    platform: Literal["firm", "client"],
    email: str,
    client_id: str | None = None,
    limit: int = 12,
) -> SearchResponse:
    """Compact dropdown results: few per category."""
    return search(
        query=query,
        platform=platform,
        email=email,
        client_id=client_id,
        category="all",
        limit=limit,
        per_category=4,
    )
