#!/home/sam/dev/perso/flsun/klippy-env/bin/python
"""Corrige la bed mesh a partir de l'epaisseur mesuree de 13 carres imprimes.

    ./fix-mesh.py            # saisie guidee, puis reecriture de printer.cfg + restart
    ./fix-mesh.py --sans-saisie   # sans re-saisir : mesures-mesh.json tel quel
    ./fix-mesh.py --sans-pointage # imprimante eteinte : position en clair (rangee, rang)

Principe : gcodes/first_layer_13_pla.gcode imprime un carre de 20 mm, une
couche de 0.2, exactement sur chacun des 13 points de la mesh. On mesure
l'epaisseur de chaque carre au pied a coulisse (centieme). Un carre plus
epais que la moyenne = la buse etait plus haute a cet endroit = le plateau y
est plus BAS que ce que la mesh croit -> on abaisse le point de la mesh de
l'ecart, et inversement. La moyenne est conservee : le zero global (offset
0.07 de START_PRINT) n'est pas touche.

Ordre de saisie : de l'AVANT vers l'ARRIERE (Y croissant, +Y = vers la tour C),
et de GAUCHE a DROITE (X croissant). Chaque carre est nomme par ses
coordonnees, comme dans la mesh.

Ecrit le bloc "#*# points =" de printer.cfg (sauvegarde prealable), puis
redemarre klippy. Refuse si une impression est en cours.
"""
import json
import os
import re
import shutil
import socket
import subprocess
import sys
import time

HERE = os.path.dirname(os.path.abspath(__file__))
CFG = os.path.join(HERE, "printer.cfg")
FICHIER = os.path.join(HERE, "mesures-mesh.json")
UDS = "/tmp/klippy_uds"

