# QGIS 3.34.4 (PyQGIS, run with the system Python 3.12), headless.
# Reads richness-10km.csv and richness-1km.csv written by the post's R chunk,
# writes the qgis-*.csv files that the post's R chunks read.
import os, csv
os.environ["QT_QPA_PLATFORM"] = "offscreen"
from qgis.core import (Qgis, QgsApplication, QgsVectorLayer, QgsFeature, QgsGeometry,
                       QgsPointXY, QgsClassificationEqualInterval, QgsClassificationQuantile,
                       QgsClassificationJenks, QgsClassificationPrettyBreaks,
                       QgsGraduatedSymbolRenderer)
app = QgsApplication([], False); app.initQgis()

def read_grid(path):
    with open(path) as f:
        return list(csv.DictReader(f))

def point_layer(rows, field, n, shift=0.0):
    # one point per grid cell, first n rows of the file, value = field + shift
    vl = QgsVectorLayer("Point?crs=EPSG:3844&field=v:double", "cells", "memory")
    feats = []
    for r in rows[:n]:
        f = QgsFeature(vl.fields())
        f.setGeometry(QgsGeometry.fromPointXY(QgsPointXY(float(r["col"]), float(r["row"]))))
        f.setAttributes([float(r[field]) + shift])
        feats.append(f)
    vl.dataProvider().addFeatures(feats)
    return vl

def bounds(c):
    # classification ranges and renderer ranges name their bounds differently
    if hasattr(c, "lowerBound"):
        return c.lowerBound(), c.upperBound()
    return c.lowerValue(), c.upperValue()

def graduated_jenks(layer, k=5):
    # what Layer Properties > Symbology > Graduated > Classify does in Natural Breaks mode
    rend = QgsGraduatedSymbolRenderer("v")
    rend.setClassificationMethod(QgsClassificationJenks())
    rend.updateClasses(layer, k)
    ranges = rend.ranges()
    keys = [r.uuid() for r in ranges]
    tally = [0] * len(keys)            # features per legend entry, by the renderer's lookup
    for feat in layer.getFeatures():
        k = rend.legendKeyForValue(feat["v"])
        if k in keys:
            tally[keys.index(k)] += 1
    return ranges, tally

def write(path, rows):
    with open(path, "w", newline="") as f:
        w = csv.DictWriter(f, fieldnames=list(rows[0].keys())); w.writeheader(); w.writerows(rows)

# 1. the 10 km layer (600 cells), four classification modes
g10 = read_grid("richness-10km.csv")
lay10 = point_layer(g10, "richness", len(g10))
out = []
for mode, cls in [("equal", QgsClassificationEqualInterval), ("quantile", QgsClassificationQuantile),
                  ("jenks", QgsClassificationJenks), ("pretty", QgsClassificationPrettyBreaks)]:
    for i, c in enumerate(cls().classes(lay10, "v", 5)):
        out.append(dict(mode=mode, klass=i + 1, lower=repr(bounds(c)[0]), upper=repr(bounds(c)[1])))
write("qgis-10km-breaks.csv", out)

# 2. Natural Breaks on the first n cells of the 1 km layer, and on shifted copies
g1 = read_grid("richness-1km.csv")
runs = [("richness", n, 0.0) for n in
        [1000, 2000, 3000, 3001, 5000, 10000, 20000, 30000, 30009, 30010, 30020, 30050,
         30100, 30200, 30500, 31000, 35000, 40000, 50000, 60000]]
runs += [("richness", 60000, s) for s in [2.5, 5.0, 10.0, 20.0, 40.0]]
runs += [("change", n, 0.0) for n in [20000, 60000]]
out = []
for field, n, shift in runs:
    ranges, tally = graduated_jenks(point_layer(g1, field, n, shift))
    for i, r in enumerate(ranges):
        out.append(dict(field=field, n=n, shift=shift, klass=i + 1, lower=repr(bounds(r)[0]),
                        upper=repr(bounds(r)[1]), features=tally[i]))
    print(field, n, shift, [round(bounds(r)[1], 3) for r in ranges], tally, flush=True)
write("qgis-jenks-runs.csv", out)

# 3. what QGIS read, so the R side can check it is looking at the same file
write("qgis-input-check.csv", [dict(
    qgis_version=Qgis.QGIS_VERSION, cells_10km=len(g10), cells_1km=len(g1),
    sum_10km=repr(sum(float(r["richness"]) for r in g10)),
    sum_richness=repr(sum(float(r["richness"]) for r in g1)),
    sum_change=repr(sum(float(r["change"]) for r in g1)),
    jenks_code_complexity=QgsClassificationJenks().codeComplexity())])

# 4. seven values with an obvious answer, classed into two and three classes
tiny = [1, 2, 3, 50, 100, 101, 102]
tiny_rows = [dict(col=i + 1, row=1, v=v) for i, v in enumerate(tiny)]
out = []
for k in [2, 3]:
    ranges, tally = graduated_jenks(point_layer(tiny_rows, "v", len(tiny)), k)
    for i, r in enumerate(ranges):
        out.append(dict(k=k, klass=i + 1, lower=repr(bounds(r)[0]), upper=repr(bounds(r)[1]),
                        features=tally[i]))
    print("tiny", k, [bounds(r) for r in ranges], tally, flush=True)
write("qgis-tiny.csv", out)
app.exitQgis()
