"""Deterministic computation engine — LLMs must not invent these numbers."""

from __future__ import annotations

import re
from dataclasses import dataclass, field
from typing import Any

from app.database import get_connection


@dataclass
class ComputedResult:
    ok: bool
    metric: str
    value: float | None = None
    unit: str | None = None
    currency: str | None = None
    row_count: int = 0
    group_rows: list[dict[str, Any]] = field(default_factory=list)
    evidence_ids: list[str] = field(default_factory=list)
    sql_fingerprint: str = ""
    message: str = ""
    confidence: float = 0.0


def _docs_for_agent(conn: Any, org_id: str, agent_id: str) -> list[str]:
    rows = conn.execute(
        """
        SELECT id FROM fdi_documents
        WHERE org_id = ? AND agent_id = ? AND status = 'ready'
        """,
        (org_id, agent_id),
    ).fetchall()
    return [str(r["id"] if isinstance(r, dict) else r[0]) for r in rows]


def lookup_identifier(
    *,
    org_id: str,
    agent_id: str,
    id_type: str,
    value_norm: str = "",
) -> ComputedResult:
    with get_connection() as conn:
        doc_ids = _docs_for_agent(conn, org_id, agent_id)
        if not doc_ids:
            return ComputedResult(False, "identifier", message="No indexed financial documents.")
        placeholders = ",".join("?" for _ in doc_ids)
        if value_norm:
            rows = conn.execute(
                f"""
                SELECT i.id, i.id_type, i.value_raw, i.value_norm, i.confidence, d.filename, d.doc_type
                FROM fdi_identifiers i
                JOIN fdi_documents d ON d.id = i.document_id
                WHERE i.org_id = ? AND i.id_type = ? AND i.value_norm = ?
                  AND i.document_id IN ({placeholders})
                ORDER BY i.confidence DESC
                LIMIT 20
                """,
                (org_id, id_type, value_norm, *doc_ids),
            ).fetchall()
        else:
            rows = conn.execute(
                f"""
                SELECT i.id, i.id_type, i.value_raw, i.value_norm, i.confidence, d.filename, d.doc_type
                FROM fdi_identifiers i
                JOIN fdi_documents d ON d.id = i.document_id
                WHERE i.org_id = ? AND i.id_type = ?
                  AND i.document_id IN ({placeholders})
                ORDER BY i.confidence DESC
                LIMIT 20
                """,
                (org_id, id_type, *doc_ids),
            ).fetchall()
        items = [dict(r) for r in rows]
        if not items:
            return ComputedResult(
                False,
                "identifier",
                message=f"No {id_type} found in the knowledge base.",
                confidence=0.9,
            )
        return ComputedResult(
            True,
            "identifier",
            value=None,
            row_count=len(items),
            group_rows=items,
            evidence_ids=[str(x["id"]) for x in items],
            sql_fingerprint=f"identifiers:{id_type}:{value_norm or '*'}",
            message=f"Found {len(items)} {id_type} value(s).",
            confidence=float(items[0].get("confidence") or 0.85),
        )


def get_assertion(
    *,
    org_id: str,
    agent_id: str,
    kind: str,
) -> ComputedResult:
    with get_connection() as conn:
        doc_ids = _docs_for_agent(conn, org_id, agent_id)
        if not doc_ids:
            return ComputedResult(False, kind, message="No indexed financial documents.")
        placeholders = ",".join("?" for _ in doc_ids)
        rows = conn.execute(
            f"""
            SELECT a.id, a.kind, a.label, a.value_numeric, a.currency, a.source, a.confidence,
                   d.filename, d.doc_type, d.holder_name
            FROM fdi_assertions a
            JOIN fdi_documents d ON d.id = a.document_id
            WHERE a.org_id = ? AND a.kind = ? AND a.document_id IN ({placeholders})
            ORDER BY a.confidence DESC, CASE a.source WHEN 'header' THEN 0 ELSE 1 END
            LIMIT 10
            """,
            (org_id, kind, *doc_ids),
        ).fetchall()
        items = [dict(r) for r in rows]
        if not items:
            return ComputedResult(False, kind, message=f"No '{kind}' metric on indexed documents.")
        best = items[0]
        return ComputedResult(
            True,
            kind,
            value=best.get("value_numeric"),
            currency=best.get("currency"),
            row_count=len(items),
            group_rows=items,
            evidence_ids=[str(x["id"]) for x in items],
            sql_fingerprint=f"assertion:{kind}",
            message=best.get("label") or kind,
            confidence=float(best.get("confidence") or 0.8),
        )


