"""Deterministic answers from uploaded KB statements (not vector-snippet guessing)."""

from __future__ import annotations

import re
from pathlib import Path
from typing import Any

from app.pdf.calc import (
    _money_out,
    _pretty_money,
    _detect_currency_symbol,
    answer_statement_metric_question,
    calculate_from_statement_totals,
    describe_pdf_job,
    format_job_context,
    is_pdf_overview_question,
)
from app.pdf.extract import extract_pdf
from app.saas import db as saas_db
from app.saas.kb import org_storage_dir

_MONTH_RE = re.compile(
    r"(?i)\b(jan(?:uary)?|feb(?:ruary)?|mar(?:ch)?|apr(?:il)?|may|jun(?:e)?|"
    r"jul(?:y)?|aug(?:ust)?|sep(?:t(?:ember)?)?|oct(?:ober)?|nov(?:ember)?|dec(?:ember)?)\b"
)
_MONTH_ALIASES = {
    "jan": "jan", "january": "jan",
    "feb": "feb", "february": "feb",
    "mar": "mar", "march": "mar",
    "apr": "apr", "april": "apr",
    "may": "may",
    "jun": "jun", "june": "jun",
    "jul": "jul", "july": "jul",
    "aug": "aug", "august": "aug",
    "sep": "sep", "sept": "sep", "september": "sep",
    "oct": "oct", "october": "oct",
    "nov": "nov", "november": "nov",
    "dec": "dec", "december": "dec",
}

# Cache extract results so chat stays fast after first hit.
# Bump CACHE_VERSION when extraction logic changes so stale bad totals are dropped.
_CACHE_VERSION = "hdfc-v2"
_JOB_CACHE: dict[str, tuple[float, str, dict[str, Any]]] = {}
_JOBS_LIST_CACHE: dict[str, tuple[float, list[dict[str, Any]]]] = {}
_JOBS_LIST_TTL_S = 60.0


def _pdf_path_for_doc(org_id: str, agent_id: str, doc: dict[str, Any]) -> Path | None:
    folder = org_storage_dir(org_id, agent_id)
    matches = sorted(folder.glob(f"{doc['id']}_*"))
    for path in matches:
        if path.suffix.lower() == ".pdf":
            return path
    return None


def _row_dicts(result) -> list[dict[str, Any]]:
    rows: list[dict[str, Any]] = []
    for r in result.rows:
        if hasattr(r, "__dict__"):
            rows.append(dict(r.__dict__))
        elif isinstance(r, dict):
            rows.append(r)
    return rows


def build_statement_job(path: Path, filename: str) -> dict[str, Any] | None:
    key = str(path.resolve())
    try:
        mtime = path.stat().st_mtime
    except OSError:
        return None
    cached = _JOB_CACHE.get(key)
    if cached and cached[0] == mtime and cached[1] == _CACHE_VERSION:
        return cached[2]

    result = extract_pdf(path, filename=filename)
    if result.document_kind not in {"bank_statement", "unknown"} and not result.statement_totals:
        # Still usable if header totals exist
        if not result.rows:
            return None

    rows = _row_dicts(result)
    totals = result.statement_totals or {}
    calc = calculate_from_statement_totals(totals, rows=rows) if totals else None
    if calc is None:
        from app.pdf.calc import calculate as calc_rows

        calc = calc_rows(rows, document_kind=result.document_kind or "bank_statement")

    job = {
        "filename": filename,
        "document_kind": result.document_kind or "bank_statement",
        "source_type": result.source_type,
        "ocr_used": result.ocr_used,
        "statement_totals": totals,
        "rows": rows,
        "calculation": calc,
        "document_text": (result.document_text or "")[:100_000],
        "text_preview": result.text_preview,
        "validation": result.validation,
    }
    _JOB_CACHE[key] = (mtime, _CACHE_VERSION, job)
    return job


