"""Structured pandas queries over the finance CSV (Phase 3).

Handles filter-style questions that vector search alone answers poorly, e.g.:
- AP transactions over $500 in Q3 2023
- Companies in oil drilling
- Average debt-to-equity from analytics
"""

from __future__ import annotations

import json
import re
from functools import lru_cache
from pathlib import Path

import pandas as pd

from app.config import get_settings

_ACCOUNT_ALIASES = {
    "ap": "Accounts Payable",
    "accounts payable": "Accounts Payable",
    "a/p": "Accounts Payable",
    "ar": "Accounts Receivable",
    "accounts receivable": "Accounts Receivable",
    "a/r": "Accounts Receivable",
    "cash": "Cash",
    "inventory": "Inventory",
    "revenue": "Revenue",
}

_TYPE_ALIASES = {
    "sale": "Sale",
    "sales": "Sale",
    "purchase": "Purchase",
    "purchases": "Purchase",
    "transfer": "Transfer",
    "transfers": "Transfer",
}

_STRUCTURED_HINT = re.compile(
    r"\b(over|above|under|below|greater than|less than|between|"
    r"q[1-4]|quarter|202[0-9]|accounts?\s+payable|accounts?\s+receivable|"
    r"\ba/?p\b|\ba/?r\b|industry|companies?\s+in|debt[- ]to[- ]equity|"
    r"profit margin|average|sum|total|count|how many)\b",
    re.IGNORECASE,
)


def is_structured_query(question: str) -> bool:
    return bool(_STRUCTURED_HINT.search(question))


@lru_cache
def _load_frames() -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:
    path = Path(get_settings().dataset_path)
    df = pd.read_csv(path, low_memory=False)

    def parse_meta(raw: str) -> dict:
        try:
            return json.loads(raw) if isinstance(raw, str) else {}
        except json.JSONDecodeError:
            return {}

    metas = df["metadata_json"].map(parse_meta)
    meta_df = pd.json_normalize(metas)
    full = pd.concat([df.reset_index(drop=True), meta_df], axis=1)

    tx = full[full["data_category"] == "accounting_transaction"].copy()
    if not tx.empty:
        tx["Date"] = pd.to_datetime(tx.get("Date"), errors="coerce")
        for col in ("Debit", "Credit"):
            if col in tx.columns:
                tx[col] = pd.to_numeric(tx[col], errors="coerce")

    companies = full[full["data_category"] == "company_financial_statements"].copy()
    analytics = full[full["data_category"] == "accounting_analytics"].copy()
    for col in ("Debt-to-Equity Ratio", "Profit Margin", "Revenue", "Net Income"):
        if col in analytics.columns:
            analytics[col] = pd.to_numeric(analytics[col], errors="coerce")

    return tx, companies, analytics


def _extract_amount_threshold(question: str) -> tuple[str | None, float | None]:
    q = question.lower().replace(",", "")
    m = re.search(
        r"(?:over|above|greater than|more than|>)\s*\$?\s*(\d+(?:\.\d+)?)",
        q,
    )
    if m:
        return "gt", float(m.group(1))
    m = re.search(
        r"(?:under|below|less than|<)\s*\$?\s*(\d+(?:\.\d+)?)",
        q,
    )
    if m:
        return "lt", float(m.group(1))
    return None, None


def _extract_account(question: str) -> str | None:
    q = question.lower()
    for key, value in sorted(_ACCOUNT_ALIASES.items(), key=lambda x: -len(x[0])):
        if re.search(rf"\b{re.escape(key)}\b", q):
            return value
    return None


def _extract_tx_type(question: str) -> str | None:
    q = question.lower()
    for key, value in _TYPE_ALIASES.items():
        if re.search(rf"\b{re.escape(key)}\b", q):
            return value
    return None


def _extract_quarter(question: str) -> tuple[int, int] | None:
    """Return (year, quarter) if present."""
    q = question.lower()
    year = None
    ym = re.search(r"\b(20[12]\d)\b", q)
    if ym:
        year = int(ym.group(1))
    qm = re.search(r"\bq([1-4])\b|\bquarter\s*([1-4])\b", q)
    if qm:
        quarter = int(qm.group(1) or qm.group(2))
        return (year or 2023), quarter
    return None


def _quarter_range(year: int, quarter: int) -> tuple[pd.Timestamp, pd.Timestamp]:
    start_month = (quarter - 1) * 3 + 1
    start = pd.Timestamp(year=year, month=start_month, day=1)
    if quarter == 4:
        end = pd.Timestamp(year=year, month=12, day=31)
    else:
        end = pd.Timestamp(year=year, month=start_month + 3, day=1) - pd.Timedelta(days=1)
    return start, end


def _extract_industry(question: str) -> str | None:
    q = question.lower()
    m = re.search(
        r"(?:companies?\s+in|industry(?:\s+of)?|in the)\s+([a-z0-9 /&\-]+?)(?:\s+industry)?(?:\?|$)",
        q,
    )
    if m:
        return m.group(1).strip(" .?")
    # common short phrases
    for phrase in ("oil drilling", "oil", "banking", "software", "pharma", "steel"):
        if phrase in q:
            return phrase
    return None


