from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy import func, extract
from sqlalchemy.orm import Session, joinedload
from typing import List, Optional

from ..database import get_db
from ..models import Expense
from ..schemas import ExpenseCreate, ExpenseUpdate, ExpenseOut
from ..dependencies import get_current_user

router = APIRouter()


def _expense_out(e: Expense) -> dict:
    d = {c.key: getattr(e, c.key) for c in e.__table__.columns}
    d['staff_name'] = e.staff_member.full_name if e.staff_member else None
    return d


@router.get("/expenses", response_model=List[ExpenseOut])
def list_expenses(
    date_from: Optional[str] = None,
    date_to: Optional[str] = None,
    category: Optional[str] = None,
    shift_id: Optional[int] = None,
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    q = db.query(Expense).options(joinedload(Expense.staff_member))
    if date_from:
        q = q.filter(Expense.expense_date >= date_from)
    if date_to:
        q = q.filter(Expense.expense_date <= date_to)
    if category:
        q = q.filter(Expense.category == category)
    if shift_id:
        q = q.filter(Expense.shift_id == shift_id)
    return [_expense_out(e) for e in q.order_by(Expense.expense_date.desc(), Expense.id.desc()).all()]


@router.post("/expenses", response_model=ExpenseOut)
def create_expense(
    body: ExpenseCreate,
    db: Session = Depends(get_db),
    user=Depends(get_current_user),
):
    expense = Expense(**body.model_dump(), created_by=user.id)
    db.add(expense)
    db.flush()
    from .audit import log_audit
    log_audit(db, user.id, 'CREATE', 'expenses', expense.id,
              f"Expense: {body.category} Rs {body.amount} on {body.expense_date}")
    db.commit()
    db.refresh(expense)
    return _expense_out(expense)


@router.put("/expenses/{expense_id}", response_model=ExpenseOut)
def update_expense(
    expense_id: int,
    body: ExpenseUpdate,
    db: Session = Depends(get_db),
    current_user=Depends(get_current_user),
):
    expense = db.query(Expense).filter(Expense.id == expense_id).first()
    if not expense:
        raise HTTPException(404, "Expense not found")
    for k, v in body.model_dump(exclude_none=True).items():
        setattr(expense, k, v)
    from .audit import log_audit
    log_audit(db, current_user.id, 'UPDATE', 'expenses', expense_id,
              f"Updated expense: {expense.category} Rs {expense.amount}")
    db.commit()
    db.refresh(expense)
    return expense


@router.delete("/expenses/{expense_id}")
def delete_expense(
    expense_id: int,
    db: Session = Depends(get_db),
    current_user=Depends(get_current_user),
):
    expense = db.query(Expense).filter(Expense.id == expense_id).first()
    if not expense:
        raise HTTPException(404, "Expense not found")
    from .audit import log_audit
    log_audit(db, current_user.id, 'DELETE', 'expenses', expense_id,
              f"Deleted expense: {expense.category} Rs {expense.amount}")
    db.delete(expense)
    db.commit()
    return {"ok": True}


@router.get("/expenses/summary")
def expense_summary(
    month: int,
    year: int,
    db: Session = Depends(get_db),
    _=Depends(get_current_user),
):
    rows = (
        db.query(Expense.category, func.sum(Expense.amount).label('total'))
        .filter(
            extract('month', Expense.expense_date) == month,
            extract('year',  Expense.expense_date) == year,
        )
        .group_by(Expense.category)
        .all()
    )
    grand_total = sum(float(r.total) for r in rows)
    return {
        "month": month,
        "year": year,
        "by_category": [{"category": r.category, "total": float(r.total)} for r in rows],
        "grand_total": grand_total,
    }
