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

from ..database import get_db
from ..models import AuditLog, User
from ..schemas import AuditLogOut
from ..dependencies import get_current_user

router = APIRouter()


def log_audit(db: Session, user_id: Optional[int], action: str,
              table_name: Optional[str], record_id: Optional[int], summary: str):
    """Call before db.commit() in the parent endpoint."""
    entry = AuditLog(
        user_id=user_id,
        action=action,
        table_name=table_name,
        record_id=record_id,
        summary=summary[:255] if summary else None,
    )
    db.add(entry)


@router.get("/audit-log", response_model=List[AuditLogOut])
def get_audit_log(
    date_from: Optional[str] = None,
    date_to: Optional[str] = None,
    table_name: Optional[str] = None,
    limit: int = 200,
    db: Session = Depends(get_db),
    user=Depends(get_current_user),
):
    if user.role not in ('OWNER', 'SUPER_ADMIN'):
        raise HTTPException(status_code=403, detail="Owner access required")

    q = db.query(AuditLog)
    if date_from:
        q = q.filter(AuditLog.created_at >= date_from)
    if date_to:
        q = q.filter(AuditLog.created_at <= date_to + ' 23:59:59')
    if table_name:
        q = q.filter(AuditLog.table_name == table_name)
    rows = q.order_by(AuditLog.created_at.desc()).limit(limit).all()

    # Bulk-load all referenced users in one query (avoids N+1)
    user_ids = {r.user_id for r in rows if r.user_id}
    user_map: dict = {}
    if user_ids:
        users = db.query(User).filter(User.id.in_(user_ids)).all()
        user_map = {u.id: u.full_name for u in users}

    return [
        AuditLogOut(
            id=row.id,
            user_id=row.user_id,
            actor_name=user_map.get(row.user_id) if row.user_id else None,
            action=row.action,
            table_name=row.table_name,
            record_id=row.record_id,
            summary=row.summary,
            created_at=row.created_at,
        )
        for row in rows
    ]