def list_agent_statement_jobs(org_id: str, agent_id: str) -> list[dict[str, Any]]:
    import time

    cache_key = f"{org_id}:{agent_id}"
    hit = _JOBS_LIST_CACHE.get(cache_key)
    now = time.monotonic()
    if hit and (now - hit[0]) < _JOBS_LIST_TTL_S:
        return hit[1]

    jobs: list[dict[str, Any]] = []
    for doc in saas_db.list_kb_documents(org_id, agent_id):
        if doc.get("status") != "ready":
            continue
        if not str(doc.get("filename") or "").lower().endswith(".pdf"):
            continue
        path = _pdf_path_for_doc(org_id, agent_id, doc)
        if not path or not path.exists():
            continue
        job = build_statement_job(path, doc["filename"])
        if job:
            jobs.append(job)
    _JOBS_LIST_CACHE[cache_key] = (now, jobs)
    return jobs


def warm_local_statement_caches(*, limit: int = 8) -> int:
    """Pre-extract recent local KB PDFs so the first spend question is fast after restart."""
    from app.config import get_settings

    root = get_settings().chroma_persist_dir.parent / "kb_uploads"
    if not root.exists():
        return 0
    pdfs = sorted(root.rglob("*.pdf"), key=lambda p: p.stat().st_mtime, reverse=True)
    warmed = 0
    for path in pdfs[:limit]:
        try:
            # filename after first underscore is original name: {doc_id}_{filename}
            name = path.name
            if "_" in name:
                name = name.split("_", 1)[1]
            if build_statement_job(path, name):
                warmed += 1
        except Exception:
            continue
    return warmed


def _month_token(question: str) -> str | None:
    m = _MONTH_RE.search(question or "")
    if not m:
        return None
    return _MONTH_ALIASES.get(m.group(1).lower())


def _count_rows_in_month(rows: list[dict[str, Any]], month: str) -> int:
    n = 0
    for r in rows:
        blob = " ".join(
            str(r.get(k) or "") for k in ("date", "description", "raw", "account")
        ).lower()
        if re.search(rf"\b{month}\b", blob) or f" {month}'" in blob or f"{month}'" in blob:
            n += 1
    return n


def describe_knowledge_base(org_id: str, agent_id: str) -> str | None:
    """Deterministic overview of uploaded KB docs — no LLM snippet guessing."""
    docs = [d for d in saas_db.list_kb_documents(org_id, agent_id) if d.get("status") == "ready"]
    if not docs:
        return None
    lines = [
        f"Your knowledge base has {len(docs)} ready document(s):",
    ]
    for d in docs[:8]:
        name = d.get("filename") or "document"
        chunks = d.get("chunk_count") or 0
        lines.append(f"- “{name}” ({chunks} chunks)")
    if len(docs) > 8:
        lines.append(f"- …and {len(docs) - 8} more")

    jobs = list_agent_statement_jobs(org_id, agent_id)
    for job in jobs[:2]:
        lines.append(describe_pdf_job(job))
    lines.append(
        "Ask about spend, received, payment counts, a category like groceries, or a merchant on the statement."
    )
    return "\n".join(lines)


_DAY_MONTH_RE = re.compile(
    r"(?i)\b(?:on\s+|dated\s+)?"
    r"(?:"
    r"(?P<d1>\d{1,2})(?:st|nd|rd|th)?\s+(?P<m1>jan(?:uary)?|feb(?:ruary)?|mar(?:ch)?|apr(?:il)?|may|jun(?:e)?|jul(?:y)?|aug(?:ust)?|sep(?:t(?:ember)?)?|oct(?:ober)?|nov(?:ember)?|dec(?:ember)?)"
    r"|"
    r"(?P<m2>jan(?:uary)?|feb(?:ruary)?|mar(?:ch)?|apr(?:il)?|may|jun(?:e)?|jul(?:y)?|aug(?:ust)?|sep(?:t(?:ember)?)?|oct(?:ober)?|nov(?:ember)?|dec(?:ember)?)\s+(?P<d2>\d{1,2})(?:st|nd|rd|th)?"
    r")\b"
)

