"""
Import country polygon GeoJSON into CountryBoundary.

Tuned for this repo's ``countries.geojson`` FeatureCollection, which uses:
  - properties.ISO3166-1-Alpha-2  (ISO alpha-2 code)
  - properties.name               (country name)
  - geometry type Polygon or MultiPolygon

By default imports **all** countries with a valid ISO alpha-2 code.
Use ``--codes`` to limit to a subset.

Large polygons are simplified before save so MySQL does not drop the connection
on remote max_allowed_packet / packet timeouts.

Example:
  python manage.py import_country_boundaries
  python manage.py import_country_boundaries --codes AE,OM,SA
"""

from __future__ import annotations

import json
import time
from pathlib import Path

from django.conf import settings
from django.core.management.base import BaseCommand, CommandError
from django.db import OperationalError, close_old_connections, connection

from core.models import CountryBoundary

# Prefer keys used by countries.geojson; keep older Natural Earth aliases as fallback.
CODE_KEYS = (
    'ISO3166-1-Alpha-2',
    'iso3166-1-alpha-2',
    'ISO_A2',
    'iso_a2',
    'ISO_A2_EH',
    'WB_A2',
    'cca2',
)
NAME_KEYS = (
    'name',
    'NAME',
    'NAME_EN',
    'ADMIN',
    'NAME_LONG',
)

DEFAULT_FILE = Path(settings.BASE_DIR) / 'countries.geojson'
DEFAULT_SIMPLIFY = 0.01
MAX_RETRIES = 3


def _feature_country_code(props: dict) -> str:
    for key in CODE_KEYS:
        raw = props.get(key)
        if raw is None:
            continue
        code = str(raw).strip().upper()
        if len(code) == 2 and code.isalpha() and code not in ('-1', '-99'):
            return code
    return ''


def _feature_country_name(props: dict, code: str) -> str:
    for key in NAME_KEYS:
        raw = props.get(key)
        if raw:
            return str(raw).strip()[:100]
    return code


def _simplify_geometry(geometry: dict, tolerance: float) -> dict:
    """Reduce vertex count for MySQL-friendly JSON storage (country-scale pricing)."""
    if tolerance <= 0:
        return geometry
    from shapely.geometry import mapping, shape

    geom = shape(geometry)
    if geom.is_empty:
        raise ValueError('empty geometry')
    simplified = geom.simplify(tolerance, preserve_topology=True)
    if simplified.is_empty:
        simplified = geom
    out = mapping(simplified)
    if out.get('type') not in ('Polygon', 'MultiPolygon'):
        return geometry
    return out


def _upsert_boundary(code: str, name: str, geometry: dict) -> bool:
    """Create or update without SELECT-ing existing geojson."""
    from django.db import IntegrityError

    existing_id = (
        CountryBoundary.objects.filter(country_code=code)
        .values_list('id', flat=True)
        .first()
    )
    if existing_id:
        CountryBoundary.objects.filter(pk=existing_id).update(
            country_name=name,
            geojson=geometry,
        )
        return False
    try:
        CountryBoundary.objects.create(
            country_code=code,
            country_name=name,
            geojson=geometry,
        )
        return True
    except IntegrityError:
        # Concurrent import or retry after a write that actually succeeded.
        CountryBoundary.objects.filter(country_code=code).update(
            country_name=name,
            geojson=geometry,
        )
        return False


def _upsert_with_retry(code: str, name: str, geometry: dict, stdout, style) -> bool:
    last_exc = None
    for attempt in range(1, MAX_RETRIES + 1):
        close_old_connections()
        try:
            with connection.cursor() as cursor:
                cursor.execute('SET SESSION net_read_timeout=600')
                cursor.execute('SET SESSION net_write_timeout=600')
                cursor.execute('SET SESSION wait_timeout=600')
            return _upsert_boundary(code, name, geometry)
        except OperationalError as exc:
            last_exc = exc
            close_old_connections()
            stdout.write(style.WARNING(
                f'{code}: MySQL error on attempt {attempt}/{MAX_RETRIES}: {exc}'
            ))
            time.sleep(attempt * 2)
    raise last_exc


