from fastapi import APIRouter, Depends, Query
from sqlalchemy import text, func
from sqlalchemy.orm import Session
from typing import Optional
import time
from datetime import datetime, timedelta, timezone

from ..database import get_db, engine
from ..models import (
    TransactionLog, DailyShift, PumpReading, CreditCustomer,
    StockDelivery, User, AuditLog,
)
from ..dependencies import require_owner as require_super_admin

router = APIRouter()


@router.get("/monitoring/health")
def health_check(
    db: Session = Depends(get_db),
    _=Depends(require_super_admin),
):
    db_ok = False
    db_error = None
    try:
        db.execute(text("SELECT 1"))
        db_ok = True
    except Exception as e:
        db_error = str(e)

    recent_errors = (
        db.query(func.count(TransactionLog.id))
        .filter(TransactionLog.status_code >= 400)
        .scalar()
    ) or 0

    # Avg response time over last 1 hour
    one_hour_ago = datetime.now(timezone.utc) - timedelta(hours=1)
    avg_response_ms = (
        db.query(func.avg(TransactionLog.duration_ms))
        .filter(TransactionLog.created_at >= one_hour_ago)
        .scalar()
    )
    avg_response_ms = round(float(avg_response_ms), 1) if avg_response_ms is not None else 0.0

    # Total requests today (since midnight UTC)
    today_midnight = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
    total_requests_today = (
        db.query(func.count(TransactionLog.id))
        .filter(TransactionLog.created_at >= today_midnight)
        .scalar()
    ) or 0

    # Requests per minute (last 5 minutes)
    five_min_ago = datetime.now(timezone.utc) - timedelta(minutes=5)
    count_5min = (
        db.query(func.count(TransactionLog.id))
        .filter(TransactionLog.created_at >= five_min_ago)
        .scalar()
    ) or 0
    requests_per_minute = round(count_5min / 5, 2)

    return {
        "database": {"ok": db_ok, "error": db_error},
        "recent_error_transactions": recent_errors,
        "status": "healthy" if db_ok else "degraded",
        "avg_response_ms": avg_response_ms,
        "total_requests_today": total_requests_today,
        "requests_per_minute": requests_per_minute,
    }


@router.get("/monitoring/transactions")
def list_transactions(
    limit: int = 50,
    db: Session = Depends(get_db),
    _=Depends(require_super_admin),
):
    rows = (
        db.query(TransactionLog)
        .order_by(TransactionLog.created_at.desc())
        .limit(limit)
        .all()
    )
    return [
        {
            "id": r.id,
            "user_id": r.user_id,
            "endpoint": r.endpoint,
            "method": r.method,
            "status_code": r.status_code,
            "duration_ms": r.duration_ms,
            "ip_address": r.ip_address,
            "created_at": r.created_at,
        }
        for r in rows
    ]


@router.get("/monitoring/stats")
def monitoring_stats(
    db: Session = Depends(get_db),
    _=Depends(require_super_admin),
):
    return {
        "users":           db.query(func.count(User.id)).scalar(),
        "shifts":          db.query(func.count(DailyShift.id)).scalar(),
        "pump_readings":   db.query(func.count(PumpReading.id)).scalar(),
        "credit_customers":db.query(func.count(CreditCustomer.id)).scalar(),
        "stock_deliveries":db.query(func.count(StockDelivery.id)).scalar(),
        "audit_entries":   db.query(func.count(AuditLog.id)).scalar(),
        "transaction_logs":db.query(func.count(TransactionLog.id)).scalar(),
    }