_SPEND_INTENT_RE = re.compile(
    r"(?i)\b(spent|spend|paid|payment|debit|withdrawn|outflow|money\s+out|how\s+much)\b"
)
_COUNT_INTENT_RE = re.compile(
    r"(?i)\b(how\s+many|number\s+of|count|payments?\s+made|transactions?\s+made)\b"
)
_PER_DATE_SPEND_RE = re.compile(
    r"(?i)\b("
    r"(?:spent|spend|paid|payment|outflow|money\s+out).{0,40}?"
    r"(?:each|every|per|by|on\s+each|for\s+each|date[\s-]?wise|day[\s-]?wise|daily)"
    r".{0,20}?(?:date|day|daily)"
    r"|(?:each|every|per|by|on\s+each|for\s+each|date[\s-]?wise|day[\s-]?wise|daily)"
    r".{0,20}?(?:date|day)"
    r".{0,40}?(?:spent|spend|paid|payment|outflow|money\s+out|how\s+much)"
    r"|spend\s+breakdown"
    r"|daily\s+spend"
    r"|day[\s-]?wise\s+spend"
    r"|date[\s-]?wise\s+spend"
    r")\b"
)
_MONTH_SORT = {
    "jan": 1, "feb": 2, "mar": 3, "apr": 4, "may": 5, "jun": 6,
    "jul": 7, "aug": 8, "sep": 9, "oct": 10, "nov": 11, "dec": 12,
}
_DATE_TOKEN_RE = re.compile(
    r"(?i)\b(\d{1,2})\s+(Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec)\b"
)


def _parse_day_month(question: str) -> tuple[int, str] | None:
    m = _DAY_MONTH_RE.search(question or "")
    if not m:
        return None
    if m.group("d1") and m.group("m1"):
        day = int(m.group("d1"))
        month = _MONTH_ALIASES.get(m.group("m1").lower())
    else:
        day = int(m.group("d2"))
        month = _MONTH_ALIASES.get(m.group("m2").lower())
    if not month or day < 1 or day > 31:
        return None
    return day, month


def _row_matches_day_month(row: dict[str, Any], day: int, month: str) -> bool:
    blob = " ".join(
        str(row.get(k) or "") for k in ("date", "description", "raw")
    ).lower()
    # Paytm-style: "26 Jul", "26 Jul'26", "26 July"
    patterns = [
        rf"\b{day}\s+{month}\b",
        rf"\b{day}\s+{month}'?\d{{0,4}}\b",
        rf"\b{day:02d}\s+{month}\b",
    ]
    return any(re.search(p, blob) for p in patterns)


def _sum_spend_from_text(text: str, day: int | None, month: str) -> tuple[float, int]:
    """Fallback: find outflow amounts near a date mention in raw PDF text."""
    total = 0.0
    hits = 0
    if not text:
        return total, hits

    month_pat = month[:3]  # jul / aug — Paytm prints 3-letter months
    # Blocks starting at a date like "26 Jul" / "06 Jul"
    chunks = re.split(
        rf"(?=(?:^|\n)\s*\d{{1,2}}\s+(?:Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec))",
        text,
        flags=re.I,
    )
    if len(chunks) <= 1:
        # Fallback: line-oriented scan
        chunks = text.splitlines()

    for chunk in chunks:
        low = chunk.lower()
        if day is not None:
            if not re.search(rf"\b0?{day}\s+{month_pat}\b", low):
                continue
        elif not re.search(rf"\b{month_pat}\b", low):
            continue
        # Skip pure credits on this chunk if only + amounts (still count - Rs)
        for m in re.finditer(r"(?i)-\s*(?:₹|rs\.?)\s*([\d,]+\.?\d*)", chunk):
            try:
                total += float(m.group(1).replace(",", ""))
                hits += 1
            except ValueError:
                pass
    return round(total, 2), hits


def _full_statement_text(job: dict[str, Any]) -> str:
    """Prefer in-memory text; if truncated/missing, re-extract from PDF path when available."""
    text = job.get("document_text") or job.get("text_preview") or ""
    if len(text) >= 500:
        return text
    return text


def is_per_date_spend_question(question: str) -> bool:
    q = (question or "").strip()
    if not q:
        return False
    if _PER_DATE_SPEND_RE.search(q):
        return True
    # Typo-tolerant: "spen on each date", "spent each day"
    if _SPEND_INTENT_RE.search(q) and re.search(
        r"(?i)\b(each|every|per|all)\s+(date|day|daily)\b|\bdate[\s-]?wise\b|\bday[\s-]?wise\b|\bdaily\b",
        q,
    ):
        return True
    return False


def _normalize_date_label(blob: str) -> str | None:
    m = _DATE_TOKEN_RE.search(blob or "")
    if not m:
        return None
    return f"{int(m.group(1)):02d} {m.group(2).title()[:3]}"


