"""Hae Luken MVMI 2023 -teemat (keskipituus, latvuspeitto) kohdealueille
/vsicurl-leikkauksina ja siivoa erikoisarvot.

MVMI-koodit (README-2023.txt):
  32767 = ei metsämaata / aineiston ulkopuolella  -> nodata (jo metatiedoissa)
  32766 = laskenta ei onnistunut                  -> muunnetaan 32767:ksi,
          koska muuten se vuotaisi resamplaukseen arvona 3276,6 dm.

Ajo VM:llä: python3 mvmi_fetch.py
"""
from pathlib import Path

import numpy as np
from osgeo import gdal, osr

gdal.UseExceptions()

BASE = "/vsicurl/https://www.nic.funet.fi/index/geodata/luke/vmi/2023/"
THEMES = {
    "keskipituus": "keskipituus_vmi1x_1923.tif",
    "latvuspeitto": "latvuspeitto_vmi1x_1923.tif",
}
OUT = Path.home() / "oh9ab/data/mvmi"

# Samat kohdepisteet kuin DEM-lehdillä (select_sheets.py), puolileveys 7 km
# jotta puusto kattaa DEM-alueen (6 km) reunoineen.
TARGETS = {
    "kemi": (24.5637, 65.7364),
    "simo": (25.0672, 65.6622),
    "rovaniemi": (25.7245, 66.4977),
}
HALF = 7000.0
NODATA = 32767

wgs = osr.SpatialReference(); wgs.ImportFromEPSG(4326)
tm35 = osr.SpatialReference(); tm35.ImportFromEPSG(3067)
wgs.SetAxisMappingStrategy(osr.OAMS_TRADITIONAL_GIS_ORDER)
tm35.SetAxisMappingStrategy(osr.OAMS_TRADITIONAL_GIS_ORDER)
tr = osr.CoordinateTransformation(wgs, tm35)

OUT.mkdir(parents=True, exist_ok=True)
for theme, fname in THEMES.items():
    for area, (lon, lat) in TARGETS.items():
        e, n, _ = tr.TransformPoint(lon, lat)
        dst = OUT / ("%s_%s.tif" % (theme, area))
        print("leikataan %-30s E=%.0f N=%.0f" % (dst.name, e, n))
        # projWin = (ulx, uly, lrx, lry); GDAL tasaa 16 m pikseligridiin.
        gdal.Translate(
            str(dst), BASE + fname,
            projWin=(e - HALF, n + HALF, e + HALF, n - HALF),
            creationOptions=["COMPRESS=DEFLATE"],
        )
        # Siivous: 32766 -> 32767 (nodata).
        ds = gdal.Open(str(dst), gdal.GA_Update)
        band = ds.GetRasterBand(1)
        arr = band.ReadAsArray()
        n_bad = int((arr == 32766).sum())
        if n_bad:
            arr[arr == 32766] = NODATA
            band.WriteArray(arr)
        band.SetNoDataValue(NODATA)
        ds.FlushCache()
        valid = arr[arr != NODATA]
        print("  %dx%d px, 32766-arvoja %d, validit %.0f..%.0f (mediaani %.0f)"
              % (arr.shape[1], arr.shape[0], n_bad,
                 valid.min() if valid.size else -1,
                 valid.max() if valid.size else -1,
                 np.median(valid) if valid.size else -1))
        ds = None

print("Valmis ->", OUT)
