"""Accountant / owner dashboard: total investment, total sales and net profit.

Definitions (also shown on the page so the owner can trust the numbers)
----------------------------------------------------------------------
Total Sales       = invoice value after discounts, minus sales returns
Cost of goods sold= buy price (at the time of sale) of the items sold, minus returns
Gross profit      = Total Sales - Cost of goods sold
Total Investment  = stock purchased + operating expenses + salaries paid
Net profit        = Gross profit - operating expenses - salaries paid
ROI               = Net profit / Total Investment
Profit margin     = Net profit / Total Sales

Everything follows the branch selector (one branch or all branches).
"""
import json
from datetime import date, datetime, time, timedelta
from decimal import Decimal

from django.db.models import (
    DecimalField, ExpressionWrapper, F, Min, OuterRef, Subquery, Sum, Value,
)
from django.db.models.functions import Coalesce, Greatest
from django.shortcuts import render
from django.utils import timezone

from accounts.branch_utils import get_branch_context, scope
from accounts.permissions import company_admin_required
from accounts.profit_utils import (
    financial_summary, return_financials, sale_line_financials,
)
from expenses.models import Expense
from inventory.models import Purchase
from inventory.stock import stock_summary
from payroll.models import Payslip
from payroll.services import payroll_cost, shift_month
from sales.models import Payment, Sale

ZERO = Decimal("0")
_MONEY = DecimalField(max_digits=14, decimal_places=2)
_ZERO_V = Value(ZERO, output_field=_MONEY)
_TOL = Decimal("0.005")

PERIODS = (
    ("today", "Today"),
    ("month", "This Month"),
    ("last_month", "Last Month"),
    ("quarter", "This Quarter"),
    ("year", "This Year"),
    ("all", "All Time"),
    ("custom", "Custom"),
)


# --------------------------------------------------------------------------
# helpers
# --------------------------------------------------------------------------

def _first_day(d):
    return d.replace(day=1)


def _last_day_of_month(y, m):
    ny, nm = shift_month(y, m, 1)
    return date(ny, nm, 1) - timedelta(days=1)