class Command(BaseCommand):
    help = (
        'Import all Polygon/MultiPolygon country boundaries from countries.geojson '
        '(ISO3166-1-Alpha-2 + name). Optionally filter with --codes.'
    )

    def add_arguments(self, parser):
        parser.add_argument(
            '--file',
            default=str(DEFAULT_FILE),
            help=f'Path to GeoJSON FeatureCollection (default: {DEFAULT_FILE}).',
        )
        parser.add_argument(
            '--codes',
            default='',
            help=(
                'Optional comma-separated ISO alpha-2 filter. '
                'Omit to import every country with a valid alpha-2 code.'
            ),
        )
        parser.add_argument(
            '--simplify',
            type=float,
            default=DEFAULT_SIMPLIFY,
            help=(
                'Shapely simplify tolerance in degrees (default: 0.01). '
                'Use 0 to store full-resolution polygons.'
            ),
        )

    def handle(self, *args, **options):
        path = Path(options['file'])
        if not path.is_file():
            raise CommandError(f'File not found: {path}')

        codes_raw = (options.get('codes') or '').strip()
        wanted = None
        if codes_raw:
            wanted = {
                c.strip().upper()
                for c in codes_raw.split(',')
                if c.strip()
            }
            if not wanted:
                raise CommandError('No country codes provided.')

        simplify = float(options['simplify'])

        try:
            payload = json.loads(path.read_text(encoding='utf-8'))
        except json.JSONDecodeError as exc:
            raise CommandError(f'Invalid JSON: {exc}') from exc

        if payload.get('type') != 'FeatureCollection':
            raise CommandError('Expected a GeoJSON FeatureCollection (type=FeatureCollection).')

        features = payload.get('features')
        if not isinstance(features, list):
            raise CommandError('Expected a GeoJSON FeatureCollection with a features array.')

        created = updated = skipped = no_code = 0
        matched_codes = set()
        scope = f'codes={",".join(sorted(wanted))}' if wanted else 'all countries'

        self.stdout.write(f'Importing {scope} from {path.name} ({len(features)} features)...')

        for feature in features:
            if not isinstance(feature, dict):
                skipped += 1
                continue
            props = feature.get('properties') or {}
            geometry = feature.get('geometry') or {}
            code = _feature_country_code(props)
            if not code:
                no_code += 1
                continue
            if wanted is not None and code not in wanted:
                continue
            # Same code can appear once; skip duplicates after first save.
            if code in matched_codes:
                skipped += 1
                continue

            geom_type = geometry.get('type')
            if geom_type not in ('Polygon', 'MultiPolygon'):
                self.stderr.write(
                    self.style.WARNING(
                        f'Skipping {code}: geometry type {geom_type!r} (need Polygon/MultiPolygon).'
                    )
                )
                skipped += 1
                continue

            name = _feature_country_name(props, code)
            try:
                geometry = _simplify_geometry(geometry, simplify)
            except Exception as exc:
                self.stderr.write(self.style.WARNING(f'Skipping {code}: simplify failed ({exc})'))
                skipped += 1
                continue

            size_kb = len(json.dumps(geometry, separators=(',', ':'))) / 1024
            try:
                was_created = _upsert_with_retry(
                    code, name, geometry, self.stdout, self.style,
                )
            except OperationalError as exc:
                raise CommandError(
                    f'Failed to save {code} after {MAX_RETRIES} retries '
                    f'({size_kb:.1f} KB geometry): {exc}'
                ) from exc

            matched_codes.add(code)
            if was_created:
                created += 1
                self.stdout.write(f'Created {code} ({name}) [{size_kb:.1f} KB]')
            else:
                updated += 1
                self.stdout.write(f'Updated {code} ({name}) [{size_kb:.1f} KB]')

        missing = sorted(wanted - matched_codes) if wanted else []
        self.stdout.write(self.style.SUCCESS(
            f'Import complete: created={created} updated={updated} skipped={skipped} '
            f'no_valid_code={no_code} matched={len(matched_codes)} from {path.name}'
        ))
        if missing:
            self.stdout.write(self.style.WARNING(f'No features found for: {", ".join(missing)}'))
