from fastapi import APIRouter, Request, Depends
from fastapi.responses import HTMLResponse
from fastapi.templating import Jinja2Templates
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, func, case
from app.core.database import get_db
from app.core.security import require_roles
from app.models.models import Customer, Invoice, Payment, Voucher, VoucherArchive, VoucherAgent
import os

router = APIRouter(dependencies=[Depends(require_roles("super_admin", "admin_cs"))])
BASE_DIR = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
templates = Jinja2Templates(directory=os.path.join(BASE_DIR, "app/templates"))


async def _collect_reports(db: AsyncSession) -> dict:
    """Build every figure the reports page needs in one pass."""
    # ---- Ringkasan pelanggan ----
    customers = await db.scalar(select(func.count(Customer.id)).where(Customer.is_archived == False))
    active_customers = await db.scalar(
        select(func.count(Customer.id)).where(Customer.status == "active", Customer.is_archived == False)
    )

    # ---- Invoice ----
    unpaid_invoices = await db.scalar(
        select(func.count(Invoice.id)).where(Invoice.status == "unpaid")
    )
    invoice_total = await db.scalar(
        select(func.coalesce(func.sum(Invoice.amount), 0)).where(Invoice.status == "unpaid")
    )

    # ---- Pembayaran (pelanggan: PPPoE & hotspot langganan) ----
    # LEFT JOIN customer supaya pembayaran tanpa pelanggan tetap ikut terhitung
    # dan ditandai, bukan hilang diam-diam dari pembukuan.
    payment_rows = (
        await db.execute(
            select(
                Customer.name,
                Customer.service_mode,
                Payment.amount,
                Payment.period,
                Payment.method,
                Payment.status,
                Payment.paid_at,
            )
            .select_from(Payment)
            .outerjoin(Customer, Customer.id == Payment.customer_id)
            .order_by(Payment.paid_at.desc())
        )
    ).all()

    payments_total = sum(r.amount or 0 for r in payment_rows)

    # Rekap per pelanggan (nama kosong = data tanpa pelanggan)
    per_customer: dict[str, dict] = {}
    for name, mode, amount, period, method, status, paid_at in payment_rows:
        key = name or "(Tanpa Pelanggan)"
        row = per_customer.setdefault(key, {"name": key, "mode": mode, "total": 0, "count": 0})
        if row["mode"] is None:
            row["mode"] = mode
        row["total"] += amount or 0
        row["count"] += 1

    # Rekap per periode
    per_period: dict[str, dict] = {}
    for _n, _m, amount, period, _me, _s, _p in payment_rows:
        row = per_period.setdefault(period or "-", {"period": period or "-", "total": 0, "count": 0})
        row["total"] += amount or 0
        row["count"] += 1

    # Rekap per metode pembayaran
    per_method: dict[str, dict] = {}
    for _n, _m, amount, _pe, method, _s, _p in payment_rows:
        key = (method or "-").upper()
        row = per_method.setdefault(key, {"method": key, "total": 0, "count": 0})
        row["total"] += amount or 0
        row["count"] += 1

    orphan_payments = sum(1 for r in payment_rows if not r[0])

    # ---- Voucher aktif ----
    voucher_available = await db.scalar(
        select(func.count(Voucher.id)).where(Voucher.status == "available")
    )
    voucher_sold = await db.scalar(select(func.count(Voucher.id)).where(Voucher.status == "sold"))
    voucher_used = await db.scalar(select(func.count(Voucher.id)).where(Voucher.status == "used"))

    # ---- Laporan keuangan voucher (dari arsip) ----
    archive_rows = (
        await db.execute(
            select(
                VoucherArchive.code,
                VoucherArchive.profile,
                VoucherArchive.price,
                VoucherArchive.status,
                VoucherArchive.reason,
                VoucherArchive.used_seconds,
                VoucherArchive.used_at,
                VoucherArchive.deleted_at,
                VoucherArchive.agent_id,
                VoucherAgent.name,
                VoucherAgent.area,
                VoucherAgent.commission_rate,
            )
            .select_from(VoucherArchive)
            .outerjoin(VoucherAgent, VoucherAgent.id == VoucherArchive.agent_id)
            .order_by(VoucherArchive.deleted_at.desc())
        )
    ).all()

    voucher_revenue = sum(r[2] or 0 for r in archive_rows)

    # Komisi dihitung dari rate agen saat ini (belum ada snapshot historis).
    per_agent: dict[str, dict] = {}
    for _c, _p, price, _s, reason, _us, _ua, _da, _aid, agent_name, area, rate in archive_rows:
        key = agent_name or "Direct / Tanpa Agen"
        row = per_agent.setdefault(
            key,
            {
                "name": key,
                "area": area or "-",
                "rate": rate or 0,
                "qty": 0,
                "gross": 0,
                "commission": 0,
            },
        )
        if row["area"] == "-" and area:
            row["area"] = area
        row["qty"] += 1
        row["gross"] += price or 0
        row["commission"] += int((price or 0) * (rate or 0) / 100)
    for row in per_agent.values():
        row["net"] = row["gross"] - row["commission"]

    # Rekap per paket
    per_profile: dict[str, dict] = {}
    for _c, profile, price, _s, _r, _us, _ua, _da, *_rest in archive_rows:
        key = profile or "-"
        row = per_profile.setdefault(key, {"profile": key, "qty": 0, "gross": 0})
        row["qty"] += 1
        row["gross"] += price or 0

    # Rekap per alasan selesai
    reason_labels = {
        "expired": "Masa Berlaku Habis",
        "duration_exhausted": "Jatah Waktu Habis",
        "gone_from_device": "Lepas dari Perangkat",
    }
    per_reason: dict[str, dict] = {}
    for _c, _p, price, _s, reason, *_rest in archive_rows:
        key = reason or "-"
        row = per_reason.setdefault(
            key, {"reason": key, "label": reason_labels.get(key, key), "qty": 0, "gross": 0}
        )
        row["qty"] += 1
        row["gross"] += price or 0

    total_commission = sum(r["commission"] for r in per_agent.values())

    return {
        # kartu ringkasan
        "customers": customers,
        "active_customers": active_customers,
        "unpaid_invoices": unpaid_invoices,
        "voucher_available": voucher_available,
        "voucher_sold": voucher_sold,
        "voucher_used": voucher_used,
        # keuangan pelanggan
        "invoice_total": invoice_total,
        "payments_total": payments_total,
        "payment_rows": payment_rows,
        "per_customer": sorted(per_customer.values(), key=lambda r: -r["total"]),
        "per_period": sorted(per_period.values(), key=lambda r: str(r["period"]), reverse=True),
        "per_method": sorted(per_method.values(), key=lambda r: -r["total"]),
        "orphan_payments": orphan_payments,
        # keuangan voucher
        "voucher_revenue": voucher_revenue,
        "voucher_archive_count": len(archive_rows),
        "per_agent": sorted(per_agent.values(), key=lambda r: -r["gross"]),
        "per_profile": sorted(per_profile.values(), key=lambda r: -r["gross"]),
        "per_reason": sorted(per_reason.values(), key=lambda r: -r["gross"]),
        "total_commission": total_commission,
        "voucher_net": voucher_revenue - total_commission,
        "grand_total": payments_total + voucher_revenue,
        "archive_rows": archive_rows,
    }


@router.get("/reports", response_class=HTMLResponse)
async def reports_page(request: Request, db: AsyncSession = Depends(get_db)):
    data = await _collect_reports(db)
    return templates.TemplateResponse(request=request, name="reports.html", context={"data": data})