#!/usr/bin/env python3
"""Quick, reproducible building-segmentation pass over the Santiago OAM mosaic.

The HOT DINOv3/UperNet model was trained on z19 OAM tiles. At Santiago's latitude,
z19 is almost exactly 0.30 m/pixel, so the 5.68 cm orthomosaic is resampled to that
resolution before inference. The output is a screening layer, not upload-ready OSM
geometry.
"""

from __future__ import annotations

import argparse
import json
import math
import time
from pathlib import Path

import cv2
import numpy as np
import onnxruntime as ort
import rasterio
from pyproj import Transformer
from rasterio.enums import Resampling
from rasterio.features import shapes
from rasterio.transform import from_origin
from shapely.geometry import mapping, shape
from shapely.ops import transform as transform_geometry


IMAGE_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)[:, None, None]
IMAGE_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)[:, None, None]


def tile_starts(length: int, tile: int, stride: int) -> list[int]:
    if length <= tile:
        return [0]
    starts = list(range(0, length - tile + 1, stride))
    if starts[-1] != length - tile:
        starts.append(length - tile)
    return starts


def write_interval(starts: list[int], index: int, tile: int, length: int) -> tuple[int, int]:
    """Partition overlapping tiles at overlap midpoints, avoiding visible seams."""
    start = starts[index]
    left = 0 if index == 0 else (starts[index - 1] + tile + start) // 2
    right = length if index == len(starts) - 1 else (start + tile + starts[index + 1]) // 2
    return left, right


def softmax_building_probability(logits: np.ndarray) -> np.ndarray:
    # The released model's class order is building, boundary, background.
    logits = logits[0]
    logits -= logits.max(axis=0, keepdims=True)
    probs = np.exp(logits)
    probs /= probs.sum(axis=0, keepdims=True)
    return probs[0]


def component_summary(probability: np.ndarray, valid: np.ndarray, threshold: float, gsd: float) -> dict:
    mask = ((probability >= round(threshold * 255)) & valid).astype(np.uint8)
    count, _, stats, _ = cv2.connectedComponentsWithStats(mask, connectivity=8)
    areas = stats[1:, cv2.CC_STAT_AREA].astype(np.float64) * gsd * gsd
    return {
        "threshold": threshold,
        "raw_components": int(count - 1),
        "counts_by_minimum_area_m2": {
            str(minimum): int(np.count_nonzero(areas >= minimum))
            for minimum in (5, 8, 12, 15, 25)
        },
        "total_detected_area_m2_at_8m2_minimum": float(areas[areas >= 8].sum()),
    }


def vectorize_default(
    probability: np.ndarray,
    valid: np.ndarray,
    threshold: float,
    minimum_area_m2: float,
    gsd: float,
    transform,
    source_crs,
) -> tuple[dict, np.ndarray]:
    mask = ((probability >= round(threshold * 255)) & valid).astype(np.uint8)
    count, labels, stats, _ = cv2.connectedComponentsWithStats(mask, connectivity=8)
    minimum_pixels = math.ceil(minimum_area_m2 / (gsd * gsd))
    keep = np.zeros(count, dtype=np.uint8)
    keep[1:] = (stats[1:, cv2.CC_STAT_AREA] >= minimum_pixels).astype(np.uint8)
    cleaned = keep[labels]
    del labels

    to_wgs84 = Transformer.from_crs(source_crs, 4326, always_xy=True).transform
    features = []
    for geometry, value in shapes(cleaned, mask=cleaned.astype(bool), transform=transform):
        if value != 1:
            continue
        projected = shape(geometry).buffer(0)
        if projected.is_empty:
            continue
        area_m2 = projected.area
        geographic = transform_geometry(to_wgs84, projected)
        features.append(
            {
                "type": "Feature",
                "properties": {
                    "screening_only": True,
                    "model": "hotosm/dinov3s-buildings",
                    "probability_threshold": threshold,
                    "minimum_area_m2": minimum_area_m2,
                    "detected_area_m2": area_m2,
                },
                "geometry": mapping(geographic),
            }
        )
    return {"type": "FeatureCollection", "features": features}, cleaned


