"""
Analytics & Reporting Router
Provides aggregated metrics, trend data, forecasting, and anomaly detection
for the Fuel Station Management System.
"""
from fastapi import APIRouter, Depends, Query
from sqlalchemy.orm import Session
from sqlalchemy import func, and_, extract, case
from datetime import date, timedelta
from typing import Optional

from ..database import get_db
from ..models import (
    DailyShift, PumpReading, Pump, Staff, Allowance,
    LeaveRequest, CashHandover,
)
from ..dependencies import get_current_user

router = APIRouter()


# ── Helpers ──────────────────────────────────────────────────────────────────

def _resolve_range(
    date_from: Optional[date],
    date_to:   Optional[date],
    default_days: int = 30,
) -> tuple[date, date]:
    today  = date.today()
    d_to   = date_to   or today
    d_from = date_from or (today - timedelta(days=default_days - 1))
    return d_from, d_to


def _linear_regression(values: list[float]) -> tuple[float, float]:
    """Returns (slope, intercept) of a least-squares line through the series."""
    n = len(values)
    if n < 2:
        return 0.0, values[-1] if values else 0.0
    x_mean = (n - 1) / 2
    y_mean = sum(values) / n
    num    = sum((i - x_mean) * (v - y_mean) for i, v in enumerate(values))
    den    = sum((i - x_mean) ** 2 for i in range(n))
    slope  = num / den if den else 0.0
    return slope, y_mean - slope * x_mean


def _moving_avg(values: list[float], window: int = 7) -> list[float]:
    result = []
    for i in range(len(values)):
        chunk = values[max(0, i - window + 1): i + 1]
        result.append(sum(chunk) / len(chunk))
    return result


def _detect_anomalies(
    rows: list[dict],
    field: str,
    threshold: float = 2.0,
) -> list[dict]:
    """Flag records where |value – mean| > threshold × std_dev."""
    values = [r[field] for r in rows]
    if len(values) < 4:
        return []
    mean = sum(values) / len(values)
    variance = sum((v - mean) ** 2 for v in values) / len(values)
    std = variance ** 0.5
    if std == 0:
        return []
    return [r for r in rows if abs(r[field] - mean) > threshold * std]


# ── Endpoints ────────────────────────────────────────────────────────────────