# Grille 5x5 : lignes Y = -110..110, colonnes X = -110..110. Les 13 points reels
# et, pour chaque ligne, les colonnes reelles (les autres sont du remplissage).
def grille_depuis_config():
    """Lit mesh_radius / round_probe_count dans printer.cfg et reproduit
    bed_mesh.generate_points() (plateau rond) : coordonnees et colonnes reelles."""
    import math
    cfg = open(CFG).read()
    r = float(re.search(r"^mesh_radius:\s*([0-9.]+)", cfg, re.M).group(1))
    n = int(re.search(r"^round_probe_count:\s*(\d+)", cfg, re.M).group(1))
    x_dist = math.floor((2 * r / (n - 1)) * 100) / 100
    new_r = (n // 2) * x_dist
    coords = [round(-new_r + i * x_dist, 3) for i in range(n)]
    reels = {i: [j for j, x in enumerate(coords) if math.hypot(x, coords[i]) <= r] for i in range(n)}
    return n, coords, reels


N, XS, REELS = grille_depuis_config()
YS = XS
ORDRE = [(r, c) for r in range(N) for c in REELS[r]]      # avant->arriere, gauche->droite


def gcode(cmd, timeout=120):
    """Envoie une commande G-code a klippy et attend son acquittement, sans rien afficher."""
    s = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
    s.settimeout(timeout)
    s.connect(UDS)
    s.sendall(json.dumps({"id": 7, "method": "gcode/script", "params": {"script": cmd}}).encode() + b"\x03")
    buf = b""
    while True:
        buf += s.recv(65536)
        while b"\x03" in buf:
            raw, buf = buf.split(b"\x03", 1)
            try:
                msg = json.loads(raw)
            except ValueError:
                continue
            if msg.get("id") == 7:
                s.close()
                if "error" in msg:
                    print("    (klippy: %s)" % msg["error"].get("message", msg["error"]))
                    return False
                return True


HAUTEUR_POINTAGE = 100     # la buse se place a cette hauteur au-dessus du carre demande


def pointer(x, y):
    """Place la buse a HAUTEUR_POINTAGE au-dessus du carre (x, y).
    Si le homing a ete perdu (timeout d'inactivite de Klipper), re-home et reessaie."""
    cmd = "G90\nG1 Z%d F3000\nG1 X%d Y%d F6000\nM400" % (HAUTEUR_POINTAGE, x, y)
    if gcode(cmd):
        return True
    print("    (homing perdu : G28 puis nouvel essai)")
    return gcode("G28", timeout=180) and gcode(cmd)


def mm(texte):
    """Convertit une saisie en mm : '0.21' (mm), '8t' ou '8.5t' (milliemes de pouce), '0.008in' (pouces)."""
    t = texte.strip().lower().replace(",", ".")
    if t.endswith("in"):
        return float(t[:-2]) * 25.4
    if t.endswith("t") or t.endswith("th"):
        return float(t.rstrip("h").rstrip("t")) * 0.0254
    return float(t)


def en_impression():
    try:
        s = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
        s.settimeout(3)
        s.connect(UDS)
        s.sendall(json.dumps({"id": 1, "method": "objects/query",
                              "params": {"objects": {"print_stats": None}}}).encode() + b"\x03")
        buf = b""
        while b"\x03" not in buf:
            buf += s.recv(65536)
        s.close()
        return json.loads(buf.split(b"\x03")[0])["result"]["status"]["print_stats"]["state"] in ("printing", "paused")
    except OSError:
        return False        # klippy arrete : pas d'impression possible


def lire_points(cfg):
    m = re.search(r"^#\*# points =\n((?:#\*# \t.*\n){%d})" % N, cfg, re.M)
    if not m:
        sys.exit("bloc '#*# points =' introuvable dans printer.cfg")
    lignes = [l.replace("#*# \t", "").strip() for l in m.group(1).strip("\n").split("\n")]
    mat = [[float(v) for v in l.split(",")] for l in lignes]
    assert all(len(r) == N for r in mat) and len(mat) == N, mat
    return m, mat


def ecrire_points(cfg, m, mat):
    bloc = "#*# points =\n" + "".join("#*# \t  " + ", ".join("%.6f" % v for v in r) + "\n" for r in mat)
    return cfg[:m.start()] + bloc + cfg[m.end():]


def repad(mat):
    """Refait le remplissage des colonnes hors disque comme Klipper : copie du bord."""
    for r, cols in REELS.items():
        lo, hi = cols[0], cols[-1]
        for c in range(N):
            if c < lo:
                mat[r][c] = mat[r][lo]
            elif c > hi:
                mat[r][c] = mat[r][hi]
    return mat


def saisir():
    try:
        vals = json.load(open(FICHIER))
    except (OSError, ValueError):
        vals = {}
    pointage = "--sans-pointage" not in sys.argv
    if pointage:
        print("\nHoming, puis la buse pointera chaque carre a %d mm au-dessus..." % HAUTEUR_POINTAGE)
        if not gcode("G28", timeout=180):
            sys.exit("homing impossible (klippy arrete ?)")
    else:
        print("\nSans pointage : chaque carre est repere par sa rangee (1 = AVANT, cote logo)"
              " et son rang dans la rangee (1 = GAUCHE).")
    print("\nEpaisseur de chaque carre (mm, ex: 0.21). Entree = garder, r = precedent.")
    print("Carre qui n'a PAS colle (extrude en l'air) : saisir 0.6 (= tres loin). Carre racle, quasi rien : mesurer ce qui reste.")
    i = 0
    while i < len(ORDRE):
        r, c = ORDRE[i]
        cle = "%g,%g" % (XS[c], YS[r])
        defaut = vals.get(cle)
        if pointage:
            pointer(XS[c], YS[r])
        ou = "" if pointage else "  rangee %d/%d, %d%s de gauche sur %d" % (
            r + 1, N, REELS[r].index(c) + 1, "er" if REELS[r].index(c) == 0 else "e", len(REELS[r]))
        try:
            s = input("  %2d. carre (X=%4g, Y=%4g)%s%s : " % (i + 1, XS[c], YS[r], ou,
                                                             ("  (Entree = %g)" % defaut) if defaut is not None else "")).strip()
        except EOFError:
            print("\n(memorise dans mesures-mesh.json)")
            sys.exit(0)
        if s == "" and defaut is not None:
            i += 1
            continue
        if s.lower() == "r":
            i = max(0, i - 1)
            continue
        try:
            v = mm(s)
            if not (0.05 <= v <= 1.0):
                raise ValueError
        except ValueError:
            print("    -> une epaisseur en mm (0.05..1.0), ou en thou avec un t (ex: 8t = 0.203 mm), ou en pouces avec in")
            continue
        vals[cle] = v
        json.dump(vals, open(FICHIER, "w"), indent=1)
        i += 1
    return vals


def main():
    print(__doc__)
    if en_impression():
        sys.exit("impression en cours : on ne touche pas a la mesh maintenant")
    cfg = open(CFG).read()
    m, mat = lire_points(cfg)
    if "--sans-saisie" in sys.argv:          # reprendre les mesures memorisees telles quelles
        vals = json.load(open(FICHIER))
        manque = [cle for r, c in ORDRE for cle in ["%g,%g" % (XS[c], YS[r])] if cle not in vals]
        if manque:
            sys.exit("mesures manquantes dans mesures-mesh.json : " + ", ".join(manque))
    else:
        vals = saisir()
    mesures = [vals["%g,%g" % (XS[c], YS[r])] for r, c in ORDRE]
    moy = sum(mesures) / len(mesures)
    print("\nepaisseur moyenne : %.3f mm  (nominal 0.20)" % moy)
    print("%-22s %8s %8s %8s" % ("carre", "mesure", "mesh av.", "mesh ap."))
    # Gain < 1 : on ne corrige qu'une fraction de l'ecart mesure, pour ne pas
    # osciller (le pied a coulisse surestime l'ecart quand les cordons sont
    # ronds et mal soudes -> la passe 1 a depasse par endroits).
    s_g = input("\nGain de correction (1 = plein, 0.7 conseille a partir de la 2e passe) [0.7] : ").strip().replace(",", ".")
    try:
        gain = float(s_g) if s_g else 0.7
        if not (0.1 <= gain <= 1.0):
            raise ValueError
    except ValueError:
        sys.exit("gain invalide (0.1..1)")
    nouveau = [row[:] for row in mat]
    for (r, c), e in zip(ORDRE, mesures):
        delta = (e - moy) * gain
        nouveau[r][c] = mat[r][c] - delta
        print("%-22s %8.3f %8.3f %8.3f   (%+.3f)" % ("(X=%g, Y=%g)" % (XS[c], YS[r]), e, mat[r][c], nouveau[r][c], -delta))
    nouveau = repad(nouveau)
    etendue = max(mesures) - min(mesures)
    print("\netendue des epaisseurs : %.3f mm  ->  apres correction, attendu ~0 au prochain test" % etendue)
    # Zero global : la moyenne conservee n'est juste que si la couche moyenne
    # etait bonne. Si la spirale montre "du manque" partout (trop loin) on
    # descend le zero ; "ecrase" partout -> on le monte. Applique dans START_PRINT.
    mz = re.search(r"^(\s*SET_GCODE_OFFSET Z=)(-?[0-9.]+)", cfg, re.M)
    z_actuel = float(mz.group(2)) if mz else None
    dz = 0.0
    if mz:
        s_dz = input("\nDecalage du zero global en plus (mm, ex: -0.05 = buse plus bas ; Entree = 0) [START_PRINT actuel Z=%g] : " % z_actuel).strip().replace(",", ".")
        if s_dz:
            try:
                dz = float(s_dz)
                if abs(dz) > 0.3:
                    raise ValueError
            except ValueError:
                sys.exit("decalage invalide (|dz| <= 0.3)")
    if input("\nEcrire cette mesh%s dans printer.cfg et redemarrer klippy (o/N) ? "
             % (" + zero global Z=%.3f" % (z_actuel + dz) if dz else "")).strip().lower() != "o":
        sys.exit("rien ecrit (mesures memorisees)")
    sauvegarde = CFG.replace(".cfg", "-avant-fixmesh-%s.cfg" % time.strftime("%Y%m%d_%H%M%S"))
    shutil.copy(CFG, sauvegarde)
    nouveau_cfg = ecrire_points(cfg, m, nouveau)
    if dz:
        nouveau_cfg = re.sub(r"^(\s*SET_GCODE_OFFSET Z=)(-?[0-9.]+)",
                             lambda mm: "%s%.3f" % (mm.group(1), z_actuel + dz), nouveau_cfg, count=1, flags=re.M)
        print("START_PRINT : SET_GCODE_OFFSET Z=%.3f -> %.3f" % (z_actuel, z_actuel + dz))
    open(CFG, "w").write(nouveau_cfg)
    print("ecrit. sauvegarde :", os.path.basename(sauvegarde))
    subprocess.run([os.path.join(HERE, "start-klipper.sh"), "restart"])
    print("klippy redemarre : attendre ~15 s, puis homing (G28) avant d'imprimer.")


if __name__ == "__main__":
    main()