def run(args: argparse.Namespace) -> None:
    started = time.time()
    output_dir = Path(args.output_dir)
    output_dir.mkdir(parents=True, exist_ok=True)

    with rasterio.open(args.imagery) as source:
        width_m = source.bounds.right - source.bounds.left
        height_m = source.bounds.top - source.bounds.bottom
        width = math.ceil(width_m / args.gsd)
        height = math.ceil(height_m / args.gsd)
        target_transform = from_origin(source.bounds.left, source.bounds.top, args.gsd, args.gsd)
        rgb = source.read(
            out_shape=(3, height, width),
            resampling=Resampling.bilinear,
        )
        source_crs = source.crs

    valid = np.any(rgb > 0, axis=0)
    probability = np.zeros((height, width), dtype=np.uint8)
    x_starts = tile_starts(width, args.tile_size, args.stride)
    y_starts = tile_starts(height, args.tile_size, args.stride)

    providers = ["CoreMLExecutionProvider", "CPUExecutionProvider"]
    session = ort.InferenceSession(args.model, providers=providers)
    processed = 0
    skipped = 0
    total_tiles = len(x_starts) * len(y_starts)

    for yi, y in enumerate(y_starts):
        global_top, global_bottom = write_interval(y_starts, yi, args.tile_size, height)
        local_top, local_bottom = global_top - y, global_bottom - y
        for xi, x in enumerate(x_starts):
            tile = rgb[:, y : y + args.tile_size, x : x + args.tile_size]
            if np.count_nonzero(np.any(tile > 0, axis=0)) < args.tile_size:
                skipped += 1
                continue

            tensor = tile.astype(np.float32) / 255.0
            tensor = ((tensor - IMAGE_MEAN) / IMAGE_STD)[None]
            logits = session.run(None, {"image": tensor})[0]
            tile_probability = softmax_building_probability(logits)

            global_left, global_right = write_interval(x_starts, xi, args.tile_size, width)
            local_left, local_right = global_left - x, global_right - x
            probability[global_top:global_bottom, global_left:global_right] = np.rint(
                tile_probability[local_top:local_bottom, local_left:local_right] * 255
            ).astype(np.uint8)
            processed += 1
            if processed % 100 == 0:
                print(
                    f"processed={processed} skipped={skipped} total_grid_tiles={total_tiles} "
                    f"elapsed_s={time.time() - started:.1f}",
                    flush=True,
                )

    del rgb

    sensitivity = [
        component_summary(probability, valid, threshold, args.gsd)
        for threshold in (0.35, args.threshold, 0.50, 0.55)
    ]
    feature_collection, cleaned = vectorize_default(
        probability,
        valid,
        args.threshold,
        args.minimum_area,
        args.gsd,
        target_transform,
        source_crs,
    )

    (output_dir / "ai_buildings_screening.geojson").write_text(json.dumps(feature_collection))
    np.save(output_dir / "ai_probability_uint8.npy", probability)
    np.save(output_dir / "ai_cleaned_mask_uint8.npy", cleaned)

    summary = {
        "imagery": str(args.imagery),
        "model": str(args.model),
        "model_huggingface_id": "hotosm/dinov3s-buildings",
        "target_gsd_m": args.gsd,
        "target_grid": {"width": width, "height": height},
        "valid_pixel_fraction": float(valid.mean()),
        "tile_size": args.tile_size,
        "stride": args.stride,
        "tiles_processed": processed,
        "tiles_skipped_as_empty": skipped,
        "default_threshold": args.threshold,
        "default_minimum_area_m2": args.minimum_area,
        "default_vector_count": len(feature_collection["features"]),
        "sensitivity": sensitivity,
        "runtime_seconds": time.time() - started,
        "providers": session.get_providers(),
        "caveat": "AI screening only; every polygon requires imagery review and clean manual geometry before OSM use.",
    }
    (output_dir / "ai_summary.json").write_text(json.dumps(summary, indent=2))
    print(json.dumps(summary, indent=2), flush=True)


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--imagery", required=True)
    parser.add_argument("--model", required=True)
    parser.add_argument("--output-dir", required=True)
    parser.add_argument("--gsd", type=float, default=0.30)
    parser.add_argument("--tile-size", type=int, default=256)
    parser.add_argument("--stride", type=int, default=192)
    parser.add_argument("--threshold", type=float, default=0.4371)
    parser.add_argument("--minimum-area", type=float, default=8.0)
    run(parser.parse_args())


if __name__ == "__main__":
    main()
