"""Poimi dem2m_direct.vrt-indeksistä karttalehdet, jotka osuvat
annettujen pisteiden ympärysalueisiin, ja lataa ne Funetista.

Ajo VM:llä: python3 select_sheets.py
"""
import re
import subprocess
import sys
from pathlib import Path

from osgeo import osr

VRT = Path.home() / "oh9ab/data/dem2m_direct.vrt"
OUT = Path.home() / "oh9ab/data/dem"
BASE = "https://www.nic.funet.fi/index/geodata/mml/dem2m/"

# VRT:n geotransform: origo (62000, 7782000), pikseli 2 m.
GT_X0, GT_Y0, GT_RES = 62000.0, 7782000.0, 2.0

# Kohdepisteet (lon, lat) ja laatikon puolileveys metreinä.
TARGETS = {
    "kemi": (24.5637, 65.7364),
    "simo": (25.0672, 65.6622),
    "rovaniemi": (25.7245, 66.4977),
}
HALF = 6000.0

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)

boxes = {}
for name, (lon, lat) in TARGETS.items():
    e, n, _ = tr.TransformPoint(lon, lat)
    boxes[name] = (e - HALF, n - HALF, e + HALF, n + HALF)
    print("%-10s E=%.0f N=%.0f" % (name, e, n))

# Parsi VRT: SourceFilename + DstRect (pikseleinä VRT:n gridissä).
text = VRT.read_text()
pat = re.compile(
    r'<SourceFilename relativeToVRT="1">([^<]+)</SourceFilename>.*?'
    r'<DstRect xOff="([\d.]+)" yOff="([\d.]+)" xSize="([\d.]+)" ySize="([\d.]+)"',
    re.S)

hits = {}
for m in pat.finditer(text):
    rel, xoff, yoff, xs, ys = m.group(1), *map(float, m.group(2, 3, 4, 5))
    minx = GT_X0 + xoff * GT_RES
    maxy = GT_Y0 - yoff * GT_RES
    maxx = minx + xs * GT_RES
    miny = maxy - ys * GT_RES
    for name, (bx0, by0, bx1, by1) in boxes.items():
        if minx < bx1 and maxx > bx0 and miny < by1 and maxy > by0:
            hits.setdefault(rel, []).append(name)

print("Lehtiä valittu: %d" % len(hits))
for rel in sorted(hits):
    print("  %-40s %s" % (rel, ",".join(hits[rel])))

if "--dry-run" in sys.argv:
    sys.exit(0)

OUT.mkdir(parents=True, exist_ok=True)
for rel in sorted(hits):
    dst = OUT / Path(rel).name
    if dst.exists() and dst.stat().st_size > 0:
        print("ohitetaan (on jo):", dst.name)
        continue
    url = BASE + rel
    print("ladataan:", url)
    subprocess.run(["curl", "-sf", "--max-time", "300", "-o", str(dst), url],
                   check=True)
print("Valmis. Ladattu ->", OUT)
