"""Multi-document compare, reconcile, timeline, duplicates."""

from __future__ import annotations

import re
from dataclasses import dataclass
from datetime import datetime
from typing import Any

from app.database import get_connection
from app.fdi.compute import ComputedResult, _docs_for_agent


@dataclass
class MatchPair:
    left_id: str
    right_id: str
    score: float
    reason: str


def _parse_loose_date(value: str | None) -> str | None:
    if not value:
        return None
    v = value.strip()
    for fmt in ("%d %b %y", "%d %b %Y", "%d-%m-%Y", "%d/%m/%Y", "%Y-%m-%d", "%d %B %Y"):
        try:
            return datetime.strptime(v.replace("'", ""), fmt).date().isoformat()
        except ValueError:
            continue
    m = re.search(r"(\d{1,2})\s+([A-Za-z]{3,9})\s*'?(\d{2,4})", v)
    if m:
        try:
            raw = f"{m.group(1)} {m.group(2)} {m.group(3)}"
            for fmt in ("%d %b %y", "%d %b %Y", "%d %B %y", "%d %B %Y"):
                try:
                    return datetime.strptime(raw, fmt).date().isoformat()
                except ValueError:
                    continue
        except Exception:
            return None
    return None


def compare_documents(*, org_id: str, agent_id: str) -> ComputedResult:
    with get_connection() as conn:
        rows = conn.execute(
            """
            SELECT d.id, d.filename, d.doc_type, d.holder_name, d.period_start, d.period_end,
                   d.parse_confidence,
                   (SELECT value_numeric FROM fdi_assertions a
                    WHERE a.document_id=d.id AND a.kind IN ('total_spend','invoice_total','net_pay','portfolio_value')
                    ORDER BY a.confidence DESC LIMIT 1) AS primary_metric,
                   (SELECT kind FROM fdi_assertions a
                    WHERE a.document_id=d.id AND a.kind IN ('total_spend','invoice_total','net_pay','portfolio_value')
                    ORDER BY a.confidence DESC LIMIT 1) AS primary_kind,
                   (SELECT COUNT(*) FROM fdi_line_items l WHERE l.document_id=d.id) AS line_count
            FROM fdi_documents d
            WHERE d.org_id=? AND d.agent_id=? AND d.status='ready'
            ORDER BY d.indexed_at DESC
            LIMIT 20
            """,
            (org_id, agent_id),
        ).fetchall()
        items = [dict(r) for r in rows]
        if len(items) < 2:
            return ComputedResult(
                False,
                "compare",
                message="Need at least two indexed documents to compare.",
                confidence=0.8,
            )
        return ComputedResult(
            True,
            "compare",
            row_count=len(items),
            group_rows=items,
            evidence_ids=[str(x["id"]) for x in items],
            sql_fingerprint="compare_docs",
            message=f"Comparison across {len(items)} documents",
            confidence=0.85,
        )


def reconcile_amounts(
    *,
    org_id: str,
    agent_id: str,
    amount_tolerance: float = 1.0,
) -> ComputedResult:
    """Match invoice/receipt totals to bank/UPI outflows by amount (+/- tolerance)."""
    with get_connection() as conn:
        doc_ids = _docs_for_agent(conn, org_id, agent_id)
        if not doc_ids:
            return ComputedResult(False, "reconcile", message="No indexed documents.")
        ph = ",".join("?" for _ in doc_ids)
        invoices = conn.execute(
            f"""
            SELECT a.id, a.value_numeric, a.label, d.filename, d.doc_type, d.id AS document_id
            FROM fdi_assertions a
            JOIN fdi_documents d ON d.id=a.document_id
            WHERE a.org_id=? AND a.document_id IN ({ph})
              AND a.kind IN ('invoice_total','receipt_total','amount_due','po_total')
              AND a.value_numeric IS NOT NULL
            """,
            (org_id, *doc_ids),
        ).fetchall()
        payments = conn.execute(
            f"""
            SELECT l.id, COALESCE(l.debit, l.amount) AS amt, l.event_date, l.description_norm,
                   d.filename, d.doc_type
            FROM fdi_line_items l
            JOIN fdi_documents d ON d.id=l.document_id
            WHERE l.org_id=? AND l.document_id IN ({ph})
              AND l.direction='out'
              AND COALESCE(l.debit, l.amount, 0) > 0
              AND COALESCE(l.debit, l.amount, 0) < 5000000
              AND d.doc_type IN ('bank_statement','upi_statement','credit_card_statement')
            """,
            (org_id, *doc_ids),
        ).fetchall()
        inv = [dict(r) for r in invoices]
        pay = [dict(r) for r in payments]
        if not inv:
            return ComputedResult(
                False,
                "reconcile",
                message="No invoice/receipt totals found to reconcile. Upload invoices first.",
                confidence=0.8,
            )
        if not pay:
            return ComputedResult(
                False,
                "reconcile",
                message="No bank/UPI outflows found to match against invoices.",
                confidence=0.8,
            )

        used_pay: set[str] = set()
        matched: list[dict[str, Any]] = []
        unmatched_inv: list[dict[str, Any]] = []
        for invoice in inv:
            target = float(invoice["value_numeric"])
            best = None
            best_delta = None
            for p in pay:
                pid = str(p["id"])
                if pid in used_pay:
                    continue
                amt = float(p["amt"] or 0)
                delta = abs(amt - target)
                if delta <= amount_tolerance and (best_delta is None or delta < best_delta):
                    best = p
                    best_delta = delta
            if best:
                used_pay.add(str(best["id"]))
                matched.append(
                    {
                        "invoice": invoice.get("filename"),
                        "invoice_amount": target,
                        "payment": best.get("description_norm"),
                        "payment_amount": float(best["amt"]),
                        "payment_date": best.get("event_date"),
                        "delta": best_delta,
                        "status": "matched",
                    }
                )
            else:
                unmatched_inv.append(
                    {
                        "invoice": invoice.get("filename"),
                        "invoice_amount": target,
                        "status": "unmatched",
                    }
                )

        rows = matched + unmatched_inv
        return ComputedResult(
            True,
            "reconcile",
            row_count=len(rows),
            group_rows=rows,
            sql_fingerprint="reconcile_amounts",
            message=(
                f"Reconciled {len(matched)}/{len(inv)} invoice totals to bank/UPI outflows "
                f"({len(unmatched_inv)} unmatched)."
            ),
            confidence=0.8,
        )


