from __future__ import annotations

from datetime import datetime, timedelta, timezone
from typing import Optional

from sqlalchemy import select, func, and_
from sqlalchemy.orm import Session

from src.apps.admin.models.scheduler_config import SchedulerConfig
from src.apps.admin.models.scheduler_log import SchedulerLog


def get_scheduler_config(db: Session, task_name: str) -> Optional[SchedulerConfig]:
    return db.execute(
        select(SchedulerConfig).where(SchedulerConfig.task_name == task_name)
    ).scalar_one_or_none()


def get_all_scheduler_configs(db: Session) -> list[SchedulerConfig]:
    return list(
        db.execute(select(SchedulerConfig).order_by(SchedulerConfig.display_name)).scalars().all()
    )


def update_scheduler_config(db: Session, task_name: str, data: dict, admin_id: int) -> SchedulerConfig:
    config = get_scheduler_config(db, task_name)
    for key, value in data.items():
        if value is not None:
            setattr(config, key, value)
    config.last_modified_by_admin_id = admin_id
    db.flush()
    db.refresh(config)
    return config


def create_scheduler_log(db: Session, data: dict) -> SchedulerLog:
    log = SchedulerLog(**data)
    db.add(log)
    db.flush()
    db.refresh(log)
    return log


def get_scheduler_log(db: Session, log_id: int) -> Optional[SchedulerLog]:
    return db.execute(
        select(SchedulerLog).where(SchedulerLog.id == log_id)
    ).scalar_one_or_none()


def update_scheduler_log(db: Session, log_id: int, data: dict) -> SchedulerLog:
    log = get_scheduler_log(db, log_id)
    for key, value in data.items():
        setattr(log, key, value)
    db.flush()
    db.refresh(log)
    return log


def get_scheduler_logs(
    db: Session,
    task_name: Optional[str] = None,
    status: Optional[str] = None,
    start_date: Optional[datetime] = None,
    end_date: Optional[datetime] = None,
    merchant_id: Optional[int] = None,
    page: int = 1,
    page_size: int = 20,
) -> tuple[list[SchedulerLog], int]:
    conditions = []
    if task_name:
        conditions.append(SchedulerLog.task_name == task_name)
    if status:
        conditions.append(SchedulerLog.run_status == status)
    if start_date:
        conditions.append(SchedulerLog.started_at >= start_date)
    if end_date:
        conditions.append(SchedulerLog.started_at <= end_date)
    if merchant_id is not None:
        conditions.append(
            SchedulerLog.merchant_ids.contains([merchant_id])
        )

    base = select(SchedulerLog)
    if conditions:
        base = base.where(and_(*conditions))

    total_stmt = select(func.count()).select_from(base.subquery())
    total: int = db.execute(total_stmt).scalar_one()

    offset = (page - 1) * page_size
    rows = list(
        db.execute(
            base.order_by(SchedulerLog.started_at.desc()).offset(offset).limit(page_size)
        ).scalars().all()
    )
    return rows, total


def get_last_run(db: Session, task_name: str) -> Optional[SchedulerLog]:
    return db.execute(
        select(SchedulerLog)
        .where(SchedulerLog.task_name == task_name)
        .order_by(SchedulerLog.started_at.desc())
        .limit(1)
    ).scalar_one_or_none()


def get_last_run_by_status(db: Session, task_name: str, status: str) -> Optional[SchedulerLog]:
    return db.execute(
        select(SchedulerLog)
        .where(SchedulerLog.task_name == task_name, SchedulerLog.run_status == status)
        .order_by(SchedulerLog.started_at.desc())
        .limit(1)
    ).scalar_one_or_none()


def get_success_rate_30d(db: Session, task_name: str) -> float:
    cutoff = datetime.now(timezone.utc) - timedelta(days=30)
    total_stmt = select(func.count()).where(
        SchedulerLog.task_name == task_name,
        SchedulerLog.started_at >= cutoff,
        SchedulerLog.run_status.in_(["SUCCESS", "FAILED"]),
    )
    success_stmt = select(func.count()).where(
        SchedulerLog.task_name == task_name,
        SchedulerLog.started_at >= cutoff,
        SchedulerLog.run_status == "SUCCESS",
    )
    total: int = db.execute(total_stmt).scalar_one() or 0
    success: int = db.execute(success_stmt).scalar_one() or 0
    if total == 0:
        return 1.0
    return round(success / total, 4)


def get_total_runs_30d(db: Session, task_name: str) -> int:
    cutoff = datetime.now(timezone.utc) - timedelta(days=30)
    return db.execute(
        select(func.count(SchedulerLog.id)).where(
            SchedulerLog.task_name == task_name,
            SchedulerLog.created_at >= cutoff,
        )
    ).scalar() or 0


def get_health_summary(db: Session) -> dict:
    configs = get_all_scheduler_configs(db)
    total = len(configs)
    healthy = 0
    failing = 0
    paused = 0

    for config in configs:
        if not config.is_enabled:
            paused += 1
            continue
        last_failed = get_last_run_by_status(db, config.task_name, "FAILED")
        last_success = get_last_run_by_status(db, config.task_name, "SUCCESS")
        if last_failed and (not last_success or last_failed.started_at > last_success.started_at):
            failing += 1
        else:
            healthy += 1

    # Recent failures (last 3)
    recent_failures = list(
        db.execute(
            select(SchedulerLog)
            .where(SchedulerLog.run_status == "FAILED")
            .order_by(SchedulerLog.started_at.desc())
            .limit(3)
        ).scalars().all()
    )

    return {
        "total": total,
        "healthy": healthy,
        "failing": failing,
        "paused": paused,
        "recent_failures": recent_failures,
    }