@router.get("/analytics/overview")
def overview(
    date_from: Optional[date] = Query(None),
    date_to:   Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """
    High-level KPI summary for the selected period, plus comparison
    against the preceding period of the same length.
    """
    d_from, d_to = _resolve_range(date_from, date_to, 30)
    delta = (d_to - d_from).days + 1
    prev_from = d_from - timedelta(days=delta)
    prev_to   = d_from - timedelta(days=1)

    def _agg(f: date, t: date) -> dict:
        rows = (
            db.query(DailyShift)
            .filter(DailyShift.record_date.between(f, t))
            .all()
        )
        n = len(rows)
        if not rows:
            return dict(total_sale=0, cash=0, card=0, credit=0,
                        shortage=0, shifts=0, avg_daily=0)
        total_sale = sum(float(r.total_sale_calc) for r in rows)
        cash       = sum(float(r.cash_collected)  for r in rows)
        card       = sum(float(r.card_visa) + float(r.card_amex) + float(r.card_touch) for r in rows)
        credit     = sum(float(r.credit_total)    for r in rows)
        shortage   = sum(float(r.shortage)        for r in rows)
        return dict(
            total_sale = round(total_sale, 2),
            cash       = round(cash, 2),
            card       = round(card, 2),
            credit     = round(credit, 2),
            shortage   = round(shortage, 2),
            shifts     = n,
            avg_daily  = round(total_sale / delta, 2),
        )

    cur  = _agg(d_from, d_to)
    prev = _agg(prev_from, prev_to)

    def _pct_change(cur_v: float, prev_v: float) -> float | None:
        if prev_v == 0:
            return None
        return round((cur_v - prev_v) / prev_v * 100, 1)

    # Best shift type in period
    shift_agg = (
        db.query(
            DailyShift.shift_type,
            func.sum(DailyShift.total_sale_calc).label('total'),
        )
        .filter(DailyShift.record_date.between(d_from, d_to))
        .group_by(DailyShift.shift_type)
        .order_by(func.sum(DailyShift.total_sale_calc).desc())
        .first()
    )

    # Total liters in period
    ltr_agg = (
        db.query(func.coalesce(func.sum(PumpReading.sale_ltr), 0))
        .join(DailyShift, PumpReading.shift_id == DailyShift.id)
        .filter(DailyShift.record_date.between(d_from, d_to))
        .scalar()
    )

    return {
        "period":          {"from": str(d_from), "to": str(d_to), "days": delta},
        "current":         cur,
        "previous":        prev,
        "pct_change":      {k: _pct_change(cur[k], prev[k]) for k in cur},
        "best_shift_type": shift_agg.shift_type if shift_agg else None,
        "total_liters":    round(float(ltr_agg or 0), 2),
    }


@router.get("/analytics/revenue-trend")
def revenue_trend(
    date_from: Optional[date] = Query(None),
    date_to:   Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """Daily revenue series with 7-day moving average."""
    d_from, d_to = _resolve_range(date_from, date_to, 30)

    rows = (
        db.query(
            DailyShift.record_date.label("day"),
            func.sum(DailyShift.total_sale_calc).label("total_sale"),
            func.sum(DailyShift.cash_collected).label("cash"),
            func.sum(DailyShift.card_visa + DailyShift.card_amex + DailyShift.card_touch).label("card"),
            func.sum(DailyShift.credit_total).label("credit"),
            func.sum(DailyShift.shortage).label("shortage"),
            func.count(DailyShift.id).label("shifts"),
        )
        .filter(DailyShift.record_date.between(d_from, d_to))
        .group_by(DailyShift.record_date)
        .order_by(DailyShift.record_date)
        .all()
    )

    # Fill missing dates with zeros
    date_map: dict[date, dict] = {}
    for r in rows:
        date_map[r.day] = {
            "date":       str(r.day),
            "total_sale": round(float(r.total_sale or 0), 2),
            "cash":       round(float(r.cash or 0), 2),
            "card":       round(float(r.card or 0), 2),
            "credit":     round(float(r.credit or 0), 2),
            "shortage":   round(float(r.shortage or 0), 2),
            "shifts":     r.shifts,
        }

    series = []
    cur = d_from
    while cur <= d_to:
        series.append(date_map.get(cur, {
            "date": str(cur), "total_sale": 0, "cash": 0, "card": 0,
            "credit": 0, "shortage": 0, "shifts": 0,
        }))
        cur += timedelta(days=1)

    sales = [r["total_sale"] for r in series]
    ma    = _moving_avg(sales, 7)
    for i, row in enumerate(series):
        row["moving_avg"] = round(ma[i], 2)

    return series


@router.get("/analytics/fuel-breakdown")
def fuel_breakdown(
    date_from: Optional[date] = Query(None),
    date_to:   Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """Sales by fuel type: liters, revenue, and daily trend per type."""
    d_from, d_to = _resolve_range(date_from, date_to, 30)

    # Totals per fuel type
    totals = (
        db.query(
            Pump.fuel_type,
            func.coalesce(func.sum(PumpReading.sale_ltr),    0).label("total_ltr"),
            func.coalesce(func.sum(PumpReading.sale_amount), 0).label("total_rs"),
            func.count(PumpReading.id).label("readings"),
        )
        .join(PumpReading, PumpReading.pump_id == Pump.id)
        .join(DailyShift,  PumpReading.shift_id == DailyShift.id)
        .filter(DailyShift.record_date.between(d_from, d_to))
        .group_by(Pump.fuel_type)
        .all()
    )

    # Daily per fuel type (for trend chart)
    daily = (
        db.query(
            DailyShift.record_date.label("day"),
            Pump.fuel_type,
            func.coalesce(func.sum(PumpReading.sale_ltr),    0).label("ltr"),
            func.coalesce(func.sum(PumpReading.sale_amount), 0).label("rs"),
        )
        .join(PumpReading, PumpReading.pump_id == Pump.id)
        .join(DailyShift,  PumpReading.shift_id == DailyShift.id)
        .filter(DailyShift.record_date.between(d_from, d_to))
        .group_by(DailyShift.record_date, Pump.fuel_type)
        .order_by(DailyShift.record_date)
        .all()
    )

    daily_map: dict[str, dict] = {}
    for r in daily:
        key = str(r.day)
        if key not in daily_map:
            daily_map[key] = {"date": key}
        daily_map[key][r.fuel_type + "_ltr"] = round(float(r.ltr), 2)
        daily_map[key][r.fuel_type + "_rs"]  = round(float(r.rs),  2)

    daily_series = sorted(daily_map.values(), key=lambda x: x["date"])

    return {
        "totals": [
            {
                "fuel_type":  r.fuel_type,
                "total_ltr":  round(float(r.total_ltr), 2),
                "total_rs":   round(float(r.total_rs),  2),
                "readings":   r.readings,
                "avg_rate":   round(float(r.total_rs) / float(r.total_ltr), 2) if float(r.total_ltr) > 0 else 0,
            }
            for r in totals
        ],
        "daily": daily_series,
    }


@router.get("/analytics/shift-comparison")
def shift_comparison(
    date_from: Optional[date] = Query(None),
    date_to:   Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """Per-shift-type aggregates and efficiency metrics."""
    d_from, d_to = _resolve_range(date_from, date_to, 30)

    rows = (
        db.query(
            DailyShift.shift_type,
            func.count(DailyShift.id).label("count"),
            func.sum(DailyShift.total_sale_calc).label("total_sale"),
            func.sum(DailyShift.cash_collected).label("cash"),
            func.sum(
                DailyShift.card_visa + DailyShift.card_amex + DailyShift.card_touch
            ).label("card"),
            func.sum(DailyShift.credit_total).label("credit"),
            func.sum(DailyShift.shortage).label("shortage"),
            func.sum(DailyShift.advance).label("advance"),
        )
        .filter(DailyShift.record_date.between(d_from, d_to))
        .group_by(DailyShift.shift_type)
        .all()
    )

    ltr_rows = (
        db.query(
            DailyShift.shift_type,
            func.coalesce(func.sum(PumpReading.sale_ltr), 0).label("ltr"),
        )
        .join(PumpReading, PumpReading.shift_id == DailyShift.id)
        .filter(DailyShift.record_date.between(d_from, d_to))
        .group_by(DailyShift.shift_type)
        .all()
    )
    ltr_map = {r.shift_type: float(r.ltr) for r in ltr_rows}

    return [
        {
            "shift_type":   r.shift_type,
            "count":        r.count,
            "total_sale":   round(float(r.total_sale or 0), 2),
            "avg_sale":     round(float(r.total_sale or 0) / r.count, 2) if r.count else 0,
            "cash":         round(float(r.cash    or 0), 2),
            "card":         round(float(r.card    or 0), 2),
            "credit":       round(float(r.credit  or 0), 2),
            "shortage":     round(float(r.shortage or 0), 2),
            "advance":      round(float(r.advance or 0), 2),
            "total_liters": round(ltr_map.get(r.shift_type, 0), 2),
            "avg_liters":   round(ltr_map.get(r.shift_type, 0) / r.count, 2) if r.count else 0,
        }
        for r in rows
    ]


@router.get("/analytics/staff-analysis")
def staff_analysis(
    date_from: Optional[date] = Query(None),
    date_to:   Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """Per-pumper metrics: shifts, liters, sales, shortage rates."""
    d_from, d_to = _resolve_range(date_from, date_to, 30)

    staff_list = db.query(Staff).filter(Staff.status == 'ACTIVE').all()
    result = []

    for s in staff_list:
        # Shifts as manager
        as_manager = (
            db.query(func.count(DailyShift.id))
            .filter(
                DailyShift.staff_id == s.id,
                DailyShift.record_date.between(d_from, d_to),
            )
            .scalar() or 0
        )

        # Pump readings assigned to this staff
        pr_agg = (
            db.query(
                func.coalesce(func.sum(PumpReading.sale_ltr),    0).label("ltr"),
                func.coalesce(func.sum(PumpReading.sale_amount), 0).label("rs"),
                func.count(PumpReading.id).label("readings"),
            )
            .join(DailyShift, PumpReading.shift_id == DailyShift.id)
            .filter(
                PumpReading.staff_id == s.id,
                DailyShift.record_date.between(d_from, d_to),
            )
            .first()
        )

        # Shortage on shifts managed
        shortage_agg = (
            db.query(
                func.coalesce(func.sum(DailyShift.shortage), 0).label("total"),
                func.count(
                    case((DailyShift.shortage > 0, 1))
                ).label("cnt"),
            )
            .filter(
                DailyShift.staff_id == s.id,
                DailyShift.record_date.between(d_from, d_to),
            )
            .first()
        )

        total_ltr = float(pr_agg.ltr)
        total_rs  = float(pr_agg.rs)

        result.append({
            "staff_id":      s.id,
            "staff_name":    s.full_name,
            "as_manager":    as_manager,
            "pump_readings": pr_agg.readings,
            "total_liters":  round(total_ltr, 2),
            "total_sale":    round(total_rs,  2),
            "avg_ltr_per_reading": round(total_ltr / pr_agg.readings, 2) if pr_agg.readings else 0,
            "shortage_total": round(float(shortage_agg.total), 2),
            "shortage_count": shortage_agg.cnt,
        })

    result.sort(key=lambda x: x["total_liters"], reverse=True)
    return result


@router.get("/analytics/payment-methods")
def payment_methods(
    date_from: Optional[date] = Query(None),
    date_to:   Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """Daily payment method breakdown and totals."""
    d_from, d_to = _resolve_range(date_from, date_to, 30)

    rows = (
        db.query(
            DailyShift.record_date.label("day"),
            func.sum(DailyShift.cash_collected).label("cash"),
            func.sum(
                DailyShift.card_visa + DailyShift.card_amex + DailyShift.card_touch
            ).label("card"),
            func.sum(DailyShift.card_visa).label("visa"),
            func.sum(DailyShift.card_amex).label("amex"),
            func.sum(DailyShift.card_touch).label("touch"),
            func.sum(DailyShift.credit_total).label("credit"),
            func.sum(DailyShift.other_income).label("other"),
            func.sum(DailyShift.total_sale_calc).label("total_sale"),
        )
        .filter(DailyShift.record_date.between(d_from, d_to))
        .group_by(DailyShift.record_date)
        .order_by(DailyShift.record_date)
        .all()
    )

    series = []
    for r in rows:
        total = float(r.total_sale or 1)
        cash  = float(r.cash   or 0)
        card  = float(r.card   or 0)
        cred  = float(r.credit or 0)
        series.append({
            "date":        str(r.day),
            "cash":        round(cash,  2),
            "card":        round(card,  2),
            "visa":        round(float(r.visa  or 0), 2),
            "amex":        round(float(r.amex  or 0), 2),
            "touch":       round(float(r.touch or 0), 2),
            "credit":      round(cred,  2),
            "other":       round(float(r.other or 0), 2),
            "total_sale":  round(float(r.total_sale or 0), 2),
            "cash_pct":    round(cash / total * 100, 1),
            "card_pct":    round(card / total * 100, 1),
            "credit_pct":  round(cred / total * 100, 1),
        })

    totals = {
        "cash":   sum(r["cash"]   for r in series),
        "card":   sum(r["card"]   for r in series),
        "credit": sum(r["credit"] for r in series),
        "other":  sum(r["other"]  for r in series),
    }
    grand = sum(totals.values()) or 1
    totals["cash_pct"]   = round(totals["cash"]   / grand * 100, 1)
    totals["card_pct"]   = round(totals["card"]   / grand * 100, 1)
    totals["credit_pct"] = round(totals["credit"] / grand * 100, 1)

    return {"series": series, "totals": totals}


@router.get("/analytics/anomalies")
def anomalies(
    date_from: Optional[date] = Query(None),
    date_to:   Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """Shifts with unusual shortages, zero sales, or excess collections."""
    d_from, d_to = _resolve_range(date_from, date_to, 30)

    shifts = (
        db.query(DailyShift, Staff.full_name)
        .outerjoin(Staff, DailyShift.staff_id == Staff.id)
        .filter(DailyShift.record_date.between(d_from, d_to))
        .order_by(DailyShift.record_date.desc())
        .all()
    )

    rows = [
        {
            "id":         s.id,
            "date":       str(s.record_date),
            "shift_type": s.shift_type,
            "staff_name": name,
            "total_sale": float(s.total_sale_calc),
            "shortage":   float(s.shortage),
            "difference": float(s.difference),
            "cash":       float(s.cash_collected),
        }
        for s, name in shifts
    ]

    shortage_flags  = _detect_anomalies(rows, "shortage",  threshold=1.5)
    diff_flags      = _detect_anomalies(rows, "difference",threshold=1.5)
    zero_sale       = [r for r in rows if r["total_sale"] == 0]

    flagged_ids = set()
    alerts = []

    for r in shortage_flags:
        if r["shortage"] > 0 and r["id"] not in flagged_ids:
            alerts.append({**r, "flag": "high_shortage",   "severity": "warning"})
            flagged_ids.add(r["id"])

    for r in diff_flags:
        if abs(r["difference"]) > 1000 and r["id"] not in flagged_ids:
            alerts.append({**r, "flag": "large_difference", "severity": "info"})
            flagged_ids.add(r["id"])

    for r in zero_sale:
        if r["id"] not in flagged_ids:
            alerts.append({**r, "flag": "zero_sale",        "severity": "error"})
            flagged_ids.add(r["id"])

    alerts.sort(key=lambda x: (x["date"], x["severity"]), reverse=True)
    return alerts


@router.get("/analytics/forecast")
def forecast(
    horizon: int = Query(14, ge=1, le=30),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """
    Revenue forecast for the next `horizon` days using linear regression
    on the past 60 days of data, adjusted by 7-day moving average.
    """
    today  = date.today()
    d_from = today - timedelta(days=59)

    rows = (
        db.query(
            DailyShift.record_date.label("day"),
            func.sum(DailyShift.total_sale_calc).label("total"),
        )
        .filter(DailyShift.record_date.between(d_from, today))
        .group_by(DailyShift.record_date)
        .order_by(DailyShift.record_date)
        .all()
    )

    # Fill gaps
    date_map = {r.day: float(r.total or 0) for r in rows}
    full_series: list[float] = []
    dates: list[date] = []
    cur = d_from
    while cur <= today:
        full_series.append(date_map.get(cur, 0))
        dates.append(cur)
        cur += timedelta(days=1)

    slope, intercept = _linear_regression(full_series)
    ma = _moving_avg(full_series, 7)

    n = len(full_series)

    # Build historical output
    historical = [
        {
            "date":       str(dates[i]),
            "actual":     round(full_series[i], 2),
            "moving_avg": round(ma[i], 2),
            "type":       "actual",
        }
        for i in range(n)
    ]

    # Project forward
    projected = []
    for i in range(1, horizon + 1):
        x   = n - 1 + i
        val = max(0, intercept + slope * x)
        # Blend with moving avg of recent actuals (dampens extreme slopes)
        blended = (val * 0.5 + ma[-1] * 0.5)
        projected.append({
            "date":     str(today + timedelta(days=i)),
            "forecast": round(blended, 2),
            "low":      round(blended * 0.85, 2),
            "high":     round(blended * 1.15, 2),
            "type":     "forecast",
        })

    # Trend label
    daily_change = slope
    if abs(daily_change) < 500:
        trend = "stable"
    elif daily_change > 0:
        trend = "growing"
    else:
        trend = "declining"

    return {
        "trend":      trend,
        "slope":      round(slope, 2),
        "historical": historical[-30:],   # last 30 days for display
        "projected":  projected,
        "note":       "Forecast uses linear regression + 7-day moving average on last 60 days of data.",
    }


@router.get("/analytics/leave-summary")
def leave_summary(
    date_from: Optional[date] = Query(None),
    date_to:   Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    """Leave utilization by type and staff for the period."""
    d_from, d_to = _resolve_range(date_from, date_to, 30)

    rows = (
        db.query(
            LeaveRequest.leave_type,
            func.count(LeaveRequest.id).label("requests"),
            func.sum(LeaveRequest.days_count).label("days"),
        )
        .filter(
            LeaveRequest.status == 'APPROVED',
            LeaveRequest.date_from.between(d_from, d_to),
        )
        .group_by(LeaveRequest.leave_type)
        .all()
    )

    by_staff = (
        db.query(
            LeaveRequest.staff_id,
            Staff.full_name,
            func.sum(LeaveRequest.days_count).label("days"),
            func.count(LeaveRequest.id).label("requests"),
        )
        .join(Staff, LeaveRequest.staff_id == Staff.id)
        .filter(
            LeaveRequest.status == 'APPROVED',
            LeaveRequest.date_from.between(d_from, d_to),
        )
        .group_by(LeaveRequest.staff_id, Staff.full_name)
        .order_by(func.sum(LeaveRequest.days_count).desc())
        .all()
    )

    return {
        "by_type": [
            {"leave_type": r.leave_type, "requests": r.requests, "days": r.days}
            for r in rows
        ],
        "by_staff": [
            {"staff_id": r.staff_id, "staff_name": r.full_name,
             "days": r.days, "requests": r.requests}
            for r in by_staff
        ],
    }