def _format_tx_rows(rows: pd.DataFrame, limit: int = 8) -> str:
    lines = []
    for _, r in rows.head(limit).iterrows():
        date = r.get("Date")
        date_s = date.strftime("%Y-%m-%d") if pd.notna(date) else "?"
        lines.append(
            f"- {date_s} | {r.get('Account', '?')} | Debit {r.get('Debit', '?')} | "
            f"Type {r.get('Transaction_Type', '?')} | {r.get('Customer_Vendor', '?')} | "
            f"Payment {r.get('Payment_Method', '?')}"
        )
    return "\n".join(lines)


def query_transactions(question: str) -> str | None:
    tx, _, _ = _load_frames()
    if tx.empty:
        return None

    filtered = tx
    account = _extract_account(question)
    if account and "Account" in filtered.columns:
        filtered = filtered[filtered["Account"].astype(str).str.lower() == account.lower()]

    tx_type = _extract_tx_type(question)
    if tx_type and "Transaction_Type" in filtered.columns:
        filtered = filtered[filtered["Transaction_Type"].astype(str) == tx_type]

    op, amount = _extract_amount_threshold(question)
    if op and amount is not None and "Debit" in filtered.columns:
        if op == "gt":
            filtered = filtered[filtered["Debit"] > amount]
        else:
            filtered = filtered[filtered["Debit"] < amount]

    quarter = _extract_quarter(question)
    if quarter and "Date" in filtered.columns:
        start, end = _quarter_range(*quarter)
        filtered = filtered[(filtered["Date"] >= start) & (filtered["Date"] <= end)]

    # Only treat as a hit if we applied at least one filter beyond empty
    applied = any([account, tx_type, amount is not None, quarter])
    if not applied:
        return None

    count = len(filtered)
    if count == 0:
        bits = []
        if account:
            bits.append(account)
        if amount is not None:
            bits.append(f"{'>' if op == 'gt' else '<'} {amount}")
        if quarter:
            bits.append(f"Q{quarter[1]} {quarter[0]}")
        return (
            f"Structured transaction search found 0 matching records"
            f" ({', '.join(bits) or 'filters applied'})."
        )

    total_debit = float(filtered["Debit"].sum()) if "Debit" in filtered.columns else 0.0
    sample = _format_tx_rows(filtered.sort_values("Debit", ascending=False))
    header = f"Structured transaction search: {count:,} matching records"
    if amount is not None:
        header += f" (amount {'>' if op == 'gt' else '<'} {amount})"
    if account:
        header += f" on {account}"
    if quarter:
        header += f" in Q{quarter[1]} {quarter[0]}"
    header += f". Total debit of matches: {total_debit:,.2f}."
    return f"{header}\nExamples:\n{sample}"


def query_companies(question: str) -> str | None:
    _, companies, _ = _load_frames()
    if companies.empty or "Industry" not in companies.columns:
        return None

    industry = _extract_industry(question)
    if not industry:
        return None

    matches = companies[
        companies["Industry"].astype(str).str.contains(industry, case=False, na=False)
    ]
    if matches.empty:
        return f"Structured company search found 0 companies matching industry '{industry}'."

    lines = []
    for _, r in matches.head(10).iterrows():
        lines.append(
            f"- {r.get('Name', '?')} | Industry: {r.get('Industry', '?')} | "
            f"Price: {r.get('Current Price', '?')} | Sales: {r.get('Sales', '?')} | "
            f"PAT: {r.get('Profit after tax', '?')} | ROCE: {r.get('Return on capital employed', '?')}"
        )
    return (
        f"Structured company search: {len(matches)} companies matching '{industry}'.\n"
        + "\n".join(lines)
    )


def query_analytics(question: str) -> str | None:
    _, _, analytics = _load_frames()
    if analytics.empty:
        return None

    q = question.lower()
    wants_dte = "debt" in q and "equity" in q
    wants_margin = "profit margin" in q or "margin" in q
    wants_avg = any(w in q for w in ("average", "avg", "mean", "typical"))

    if not (wants_dte or wants_margin):
        return None

    lines = []
    if wants_dte and "Debt-to-Equity Ratio" in analytics.columns:
        series = analytics["Debt-to-Equity Ratio"].dropna()
        if not series.empty:
            lines.append(
                f"Debt-to-Equity across analytics records: count={len(series)}, "
                f"average={series.mean():.3f}, median={series.median():.3f}, "
                f"min={series.min():.3f}, max={series.max():.3f}."
            )
    if wants_margin and "Profit Margin" in analytics.columns:
        series = analytics["Profit Margin"].dropna()
        if not series.empty:
            lines.append(
                f"Profit Margin across analytics records: count={len(series)}, "
                f"average={series.mean():.3f}, median={series.median():.3f}, "
                f"min={series.min():.3f}, max={series.max():.3f}."
            )

    if not lines:
        return None
    if wants_avg and not lines:
        return None
    return "Structured analytics summary:\n" + "\n".join(lines)


def run_structured_query(question: str) -> str | None:
    """Return a compact structured summary, or None if not applicable."""
    if not is_structured_query(question):
        return None

    parts: list[str] = []
    for fn in (query_transactions, query_companies, query_analytics):
        try:
            result = fn(question)
        except Exception:
            result = None
        if result:
            parts.append(result)

    if not parts:
        return None
    return "\n\n".join(parts)
