from django.db import migrations, models

from apps.core.countries import ALL_COUNTRY_MAP, to_country_name


def _forwards(apps, schema_editor):
    """Convert stored ISO alpha-2 geo codes (UG) to full country names (Uganda)."""
    AccessEvent = apps.get_model("repository_analytics", "AccessEvent")
    DownloadEvent = apps.get_model("repository_analytics", "DownloadEvent")
    UsageRollup = apps.get_model("repository_analytics", "UsageRollup")

    for Model in (AccessEvent, DownloadEvent):
        for obj in Model.objects.exclude(country="").only("id", "country").iterator():
            name = to_country_name(obj.country)
            if name != obj.country:
                Model.objects.filter(pk=obj.pk).update(country=name)

    # Rollup rows key the COUNTRY dimension on the raw geo value — convert those
    # keys too so weekly aggregation stays consistent post-migration.
    for obj in (
        UsageRollup.objects.filter(dimension="country")
        .exclude(key="")
        .only("id", "key")
        .iterator()
    ):
        name = to_country_name(obj.key)
        if name != obj.key:
            UsageRollup.objects.filter(pk=obj.pk).update(key=name)


def _backwards(apps, schema_editor):
    """Best-effort reverse: full names back to ISO alpha-2 codes."""
    AccessEvent = apps.get_model("repository_analytics", "AccessEvent")
    DownloadEvent = apps.get_model("repository_analytics", "DownloadEvent")
    UsageRollup = apps.get_model("repository_analytics", "UsageRollup")
    name_to_code = {name: code for code, name in ALL_COUNTRY_MAP.items()}

    for Model in (AccessEvent, DownloadEvent):
        for obj in Model.objects.exclude(country="").only("id", "country").iterator():
            code = name_to_code.get(obj.country)
            if code is not None and code != obj.country:
                Model.objects.filter(pk=obj.pk).update(country=code)

    for obj in (
        UsageRollup.objects.filter(dimension="country")
        .exclude(key="")
        .only("id", "key")
        .iterator()
    ):
        code = name_to_code.get(obj.key)
        if code is not None and code != obj.key:
            UsageRollup.objects.filter(pk=obj.pk).update(key=code)


class Migration(migrations.Migration):

    dependencies = [
        ("repository_analytics", "0003_initial"),
    ]

    operations = [
        migrations.AlterField(
            model_name="accessevent",
            name="country",
            field=models.CharField(blank=True, max_length=100),
        ),
        migrations.AlterField(
            model_name="downloadevent",
            name="country",
            field=models.CharField(blank=True, max_length=100),
        ),
        migrations.RunPython(_forwards, _backwards),
    ]
