from __future__ import annotations

import argparse
import logging
import os
import time
from datetime import datetime, timedelta
from typing import Iterable, Optional, Set

from sqlalchemy import select

from app.jobs.channel_jobs import enqueue_channel_job
from db.models import ChannelSchedule
from db.session import get_session


logger = logging.getLogger(__name__)


def next_run_after(cron_expression: str, after: Optional[datetime] = None) -> datetime:
    """Return the next UTC datetime for a simple five-field cron expression."""
    base = (after or datetime.utcnow()).replace(second=0, microsecond=0) + timedelta(minutes=1)
    minute_set = _field(cron_expression, 0, 0, 59)
    hour_set = _field(cron_expression, 1, 0, 23)
    day_set = _field(cron_expression, 2, 1, 31)
    month_set = _field(cron_expression, 3, 1, 12)
    weekday_set = _field(cron_expression, 4, 0, 6)

    current = base
    for _ in range(366 * 24 * 60):
        cron_weekday = (current.weekday() + 1) % 7
        if (
            current.minute in minute_set
            and current.hour in hour_set
            and current.day in day_set
            and current.month in month_set
            and cron_weekday in weekday_set
        ):
            return current
        current += timedelta(minutes=1)
    raise ValueError(f"Could not find next run for cron expression: {cron_expression}")


def enqueue_due_schedules(*, limit: int = 20, worker_id: Optional[str] = None) -> dict:
    worker = worker_id or f"channel-scheduler-{os.getpid()}"
    now = datetime.utcnow()
    enqueued = 0
    with get_session() as session:
        rows = session.scalars(
            select(ChannelSchedule)
            .where(ChannelSchedule.is_enabled.is_(True))
            .where(ChannelSchedule.next_run_at.is_not(None))
            .where(ChannelSchedule.next_run_at <= now)
            .order_by(ChannelSchedule.next_run_at, ChannelSchedule.id)
            .limit(limit)
            .with_for_update(skip_locked=True)
        ).all()
        for schedule in rows:
            schedule.locked_at = now
            schedule.locked_by = worker
            job = enqueue_channel_job(
                session,
                schedule_id=schedule.id,
                channel_connection_id=schedule.channel_connection_id,
                channel_type=schedule.channel_type,
                native_connection_id=schedule.native_connection_id,
                channel_code=schedule.channel_code,
                job_type=schedule.task_type,
                dry_run=schedule.dry_run,
                mode=schedule.mode,
            )
            schedule.last_enqueued_job_id = job.id
            schedule.last_status = "queued"
            schedule.last_error = None
            schedule.next_run_at = next_run_after(schedule.cron_expression, now)
            schedule.locked_at = None
            schedule.locked_by = None
            enqueued += 1
        session.commit()
    return {"status": "ok", "enqueued": enqueued}


def _field(expression: str, index: int, min_value: int, max_value: int) -> Set[int]:
    parts = str(expression or "").split()
    if len(parts) != 5:
        raise ValueError("cron_expression must have 5 fields: minute hour day month weekday")
    return _values(parts[index], min_value, max_value)


def _values(token: str, min_value: int, max_value: int) -> Set[int]:
    values: Set[int] = set()
    for part in token.split(","):
        part = part.strip()
        if not part:
            continue
        if "/" in part:
            base, step_text = part.split("/", 1)
            step = int(step_text)
            candidates = range(min_value, max_value + 1) if base == "*" else _range(base, min_value, max_value)
            values.update(v for v in candidates if (v - min_value) % step == 0)
            continue
        values.update(_range(part, min_value, max_value))
    if not values:
        raise ValueError(f"Invalid cron field: {token}")
    return values


def _range(token: str, min_value: int, max_value: int) -> Iterable[int]:
    if token == "*":
        return range(min_value, max_value + 1)
    if "-" in token:
        start_text, end_text = token.split("-", 1)
        start = int(start_text)
        end = int(end_text)
    else:
        start = end = int(token)
    if start < min_value or end > max_value or end < start:
        raise ValueError(f"Cron value out of range: {token}")
    return range(start, end + 1)


def main() -> int:
    logging.basicConfig(level=logging.INFO)
    parser = argparse.ArgumentParser(description="Enqueue due channel schedules")
    parser.add_argument("--once", action="store_true", help="Run one polling pass and exit")
    parser.add_argument("--poll", type=int, default=60, help="Seconds between polling passes")
    args = parser.parse_args()

    while True:
        result = enqueue_due_schedules()
        if result["enqueued"]:
            logger.info("channel-scheduler: enqueued=%s", result["enqueued"])
        if args.once:
            return 0
        time.sleep(args.poll)


if __name__ == "__main__":
    raise SystemExit(main())
