from datetime import date

from django.core.management.base import BaseCommand
from django.db import transaction, connection

from api.models import TtStaff, TtLocation, TtStudentSet, TtActivityTemplate, TtActivity, TtModule, TtAcademicTerm, \
    TtActivityTemplateWeek
from api.utils import convert_to_redis_week, bulk_sync_to_redis
from backend.redis_client import redis_client

class Command(BaseCommand):
    def handle(self, *args, **options):
        if date.today() > date(2026, 7, 6):
            return  # do nothing
        tables_with_cascade = [
            "tt_department",
            "tt_zone",
            "tt_activity_type",
            "tt_activity",
            "tt_academic_term",
            "tt_staff",
            "tt_location",
            "tt_suitability"
        ]
        tables = [
            "emx_tt_activity_final",
            "emx_tt_activity_history",
            "emx_tt_activity_last",
            "emx_tt_activity_static",
            "emx_tt_activity_date_final",
            "emx_tt_activity_date_history",
            "emx_tt_activity_date_last",
            "emx_tt_activity_date_static",
            "emx_tt_activity_jta_variant_final",
            "emx_tt_activity_jta_variant_history",
            "emx_tt_activity_jta_variant_last",
            "emx_tt_activity_jta_variant_static",
            "emx_tt_activity_location_final",
            "emx_tt_activity_location_history",
            "emx_tt_activity_location_last",
            "emx_tt_activity_location_static",
            "emx_tt_activity_pos_final",
            "emx_tt_activity_pos_history",
            "emx_tt_activity_pos_last",
            "emx_tt_activity_pos_static",
            "emx_tt_activity_staff_final",
            "emx_tt_activity_staff_history",
            "emx_tt_activity_staff_last",
            "emx_tt_activity_staff_static",
        ]

        with transaction.atomic():
            with connection.cursor() as cursor:
                if tables_with_cascade:
                    cascade_str = ", ".join(tables_with_cascade)
                    cascade_query = (
                        f"TRUNCATE TABLE {cascade_str} RESTART IDENTITY CASCADE;"
                    )
                    cursor.execute(cascade_query)

                if tables:
                    standard_str = ", ".join(tables)
                    standard_query = (
                        f"TRUNCATE TABLE {standard_str} RESTART IDENTITY RESTRICT;"
                    )
                    cursor.execute(standard_query)