# -*- coding: utf-8 -*-
"""m3_umum.py: fungsi bantu bersama untuk skrip Seri M3. Penulis: Badar Mubarok Yogaswara
Letakkan berkas ini satu folder dengan skrip lain. Ubah PAKET dan HASIL sesuai komputer Anda,
atau isi variabel lingkungan M3_PAKET dan M3_HASIL."""
import csv
import datetime as dt
import os

import numpy as np
from osgeo import gdal

gdal.UseExceptions()

PAKET = os.environ.get("M3_PAKET", r"D:/Latihan_M3/paket-m3")     # folder paket data sintetis
HASIL = os.environ.get("M3_HASIL", r"D:/Latihan_M3/hasil")        # folder keluaran Anda
os.makedirs(HASIL, exist_ok=True)

BAND = {"B02": 0, "B03": 1, "B04": 2, "B08": 3, "B11": 4, "B12": 5}   # urutan band di berkas *_refl.tif
BAIK = (4, 5, 6)                                                       # kode SCL: vegetasi, bukan vegetasi, air


def daftar_citra():
    """Baca Daftar_Citra.csv: kembalikan daftar (tanggal, jalur_refl, jalur_scl)."""
    hasil = []
    with open(os.path.join(PAKET, "citra", "Daftar_Citra.csv"), encoding="utf8") as f:
        for r in csv.DictReader(f):
            hasil.append((dt.date.fromisoformat(r["tanggal"]),
                          os.path.join(PAKET, "citra", r["berkas_refl"]),
                          os.path.join(PAKET, "citra", r["berkas_scl"])))
    return hasil


def baca_deret():
    """Kembalikan tanggal (daftar), refl (T,6,H,W), scl (T,H,W), geotransform, proyeksi."""
    tgl, refl, scl = [], [], []
    for d, pr, ps in daftar_citra():
        a = gdal.Open(pr)
        refl.append(a.ReadAsArray().astype("float32"))
        scl.append(gdal.Open(ps).ReadAsArray())
        gt, prj = a.GetGeoTransform(), a.GetProjection()
        tgl.append(d)
    return tgl, np.array(refl), np.array(scl), gt, prj


def nilai_ndvi(refl):
    nir, red = refl[:, BAND["B08"]], refl[:, BAND["B04"]]
    return (nir - red) / (nir + red + 1e-9)


def nilai_nbr(refl):
    nir, swir = refl[:, BAND["B08"]], refl[:, BAND["B12"]]
    return (nir - swir) / (nir + swir + 1e-9)


def bersih(scl):
    """True di piksel yang boleh dipakai (bukan awan, bayangan, atau cirrus)."""
    return np.isin(scl, BAIK)


def tulis_tif(path, arr, gt, prj, nodata=None, deskripsi=None, tipe=None):
    """Tulis larik (H,W) atau (B,H,W) ke GeoTIFF."""
    arr = np.asarray(arr)
    if arr.ndim == 2:
        arr = arr[None]
    if tipe is None:
        tipe = gdal.GDT_Float32 if arr.dtype.kind == "f" else gdal.GDT_Byte
    os.makedirs(os.path.dirname(path), exist_ok=True)
    ds = gdal.GetDriverByName("GTiff").Create(path, arr.shape[2], arr.shape[1], arr.shape[0], tipe, ["COMPRESS=LZW"])
    ds.SetGeoTransform(gt)
    ds.SetProjection(prj)
    for i in range(arr.shape[0]):
        b = ds.GetRasterBand(i + 1)
        b.WriteArray(arr[i])
        if nodata is not None:
            b.SetNoDataValue(nodata)
        if deskripsi:
            b.SetDescription(deskripsi[i])
    ds = None


def baca_tif(path):
    ds = gdal.Open(path)
    return ds.ReadAsArray(), ds.GetGeoTransform(), ds.GetProjection()


def piksel_ke_xy(gt, baris, kolom):
    return gt[0] + (kolom + 0.5) * gt[1], gt[3] + (baris + 0.5) * gt[5]


def xy_ke_piksel(gt, x, y):
    return int((gt[3] - y) / -gt[5]), int((x - gt[0]) / gt[1])


def isi_waktu(x, baik, hari):
    """Interpolasi linear sepanjang waktu. x dan baik berbentuk (T,H,W); hari = hari ke- tiap tanggal."""
    T, H, W = x.shape
    hasil = np.full(x.shape, np.nan, dtype="float64")
    for i in range(H):
        for j in range(W):
            ok = baik[:, i, j]
            if ok.any():
                hasil[:, i, j] = np.interp(hari, hari[ok], x[ok, i, j])
    return hasil