def _date_sort_key(label: str) -> tuple[int, int]:
    parts = label.split()
    if len(parts) != 2:
        return (99, 99)
    try:
        day = int(parts[0])
    except ValueError:
        day = 99
    return (_MONTH_SORT.get(parts[1].lower()[:3], 99), day)


def _spend_by_date_from_rows(rows: list[dict[str, Any]]) -> dict[str, dict[str, float | int]]:
    by: dict[str, dict[str, float | int]] = {}
    for r in rows:
        out = _money_out(r)
        if out <= 0:
            continue
        label = _normalize_date_label(str(r.get("date") or r.get("raw") or ""))
        if not label:
            continue
        bucket = by.setdefault(label, {"total": 0.0, "n": 0})
        bucket["total"] = float(bucket["total"]) + out
        bucket["n"] = int(bucket["n"]) + 1
    for v in by.values():
        v["total"] = round(float(v["total"]), 2)
    return by


def _spend_by_date_from_text(text: str) -> dict[str, dict[str, float | int]]:
    """Group Paytm-style statement outflows by calendar date from raw PDF text."""
    by: dict[str, dict[str, float | int]] = {}
    if not text:
        return by
    months = "Jan|Feb|Mar|Apr|May|Jun|Jul|Aug|Sep|Oct|Nov|Dec"
    chunks = re.split(
        rf"(?=(?:^|\n)\s*\d{{1,2}}\s+(?:{months})\b)",
        text,
        flags=re.I,
    )
    for chunk in chunks:
        m = re.match(rf"\s*(\d{{1,2}})\s+({months})\b", chunk, re.I)
        if not m:
            continue
        low = chunk.lower()
        # Skip header totals / notes, not transaction blocks.
        if "total money paid" in low or re.search(r"\bpayments?\s+made\b", low[:100]):
            continue
        if not any(
            k in low
            for k in (
                "upi ref",
                "money sent",
                "paid to",
                "received from",
                "recharge",
                "money received",
                "transferred to",
            )
        ):
            continue
        amts = [
            float(a.group(1).replace(",", ""))
            for a in re.finditer(r"(?i)-\s*(?:₹|rs\.?)\s*([\d,]+\.?\d*)", chunk)
        ]
        if not amts:
            continue
        # One primary outflow per transaction block (amount is usually last).
        amt = amts[0] if len(set(amts)) == 1 else amts[-1]
        label = f"{int(m.group(1)):02d} {m.group(2).title()[:3]}"
        bucket = by.setdefault(label, {"total": 0.0, "n": 0})
        bucket["total"] = float(bucket["total"]) + amt
        bucket["n"] = int(bucket["n"]) + 1
    for v in by.values():
        v["total"] = round(float(v["total"]), 2)
    return by


def answer_spend_by_date(job: dict[str, Any], question: str) -> str | None:
    """Deterministic daily spend breakdown — never leave this to the LLM."""
    if not is_per_date_spend_question(question):
        return None

    cur = _detect_currency_symbol(job)
    filename = job.get("filename") or "statement"
    header = job.get("statement_totals") or {}
    metrics = (job.get("calculation") or {}).get("metrics") or {}

    from_text = _spend_by_date_from_text(_full_statement_text(job))
    from_rows = _spend_by_date_from_rows(job.get("rows") or [])

    # Prefer the source that covers more payments / closer to statement header total.
    header_out = header.get("money_out")
    header_n = header.get("payments_made")
    if header_out is None:
        header_out = metrics.get("money_out")
    if header_n is None:
        header_n = metrics.get("payments_made")

    def _score(by: dict[str, dict[str, float | int]]) -> tuple[int, float]:
        n = sum(int(v["n"]) for v in by.values())
        total = round(sum(float(v["total"]) for v in by.values()), 2)
        return n, total

    text_n, text_total = _score(from_text)
    row_n, row_total = _score(from_rows)

    use_text = False
    if text_n and (not from_rows or text_n >= row_n):
        use_text = True
    if header_out is not None and from_text:
        try:
            target = float(header_out)
            if abs(text_total - target) <= abs(row_total - target):
                use_text = True
        except (TypeError, ValueError):
            pass

    by = from_text if use_text and from_text else from_rows
    if not by:
        return (
            f"I couldn't build a per-date spend breakdown from “{filename}”. "
            f"Try asking about a specific day, e.g. “how much did I spend on 16 July?”."
        )

    n_pay = sum(int(v["n"]) for v in by.values())
    total = round(sum(float(v["total"]) for v in by.values()), 2)
    lines = [
        f"Spend by date from “{filename}” ({n_pay} payment(s), {_pretty_money(total, cur)} total):",
    ]
    for label in sorted(by.keys(), key=_date_sort_key):
        bucket = by[label]
        count = int(bucket["n"])
        amt = _pretty_money(bucket["total"], cur)
        unit = "payment" if count == 1 else "payments"
        lines.append(f"- {label} — {amt} ({count} {unit})")

    if header_out is not None and header_n is not None:
        try:
            h_out = float(header_out)
            if abs(total - h_out) > 0.05 or int(header_n) != n_pay:
                lines.append(
                    f"Statement header lists {int(header_n)} payments / {_pretty_money(h_out, cur)} "
                    f"overall; this breakdown uses dated outflows found in the PDF."
                )
        except (TypeError, ValueError):
            pass
    lines.append('Ask about a specific day (e.g. "spend on 16 July") for merchant-level detail.')
    return "\n".join(lines)


