from fastapi import APIRouter, Depends, Query
from sqlalchemy.orm import Session
from sqlalchemy import func, and_
from datetime import date, timedelta
from typing import Optional
from ..database import get_db
from ..models import DailyShift, PumpReading, Pump, Staff
from ..schemas import DashboardToday, FuelBreakdown
from ..dependencies import get_current_user

router = APIRouter()


@router.get("/today", response_model=DashboardToday)
def today_summary(
    for_date: Optional[date] = Query(None),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    target = for_date or date.today()

    shifts = (
        db.query(DailyShift)
        .filter(DailyShift.record_date == target)
        .all()
    )

    total_sale   = sum(float(s.total_sale_calc) for s in shifts)
    cash_total   = sum(float(s.cash_collected)  for s in shifts)
    card_total   = sum(float(s.card_visa) + float(s.card_amex) + float(s.card_touch) for s in shifts)
    credit_total = sum(float(s.credit_total)    for s in shifts)
    shortage_total = sum(float(s.shortage)      for s in shifts)

    # Fuel breakdown from pump_readings
    fuel_rows = (
        db.query(
            Pump.fuel_type,
            func.sum(PumpReading.sale_ltr).label('total_ltr'),
            func.sum(PumpReading.sale_amount).label('total_rs'),
        )
        .join(PumpReading, PumpReading.pump_id == Pump.id)
        .join(DailyShift, PumpReading.shift_id == DailyShift.id)
        .filter(DailyShift.record_date == target)
        .group_by(Pump.fuel_type)
        .all()
    )
    fuel_breakdown = [
        FuelBreakdown(
            fuel_type=r.fuel_type,
            total_ltr=float(r.total_ltr or 0),
            total_rs=float(r.total_rs or 0),
        )
        for r in fuel_rows
    ]
    total_ltr_today = sum(float(r.total_ltr or 0) for r in fuel_rows)

    # Recent shifts (last 10 across all dates)
    recent = (
        db.query(DailyShift, Staff.full_name)
        .outerjoin(Staff, DailyShift.staff_id == Staff.id)
        .order_by(DailyShift.record_date.desc(), DailyShift.id.desc())
        .limit(10)
        .all()
    )
    recent_shifts = [
        {
            "id":             s.id,
            "record_date":    str(s.record_date),
            "shift_type":     s.shift_type,
            "staff_name":     name,
            "total_sale_calc": float(s.total_sale_calc),
            "difference":     float(s.difference),
            "is_locked":      s.is_locked,
            "status":         s.status,
        }
        for s, name in recent
    ]

    return DashboardToday(
        date             = target,
        total_sale       = total_sale,
        cash_total       = cash_total,
        card_total       = card_total,
        credit_total     = credit_total,
        shortage_total   = shortage_total,
        shifts_completed = len(shifts),
        total_ltr_today  = total_ltr_today,
        fuel_breakdown   = fuel_breakdown,
        recent_shifts    = recent_shifts,
    )


@router.get("/summary")
def weekly_summary(
    days: int = Query(7, ge=1, le=90),
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    from_date = date.today() - timedelta(days=days - 1)

    rows = (
        db.query(
            DailyShift.record_date,
            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.count(DailyShift.id).label('shift_count'),
        )
        .filter(DailyShift.record_date >= from_date)
        .group_by(DailyShift.record_date)
        .order_by(DailyShift.record_date)
        .all()
    )
    return [
        {
            "date":        str(r.record_date),
            "total_sale":  float(r.total_sale or 0),
            "cash":        float(r.cash or 0),
            "card":        float(r.card or 0),
            "credit":      float(r.credit or 0),
            "shift_count": r.shift_count,
        }
        for r in rows
    ]


@router.get("/mtd")
def mtd_summary(
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    today = date.today()
    month_start = today.replace(day=1)
    prev_month_end = month_start - timedelta(days=1)
    prev_month_start = prev_month_end.replace(day=1)

    mtd_total = db.query(func.sum(DailyShift.total_sale_calc)).filter(
        DailyShift.record_date >= month_start,
        DailyShift.record_date <= today,
    ).scalar() or 0

    prev_total = db.query(func.sum(DailyShift.total_sale_calc)).filter(
        DailyShift.record_date >= prev_month_start,
        DailyShift.record_date <= prev_month_end,
    ).scalar() or 0

    return {
        "mtd_total":        float(mtd_total),
        "prev_month_total": float(prev_total),
        "month":            today.strftime("%B"),
        "prev_month":       prev_month_end.strftime("%B"),
        "days_elapsed":     today.day,
    }