def resolve_period(get, today, earliest):
    """Return (key, start, end, label) for the requested period."""
    key = get.get("period", "all")
    if key not in dict(PERIODS):
        key = "all"
    if key == "today":
        start = end = today
    elif key == "month":
        start, end = _first_day(today), today
    elif key == "last_month":
        y, m = shift_month(today.year, today.month, -1)
        start, end = date(y, m, 1), _last_day_of_month(y, m)
    elif key == "quarter":
        qm = 3 * ((today.month - 1) // 3) + 1
        start, end = date(today.year, qm, 1), today
    elif key == "year":
        start, end = date(today.year, 1, 1), today
    elif key == "custom":
        from django.utils.dateparse import parse_date
        start = parse_date(get.get("from", "") or "") or _first_day(today)
        end = parse_date(get.get("to", "") or "") or today
        if start > end:
            start, end = end, start
    else:  # all
        start, end = earliest or today, today

    if key == "today":
        label = today.strftime("%d %b %Y")
    elif key == "all":
        label = f"All time (since {start:%d %b %Y})"
    else:
        label = f"{start:%d %b %Y} \u2013 {end:%d %b %Y}"
    return key, start, end, label


def _day_bounds(start, end):
    """Aware [start, end+1day) datetimes - safe on MySQL without timezone tables."""
    tz = timezone.get_current_timezone()
    lo = timezone.make_aware(datetime.combine(start, time.min), tz)
    hi = timezone.make_aware(datetime.combine(end + timedelta(days=1), time.min), tz)
    return lo, hi


def _earliest_date(company, branch):
    dates = []
    for qs in (
        Sale.objects.filter(company=company),
        Purchase.objects.filter(company=company),
        Expense.objects.filter(company=company),
    ):
        d = scope(qs, branch).aggregate(d=Min("date"))["d"]
        if d:
            dates.append(d)
    return min(dates) if dates else None


def _core_totals(company, start, end, branch):
    """Sales, cost, profit, investment for one period and scope."""
    s = financial_summary(company, start, end, branch)
    purchases = scope(
        Purchase.objects.filter(company=company, date__range=[start, end]), branch
    ).aggregate(t=Sum("net_amount"))["t"] or ZERO
    expenses = scope(
        Expense.objects.filter(company=company, date__range=[start, end]), branch
    ).aggregate(t=Sum("amount"))["t"] or ZERO
    salaries = payroll_cost(company, start, end, branch)
    purchases, expenses, salaries = Decimal(purchases), Decimal(expenses), Decimal(salaries)

    investment = purchases + expenses + salaries
    net_profit = s["gross_profit"] - expenses - salaries
    return {
        "sales": s["revenue"], "cogs": s["cogs"], "gross": s["gross_profit"],
        "invoices": s["sales_count"], "returns": -s["return_revenue"],
        "purchases": purchases, "expenses": expenses, "salaries": salaries,
        "investment": investment, "net": net_profit,
        "roi": (net_profit / investment * 100) if investment > 0 else None,
        "margin": (net_profit / s["revenue"] * 100) if s["revenue"] > 0 else None,
    }


def _receivables(company, branch):
    """Money customers still owe (all time - a balance, not a period figure)."""
    paid_sq = Payment.objects.filter(sale=OuterRef("pk")).order_by().values("sale").annotate(
        t=Sum("amount")
    ).values("t")
    qs = (
        scope(Sale.objects.filter(company=company, customer__isnull=False), branch)
        .annotate(paid_total=Coalesce(Subquery(paid_sq, output_field=_MONEY), _ZERO_V, output_field=_MONEY))
        .annotate(due_total=Greatest(
            ExpressionWrapper(F("net_amount") - F("paid_total"), output_field=_MONEY),
            _ZERO_V, output_field=_MONEY,
        ))
    )
    total = qs.aggregate(t=Sum("due_total"))["t"] or ZERO
    top = (
        qs.filter(due_total__gt=_TOL)
        .values("customer_id", "customer__name")
        .annotate(d=Sum("due_total")).order_by("-d")[:5]
    )
    return Decimal(total), [{"name": r["customer__name"], "due": r["d"]} for r in top]


def _collections(company, start, end, branch):
    lo, hi = _day_bounds(start, end)
    qs = scope(
        Payment.objects.filter(company=company, paid_at__gte=lo, paid_at__lt=hi),
        branch, field="sale__branch",
    )
    names = dict(Payment.PAYMENT_CHOICES)
    rows = [
        {"method": names.get(r["payment_method"], r["payment_method"]), "amount": r["t"]}
        for r in qs.values("payment_method").annotate(t=Sum("amount")).order_by("-t")
    ]
    top = max((r["amount"] for r in rows), default=ZERO)
    for r in rows:
        r["pct"] = int(r["amount"] / top * 100) if top else 0
    return sum((r["amount"] for r in rows), ZERO), rows


def _expense_categories(company, start, end, branch):
    qs = scope(Expense.objects.filter(company=company, date__range=[start, end]), branch)
    rows = [
        {"name": r["category__name"] or "Uncategorised", "amount": r["t"]}
        for r in qs.values("category__name").annotate(t=Sum("amount")).order_by("-t")[:6]
    ]
    top = max((r["amount"] for r in rows), default=ZERO)
    for r in rows:
        r["pct"] = int(r["amount"] / top * 100) if top else 0
    return rows


def _monthly_trend(company, branch, today, months=12):
    """Sales / investment / net profit for each of the last ``months`` months."""
    keys = [shift_month(today.year, today.month, -i) for i in range(months - 1, -1, -1)]
    start = date(keys[0][0], keys[0][1], 1)
    bucket = {k: {"sales": ZERO, "gross": ZERO, "inv": ZERO, "opex": ZERO} for k in keys}

    def add(d, field, amount):
        k = (d.year, d.month)
        if k in bucket:
            bucket[k][field] += Decimal(amount or 0)

    for r in sale_line_financials(company, start, today, branch):
        add(r["date"], "sales", r["net_revenue"]); add(r["date"], "gross", r["profit"])
    for r in return_financials(company, start, today, branch):
        add(r["date"], "sales", r["revenue"]); add(r["date"], "gross", r["profit"])
    for d, amt in scope(
        Purchase.objects.filter(company=company, date__range=[start, today]), branch
    ).values_list("date", "net_amount"):
        add(d, "inv", amt)
    for d, amt in scope(
        Expense.objects.filter(company=company, date__range=[start, today]), branch
    ).values_list("date", "amount"):
        add(d, "inv", amt); add(d, "opex", amt)
    for d, net, adv in scope(
        Payslip.objects.filter(
            company=company, status=Payslip.PAID, paid_date__range=[start, today]
        ), branch,
    ).values_list("paid_date", "net_salary", "advance_deduction"):
        cost = Decimal(net or 0) + Decimal(adv or 0)
        add(d, "inv", cost); add(d, "opex", cost)

    labels = [date(y, m, 1).strftime("%b %y") for y, m in keys]
    sales = [float(bucket[k]["sales"]) for k in keys]
    invest = [float(bucket[k]["inv"]) for k in keys]
    net = [float(bucket[k]["gross"] - bucket[k]["opex"]) for k in keys]
    return labels, sales, invest, net


# --------------------------------------------------------------------------
# the view
# --------------------------------------------------------------------------

@company_admin_required
def accountant_dashboard(request):
    company = request.user.company
    ctx = get_branch_context(request)
    branch = ctx.active
    today = timezone.now().date()   # same definition of 'today' as the rest of the app

    earliest = _earliest_date(company, branch)
    key, start, end, label = resolve_period(request.GET, today, earliest)

    t = _core_totals(company, start, end, branch)
    receivable, top_dues = _receivables(company, branch)
    collected, by_method = _collections(company, start, end, branch)
    stock = stock_summary(company, branch)
    categories = _expense_categories(company, start, end, branch)
    labels, tr_sales, tr_invest, tr_net = _monthly_trend(company, branch, today)

    branch_rows = []
    if branch is None and ctx.is_multi:
        for b in ctx.branches:
            bt = _core_totals(company, start, end, b)
            bdue, _ = _receivables(company, b)
            bt.update({
                "branch": b, "due": bdue,
                "stock_value": stock_summary(company, b)["total_value"],
            })
            branch_rows.append(bt)

    return render(request, "reports/accountant_dashboard.html", {
        "periods": PERIODS, "period": key, "period_label": label,
        "date_from": start, "date_to": end,
        "t": t, "is_loss": t["net"] < 0,
        "receivable": receivable, "top_dues": top_dues,
        "collected": collected, "by_method": by_method,
        "stock_value": stock["total_value"],
        "low_stock_count": len(stock["alert_products"]),
        "categories": categories,
        "branch_rows": branch_rows,
        "chart_labels": json.dumps(labels),
        "chart_sales": json.dumps(tr_sales),
        "chart_invest": json.dumps(tr_invest),
        "chart_net": json.dumps(tr_net),
        "invest_split": json.dumps([float(t["purchases"]), float(t["expenses"]), float(t["salaries"])]),
        "method_labels": json.dumps([r["method"] for r in by_method]),
        "method_values": json.dumps([float(r["amount"]) for r in by_method]),
    })