def sum_direction(
    *,
    org_id: str,
    agent_id: str,
    direction: str,
    counterparty_like: str | None = None,
) -> ComputedResult:
    with get_connection() as conn:
        doc_ids = _docs_for_agent(conn, org_id, agent_id)
        if not doc_ids:
            return ComputedResult(False, f"sum_{direction}", message="No indexed financial documents.")
        placeholders = ",".join("?" for _ in doc_ids)
        params: list[Any] = [org_id, direction, *doc_ids]
        where_extra = ""
        if counterparty_like:
            where_extra = " AND (l.description_norm LIKE ? OR e.canonical_name LIKE ?)"
            like = f"%{counterparty_like.lower()}%"
            params.extend([like, like])
        row = conn.execute(
            f"""
            SELECT COALESCE(SUM(
                     COALESCE(
                       CASE WHEN COALESCE(l.debit, 0) > 0 AND COALESCE(l.debit, 0) < 5000000
                            THEN l.debit END,
                       CASE WHEN COALESCE(l.credit, 0) > 0 AND COALESCE(l.credit, 0) < 5000000
                            THEN l.credit END,
                       CASE WHEN COALESCE(l.amount, 0) > 0 AND COALESCE(l.amount, 0) < 5000000
                            THEN l.amount END,
                       0
                     )
                   ), 0) AS total,
                   COUNT(*) AS n
            FROM fdi_line_items l
            LEFT JOIN fdi_entities e ON e.id = l.counterparty_entity_id
            WHERE l.org_id = ? AND l.direction = ? AND l.document_id IN ({placeholders})
            {where_extra}
            """,
            tuple(params),
        ).fetchone()
        total = float(row["total"] if isinstance(row, dict) else row[0])
        n = int(row["n"] if isinstance(row, dict) else row[1])
        return ComputedResult(
            True,
            f"sum_{direction}",
            value=round(total, 2),
            currency="INR",
            row_count=n,
            sql_fingerprint=f"sum:{direction}:{counterparty_like or '*'}",
            message=f"Sum of {direction} flows" + (f" for {counterparty_like}" if counterparty_like else ""),
            confidence=0.8 if n else 0.4,
        )