def answer_dated_spend(job: dict[str, Any], question: str) -> str | None:
    """Answer spend for a specific day or month from extracted rows / statement text."""
    if not _SPEND_INTENT_RE.search(question or ""):
        return None
    # "each date" / daily breakdown is handled separately.
    if is_per_date_spend_question(question):
        return None

    day_month = _parse_day_month(question)
    month_only = _month_token(question) if not day_month else None
    if not day_month and not month_only:
        return None

    rows = job.get("rows") or []
    cur = _detect_currency_symbol(job)
    filename = job.get("filename") or "statement"

    if day_month:
        day, month = day_month
        label = f"{day} {month[:1].upper() + month[1:]}"
        date_key = f"{day:02d} {month.title()[:3]}"
        # Prefer full-text daily bucket (matches header coverage) over partial table rows.
        by_text = _spend_by_date_from_text(_full_statement_text(job))
        text_bucket = by_text.get(date_key)
        matched = [r for r in rows if _row_matches_day_month(r, day, month) and _money_out(r) > 0]
        if text_bucket and (
            not matched or int(text_bucket["n"]) >= len(matched)
        ):
            total = float(text_bucket["total"])
            n = int(text_bucket["n"])
            examples = []
            for r in matched[:4]:
                desc = str(r.get("description") or "payment").split("\n")[0][:50]
                examples.append(f"{desc} ({_pretty_money(_money_out(r), cur)})")
            extra = (" Examples: " + "; ".join(examples) + ".") if examples else ""
            return (
                f"Spent on {label} in “{filename}”: {_pretty_money(total, cur)} "
                f"across {n} payment(s).{extra}"
            )
        if matched:
            total = round(sum(_money_out(r) for r in matched), 2)
            examples = []
            for r in matched[:4]:
                desc = str(r.get("description") or "payment").split("\n")[0][:50]
                examples.append(f"{desc} ({_pretty_money(_money_out(r), cur)})")
            extra = (" Examples: " + "; ".join(examples) + ".") if examples else ""
            return (
                f"Spent on {label} in “{filename}”: {_pretty_money(total, cur)} "
                f"across {len(matched)} extracted payment(s).{extra}"
            )
        text_total, text_hits = _sum_spend_from_text(_full_statement_text(job), day, month)
        if text_hits:
            return (
                f"Spent on {label} in “{filename}”: {_pretty_money(text_total, cur)} "
                f"from {text_hits} outflow amount(s) found in the statement text."
            )
        return f"I couldn't find any money-out transactions dated {label} in “{filename}”."

    # Month-only spend (e.g. "how much did I spend in July")
    month = month_only  # type: ignore[assignment]
    assert month is not None
    month_label = month[:1].upper() + month[1:]
    matched = []
    for r in rows:
        blob = " ".join(str(r.get(k) or "") for k in ("date", "description", "raw")).lower()
        if re.search(rf"\b{month}\b", blob) and _money_out(r) > 0:
            matched.append(r)
    if matched:
        total = round(sum(_money_out(r) for r in matched), 2)
        return (
            f"Spent in {month_label} in “{filename}”: {_pretty_money(total, cur)} "
            f"across {len(matched)} extracted payment(s)."
        )
    text_total, text_hits = _sum_spend_from_text(_full_statement_text(job), None, month)
    if text_hits:
        return (
            f"Spent in {month_label} in “{filename}”: {_pretty_money(text_total, cur)} "
            f"from {text_hits} outflow amount(s) found in the statement text."
        )
    return None


