from datetime import timedelta

from celery import Celery
from django.core.exceptions import ValidationError
from django.test import TransactionTestCase

from task_scheduler.models import PeriodicTask
from task_scheduler.scheduler import DatabaseScheduler


class DatabaseSchedulerTests(TransactionTestCase):
    def setUp(self):
        PeriodicTask.objects.all().delete()
        self.celery_app = Celery("scheduler-tests", broker="memory://")
        self.celery_app.conf.update(
            beat_max_loop_interval=5,
            timezone="UTC",
        )

    def test_loads_only_enabled_tasks_from_database(self):
        enabled = self._create_task(name="enabled", interval_seconds=60)
        self._create_task(
            name="disabled",
            interval_seconds=120,
            enabled=False,
        )

        scheduler = DatabaseScheduler(app=self.celery_app)

        self.assertEqual(set(scheduler.schedule), {"enabled"})
        entry = scheduler.schedule["enabled"]
        self.assertEqual(entry.task, enabled.task)
        self.assertEqual(entry.schedule.run_every, timedelta(seconds=60))
        self.assertEqual(entry.options, {"queue": "test"})

    def test_sync_persists_last_run_state(self):
        task = self._create_task(name="sync-state", interval_seconds=60)
        scheduler = DatabaseScheduler(app=self.celery_app)

        new_entry = scheduler.reserve(scheduler.schedule[task.name])
        scheduler.sync()

        task.refresh_from_db()
        self.assertEqual(task.lastRunAt, new_entry.last_run_at)
        self.assertEqual(task.totalRunCount, 1)

    def test_refresh_applies_database_changes(self):
        task = self._create_task(name="refresh", interval_seconds=60)
        scheduler = DatabaseScheduler(app=self.celery_app)
        task.intervalSeconds = 120
        task.options = {"queue": "updated"}
        task.save()

        scheduler._refresh_schedule()

        entry = scheduler.schedule[task.name]
        self.assertEqual(entry.schedule.run_every, timedelta(seconds=120))
        self.assertEqual(entry.options, {"queue": "updated"})

    def test_refresh_removes_disabled_task(self):
        task = self._create_task(name="disable", interval_seconds=60)
        scheduler = DatabaseScheduler(app=self.celery_app)
        task.enabled = False
        task.save()

        scheduler._refresh_schedule()

        self.assertNotIn(task.name, scheduler.schedule)

    def test_model_rejects_invalid_json_shapes(self):
        task = self._create_task(name="invalid", interval_seconds=60)
        task.args = {}
        task.kwargs = []
        task.options = []

        with self.assertRaises(ValidationError) as context:
            task.full_clean()

        self.assertEqual(
            set(context.exception.message_dict),
            {"args", "kwargs", "options"},
        )

    @staticmethod
    def _create_task(name, interval_seconds, enabled=True):
        return PeriodicTask.objects.create(
            name=name,
            task="test.task",
            intervalSeconds=interval_seconds,
            options={"queue": "test"},
            enabled=enabled,
        )