@router.get("/monitoring/response-times")
def response_times(
    hours: int = Query(24, ge=1, le=720),
    db: Session = Depends(get_db),
    _=Depends(require_super_admin),
):
    since = datetime.now(timezone.utc) - timedelta(hours=hours)
    db.execute(text("SET SESSION group_concat_max_len = 65536"))
    rows = db.execute(text("""
        SELECT DATE_FORMAT(created_at, '%Y-%m-%d %H:00:00') AS hour,
               AVG(duration_ms) AS avg_ms,
               MAX(duration_ms) AS max_ms,
               COUNT(*) AS request_count,
               GROUP_CONCAT(duration_ms ORDER BY duration_ms) AS sorted_durations
        FROM transaction_logs
        WHERE created_at >= :since
        GROUP BY DATE_FORMAT(created_at, '%Y-%m-%d %H:00:00')
        ORDER BY hour ASC
    """), {"since": since}).fetchall()

    result = []
    for row in rows:
        avg_ms = round(float(row.avg_ms), 1) if row.avg_ms is not None else 0.0
        max_ms = int(row.max_ms) if row.max_ms is not None else 0
        p95_ms = 0
        if row.sorted_durations:
            durations = [int(x) for x in row.sorted_durations.split(',') if x]
            if durations:
                idx = int(len(durations) * 0.95)
                idx = min(idx, len(durations) - 1)
                p95_ms = durations[idx]
        result.append({
            "hour": row.hour,
            "avg_ms": avg_ms,
            "max_ms": max_ms,
            "p95_ms": p95_ms,
            "request_count": int(row.request_count),
        })
    return result


@router.get("/monitoring/error-rates")
def error_rates(
    hours: int = Query(24, ge=1, le=720),
    db: Session = Depends(get_db),
    _=Depends(require_super_admin),
):
    since = datetime.now(timezone.utc) - timedelta(hours=hours)
    rows = db.execute(text("""
        SELECT DATE_FORMAT(created_at, '%Y-%m-%d %H:00:00') AS hour,
               COUNT(*) AS total,
               SUM(CASE WHEN status_code >= 400 THEN 1 ELSE 0 END) AS errors
        FROM transaction_logs
        WHERE created_at >= :since
        GROUP BY DATE_FORMAT(created_at, '%Y-%m-%d %H:00:00')
        ORDER BY hour ASC
    """), {"since": since}).fetchall()

    result = []
    for row in rows:
        total = int(row.total)
        errors = int(row.errors) if row.errors else 0
        error_rate_pct = round(errors / total * 100, 1) if total > 0 else 0.0
        result.append({
            "hour": row.hour,
            "total": total,
            "errors": errors,
            "error_rate_pct": error_rate_pct,
        })
    return result


@router.get("/monitoring/endpoint-stats")
def endpoint_stats(
    db: Session = Depends(get_db),
    _=Depends(require_super_admin),
):
    rows = db.execute(text("""
        SELECT endpoint, method, COUNT(*) AS call_count,
               AVG(duration_ms) AS avg_ms, MAX(duration_ms) AS max_ms,
               SUM(CASE WHEN status_code >= 400 THEN 1 ELSE 0 END) AS error_count
        FROM transaction_logs
        GROUP BY endpoint, method
        ORDER BY call_count DESC
        LIMIT 10
    """)).fetchall()

    result = []
    for row in rows:
        call_count = int(row.call_count)
        error_count = int(row.error_count) if row.error_count else 0
        error_rate_pct = round(error_count / call_count * 100, 1) if call_count > 0 else 0.0
        result.append({
            "endpoint": row.endpoint,
            "method": row.method,
            "call_count": call_count,
            "avg_ms": round(float(row.avg_ms), 1) if row.avg_ms is not None else 0.0,
            "max_ms": int(row.max_ms) if row.max_ms is not None else 0,
            "error_count": error_count,
            "error_rate_pct": error_rate_pct,
        })
    return result


@router.post("/monitoring/run-checks")
def run_health_checks(
    db: Session = Depends(get_db),
    _=Depends(require_super_admin),
):
    results = {}
    # DB ping
    try:
        db.execute(text("SELECT 1"))
        results["db_ping"] = "ok"
    except Exception as e:
        results["db_ping"] = f"error: {e}"

    # Check for unclosed shifts
    from ..services.notification_service import check_unclosed_shifts, check_credit_limits, check_low_stock
    try:
        results["unclosed_shifts"] = check_unclosed_shifts(db)
    except Exception as e:
        results["unclosed_shifts"] = f"error: {e}"
    try:
        results["credit_limits"] = check_credit_limits(db)
    except Exception as e:
        results["credit_limits"] = f"error: {e}"
    try:
        results["low_stock"] = check_low_stock(db)
    except Exception as e:
        results["low_stock"] = f"error: {e}"

    return results