def timeline_activity(
    *,
    org_id: str,
    agent_id: str,
    limit: int = 30,
) -> ComputedResult:
    with get_connection() as conn:
        doc_ids = _docs_for_agent(conn, org_id, agent_id)
        if not doc_ids:
            return ComputedResult(False, "timeline", message="No indexed documents.")
        ph = ",".join("?" for _ in doc_ids)
        rows = conn.execute(
            f"""
            SELECT l.id, l.event_date, l.description_norm, l.direction,
                   COALESCE(l.debit, l.credit, l.amount) AS amt, d.filename, d.doc_type
            FROM fdi_line_items l
            JOIN fdi_documents d ON d.id=l.document_id
            WHERE l.org_id=? AND l.document_id IN ({ph})
              AND l.event_date IS NOT NULL
              AND COALESCE(l.debit, l.credit, l.amount, 0) > 0
              AND COALESCE(l.debit, l.credit, l.amount, 0) < 5000000
            ORDER BY l.event_date DESC
            LIMIT ?
            """,
            (org_id, *doc_ids, limit),
        ).fetchall()
        items = [dict(r) for r in rows]
        # Attach ISO sort key when possible
        for it in items:
            it["iso_date"] = _parse_loose_date(str(it.get("event_date") or ""))
        items.sort(key=lambda x: x.get("iso_date") or "", reverse=True)
        return ComputedResult(
            True,
            "timeline",
            row_count=len(items),
            group_rows=items,
            evidence_ids=[str(x["id"]) for x in items],
            sql_fingerprint="timeline",
            message=f"Timeline of {len(items)} dated line items",
            confidence=0.8 if items else 0.4,
        )


def find_duplicates(
    *,
    org_id: str,
    agent_id: str,
    amount_tolerance: float = 0.01,
) -> ComputedResult:
    with get_connection() as conn:
        doc_ids = _docs_for_agent(conn, org_id, agent_id)
        if not doc_ids:
            return ComputedResult(False, "duplicates", message="No indexed documents.")
        ph = ",".join("?" for _ in doc_ids)
        rows = conn.execute(
            f"""
            SELECT l.event_date,
                   COALESCE(l.debit, l.credit, l.amount) AS amt,
                   l.description_norm
            FROM fdi_line_items l
            WHERE l.org_id=? AND l.document_id IN ({ph})
              AND COALESCE(l.debit, l.credit, l.amount, 0) > 0
              AND COALESCE(l.debit, l.credit, l.amount, 0) < 5000000
            """,
            (org_id, *doc_ids),
        ).fetchall()
        buckets: dict[tuple[str, str], dict[str, Any]] = {}
        for r in rows:
            d = dict(r)
            amt = round(float(d.get("amt") or 0), 2)
            key = (str(d.get("event_date") or ""), f"{amt:.2f}")
            slot = buckets.setdefault(
                key,
                {"event_date": d.get("event_date"), "amt": amt, "n": 0, "sample_desc": d.get("description_norm")},
            )
            slot["n"] += 1
        items = [v for v in buckets.values() if v["n"] > 1]
        items.sort(key=lambda x: (-int(x["n"]), -float(x["amt"])))
        return ComputedResult(
            True,
            "duplicates",
            row_count=len(items),
            group_rows=items[:30],
            sql_fingerprint="duplicates",
            message=f"Found {len(items)} possible duplicate amount/date groups",
            confidence=0.75 if items else 0.7,
        )
