from decimal import Decimal

from core.countries import normalize_country_code
from core.models import CountryBoundary, FreightRoute, Zone

from .backends import ShapelyBackend

DEFAULT_BACKEND = ShapelyBackend()


def country_code_for_point(lat, lon) -> str | None:
    """Resolve ISO alpha-2 from a lat/lon via country boundaries, then Nominatim."""
    try:
        lat_f = float(lat)
        lon_f = float(lon)
    except (TypeError, ValueError):
        return None

    from shapely.geometry import Point, shape

    point = Point(lon_f, lat_f)
    for boundary in CountryBoundary.objects.all().iterator():
        try:
            geom = shape(boundary.geojson)
        except Exception:
            continue
        if geom.covers(point) or geom.contains(point):
            return normalize_country_code(boundary.country_code)

    from core.geocoding import reverse_geocode_country_code

    return reverse_geocode_country_code(lat, lon)


def freight_route_price(origin, destination) -> tuple[Decimal, str] | None:
    """Directional min freight for origin→destination. None if the pair is missing."""
    origin = normalize_country_code(origin)
    destination = normalize_country_code(destination)
    if not origin or not destination or origin == destination:
        return None
    route = FreightRoute.objects.filter(
        origin_country_code=origin,
        destination_country_code=destination,
    ).first()
    if route is None:
        return None
    return route.min_freight, route.currency


def route_distance_km(route_coords: list[tuple[float, float]]) -> Decimal:
    """Geodesic-approx polyline length in km (EPSG:3857), matching zone intersection."""
    from shapely.geometry import LineString
    from shapely.ops import transform
    import pyproj

    if not route_coords or len(route_coords) < 2:
        return Decimal('0.00')
    project = pyproj.Transformer.from_crs(
        'EPSG:4326', 'EPSG:3857', always_xy=True,
    ).transform
    length_m = transform(project, LineString(route_coords)).length
    return Decimal(str(length_m / 1000)).quantize(Decimal('0.01'))


def estimate_suggested_price(
    pickup_lat,
    pickup_lon,
    delivery_lat,
    delivery_lon,
    route_coords: list[tuple[float, float]],
    backend=None,
) -> dict:
    """
    Same-country: Zone.rate_per_km × intersecting km (raises ValueError if unpriced).
    Cross-border: directional FreightRoute.min_freight, or suggested_price None.
    Currency is always the pickup (origin) country currency.
    """
    from core.currencies import currency_code_for_country

    origin = country_code_for_point(pickup_lat, pickup_lon)
    destination = country_code_for_point(delivery_lat, delivery_lon)
    currency = currency_code_for_country(origin)

    if origin and destination and origin != destination:
        distance_km = route_distance_km(route_coords)
        priced = freight_route_price(origin, destination)
        if priced is None:
            return {
                'suggested_price': None,
                'distance_km': str(distance_km),
                'currency': currency,
                'origin_country_code': origin,
                'destination_country_code': destination,
                'price_breakdown': [],
            }
        min_freight, _route_currency = priced
        return {
            'suggested_price': str(min_freight.quantize(Decimal('0.01'))),
            'distance_km': str(distance_km),
            'currency': currency,
            'origin_country_code': origin,
            'destination_country_code': destination,
            'price_breakdown': [],
        }

    price, breakdown = calculate_suggested_price(route_coords, backend=backend)
    distance_km = sum(
        (Decimal(str(seg.get('km') or 0)) for seg in breakdown),
        Decimal('0'),
    ).quantize(Decimal('0.01'))
    return {
        'suggested_price': str(price),
        'distance_km': str(distance_km),
        'currency': currency,
        'price_breakdown': breakdown,
    }


def calculate_suggested_price(route_coords: list[tuple[float, float]], backend=None) -> tuple[Decimal, list[dict]]:
    """
    route_coords: [(lng, lat), ...] — the actual driven route polyline
    from the routing provider, NOT just pickup/dropoff.

    Returns (total_price, breakdown). Raises ValueError if the route
    doesn't intersect any zone with a configured boundary and rate.
    """
    backend = backend or DEFAULT_BACKEND

    if not route_coords or len(route_coords) < 2:
        raise ValueError('route_coords must contain at least two (lng, lat) points')

    relevant_codes = list(Zone.objects.values_list('country_code', flat=True))
    boundaries = list(CountryBoundary.objects.filter(country_code__in=relevant_codes))

    if not boundaries:
        raise ValueError('No country boundaries configured for any active zone')

    segments = backend.intersect_route_with_boundaries(route_coords, boundaries)

    if not segments:
        raise ValueError('Route does not intersect any configured zone boundary')

    zones_by_code = {
        z.country_code: z
        for z in Zone.objects.filter(country_code__in=[s['country_code'] for s in segments])
    }

    total_price = Decimal('0')
    breakdown = []

    for seg in segments:
        zone = zones_by_code.get(seg['country_code'])
        if not zone or zone.rate_per_km is None:
            continue
        segment_price = seg['km'] * zone.rate_per_km
        total_price += segment_price
        breakdown.append({
            'country_code': seg['country_code'],
            'km': float(seg['km']),
            'rate_per_km': float(zone.rate_per_km),
            'price': float(segment_price),
        })

    if not breakdown:
        raise ValueError('Route intersects zones but none have rate_per_km configured')

    return total_price.quantize(Decimal('0.01')), breakdown