def _month_aware_count_answer(job: dict[str, Any], question: str) -> str | None:
    # Do not answer spend questions with payment-count boilerplate.
    if _SPEND_INTENT_RE.search(question or "") and not _COUNT_INTENT_RE.search(question or ""):
        return None
    # Day-specific "how many on 26 July" can be handled later; skip full-period dump for day queries.
    if _parse_day_month(question):
        return None

    month = _month_token(question)
    if not month:
        return None
    rows = job.get("rows") or []
    header = job.get("statement_totals") or {}
    metrics = (job.get("calculation") or {}).get("metrics") or {}
    made = metrics.get("payments_made")
    if made is None:
        made = header.get("payments_made")
    recv = metrics.get("payments_received")
    if recv is None:
        recv = header.get("payments_received")
    month_label = month[:1].upper() + month[1:]
    n = _count_rows_in_month(rows, month) if rows else 0

    if not _COUNT_INTENT_RE.search(question or ""):
        # Mentions a month but isn't clearly a count question — don't hijack.
        return None

    bits = []
    if made is not None:
        bits.append(f"{made} payments made")
    if recv is not None:
        bits.append(f"{recv} payments received")
    if bits:
        base = (
            f"This statement lists {' and '.join(bits)} "
            f"for the full period printed on the uploaded PDF."
        )
        if n:
            base += f" Among extracted dated lines, {n} fall in {month_label}."
        return base
    if n == 0:
        return None
    return f"I counted {n} line items dated in {month_label} in “{job.get('filename')}”."


def try_deterministic_kb_answer(org_id: str, agent_id: str, question: str) -> str | None:
    """Answer spend/count/totals from uploaded statement PDFs when possible."""
    q = (question or "").strip()
    if not q:
        return None

    if re.search(r"(?i)\bknowledge\s*base\b", q) or is_pdf_overview_question(q):
        if re.search(r"(?i)\bknowledge\s*base\b", q):
            overview = describe_knowledge_base(org_id, agent_id)
            if overview:
                return overview

    jobs = list_agent_statement_jobs(org_id, agent_id)
    if not jobs:
        if re.search(r"(?i)\bknowledge\s*base\b", q):
            docs = saas_db.list_kb_documents(org_id, agent_id)
            ready = [d for d in docs if d.get("status") == "ready"]
            if not ready:
                return "Your knowledge base has no ready documents yet. Upload a PDF, TXT, MD, or CSV in Feed data."
        return None

    if is_pdf_overview_question(q) and not re.search(r"(?i)\bknowledge\s*base\b", q):
        return describe_pdf_job(jobs[0])

    # Daily / per-date spend before single-day, month, or whole-statement totals.
    for job in jobs:
        by_date = answer_spend_by_date(job, q)
        if by_date:
            return by_date

    # Date/month spend before count boilerplate.
    for job in jobs:
        dated = answer_dated_spend(job, q)
        if dated:
            return dated

    if _COUNT_INTENT_RE.search(q) and _month_token(q):
        for job in jobs:
            ans = _month_aware_count_answer(job, q)
            if ans:
                return ans

    for job in jobs:
        # Don't answer "each date" with whole-statement spend.
        if is_per_date_spend_question(q):
            continue
        ans = answer_statement_metric_question(job, q)
        if ans:
            return ans
    return None


def statement_context_notes(org_id: str, agent_id: str, *, max_jobs: int = 2) -> str:
    """Compact structured totals to prepend to RAG context."""
    jobs = list_agent_statement_jobs(org_id, agent_id)[:max_jobs]
    if not jobs:
        return ""
    blocks = []
    for job in jobs:
        # Keep this short for phi3 context
        ctx = format_job_context(job)
        blocks.append(ctx[:3500])
    return "\n\n---\n\n".join(blocks)