def group_by_counterparty(
    *,
    org_id: str,
    agent_id: str,
    direction: str = "out",
    limit: int = 10,
) -> ComputedResult:
    with get_connection() as conn:
        doc_ids = _docs_for_agent(conn, org_id, agent_id)
        if not doc_ids:
            return ComputedResult(False, "group_counterparty", message="No indexed financial documents.")
        placeholders = ",".join("?" for _ in doc_ids)
        rows = conn.execute(
            f"""
            SELECT COALESCE(e.canonical_name, l.description_norm, 'Unknown') AS name,
                   SUM(
                     COALESCE(
                       CASE WHEN COALESCE(l.debit, 0) > 0 AND COALESCE(l.debit, 0) < 5000000
                            THEN l.debit END,
                       CASE WHEN COALESCE(l.credit, 0) > 0 AND COALESCE(l.credit, 0) < 5000000
                            THEN l.credit END,
                       CASE WHEN COALESCE(l.amount, 0) > 0 AND COALESCE(l.amount, 0) < 5000000
                            THEN l.amount END,
                       0
                     )
                   ) AS total,
                   COUNT(*) AS n
            FROM fdi_line_items l
            LEFT JOIN fdi_entities e ON e.id = l.counterparty_entity_id
            WHERE l.org_id = ? AND l.direction = ? AND l.document_id IN ({placeholders})
              AND LOWER(COALESCE(e.canonical_name, l.description_norm, '')) NOT LIKE 'upi ref%'
            GROUP BY COALESCE(e.canonical_name, l.description_norm, 'Unknown')
            HAVING SUM(
                     COALESCE(
                       CASE WHEN COALESCE(l.debit, 0) > 0 AND COALESCE(l.debit, 0) < 5000000
                            THEN l.debit END,
                       CASE WHEN COALESCE(l.credit, 0) > 0 AND COALESCE(l.credit, 0) < 5000000
                            THEN l.credit END,
                       CASE WHEN COALESCE(l.amount, 0) > 0 AND COALESCE(l.amount, 0) < 5000000
                            THEN l.amount END,
                       0
                     )
                   ) > 0
            ORDER BY total DESC
            LIMIT ?
            """,
            (org_id, direction, *doc_ids, limit),
        ).fetchall()
        items = [dict(r) for r in rows]
        return ComputedResult(
            True,
            "group_counterparty",
            row_count=len(items),
            group_rows=items,
            sql_fingerprint=f"group:{direction}",
            message=f"Top counterparties by {direction}",
            confidence=0.8 if items else 0.4,
        )


def document_inventory(*, org_id: str, agent_id: str) -> ComputedResult:
    with get_connection() as conn:
        rows = conn.execute(
            """
            SELECT id, filename, doc_type, doc_type_confidence, page_count,
                   holder_name, period_start, period_end, parse_confidence, status
            FROM fdi_documents
            WHERE org_id = ? AND agent_id = ? AND status = 'ready'
            ORDER BY indexed_at DESC
            """,
            (org_id, agent_id),
        ).fetchall()
        items = [dict(r) for r in rows]
        return ComputedResult(
            True,
            "inventory",
            row_count=len(items),
            group_rows=items,
            sql_fingerprint="inventory",
            message=f"{len(items)} ready document(s)",
            confidence=1.0,
        )


def document_profile(*, org_id: str, agent_id: str) -> ComputedResult:
    with get_connection() as conn:
        rows = conn.execute(
            """
            SELECT id, filename, doc_type, holder_name, issuer_name,
                   period_start, period_end, parse_confidence, currency_primary
            FROM fdi_documents
            WHERE org_id = ? AND agent_id = ? AND status = 'ready'
            ORDER BY indexed_at DESC
            LIMIT 20
            """,
            (org_id, agent_id),
        ).fetchall()
        items = [dict(r) for r in rows]
        if not items:
            return ComputedResult(False, "profile", message="No indexed financial documents.")
        return ComputedResult(
            True,
            "profile",
            row_count=len(items),
            group_rows=items,
            evidence_ids=[str(x["id"]) for x in items],
            sql_fingerprint="profile",
            message="Document profile",
            confidence=float(items[0].get("parse_confidence") or 0.7),
        )


def detect_merchant_in_question(question: str) -> str | None:
    q = (question or "").strip()
    m = re.search(
        r"(?i)(?:on|at|to|for)\s+([A-Za-z][A-Za-z0-9 &.'\-]{1,40})\s*[?.!]*$",
        q,
    )
    if m:
        name = m.group(1).strip()
        if name.lower() not in {"total", "overall", "everything", "this", "that"}:
            return name
    m = re.search(
        r"(?i)\b(?:merchant|vendor|paid\s+to|spent\s+on)\s+([A-Za-z][A-Za-z0-9 &.'\-]{1,40})",
        q,
    )
    if m:
        return m.group(1).strip()
    return None
