"""Generatore di camminata a oscillatori (CPG), identico alla semantica di SPEC §5.1.

Il CPG è l'alternativa "leggera" alla rete neurale: 12 numeri (gait.json) invece di ~1800 pesi.
train_ars.py --cpg ottimizza questi 12 numeri con una Evolution Strategy in simulazione.

Convenzioni IMU usate per il feedback (vedi env.PulcinoEnv.imu_roll_pitch):
    roll  > 0 = fianco sinistro in su   (rotazione positiva attorno a +x)
    pitch > 0 = muso in giù             (rotazione positiva attorno a +y)
"""
import json

import numpy as np

Q_MIN = np.array([-0.35, -0.8, -0.8, -0.35, -0.8, -0.8])
Q_MAX = -Q_MIN
TIPI = ["hip_roll", "hip_pitch", "ankle_pitch"]

# gait di esempio di SPEC §5.1 (punto di partenza dell'ottimizzazione)
GAIT_ESEMPIO = {
    "version": 1, "type": "cpg", "name": "esempio-spec",
    "freq": 1.6,
    "joints": {
        "hip_roll": {"amp": 0.18, "phase": 0.0, "offset": 0.02},
        "hip_pitch": {"amp": 0.30, "phase": 1.57, "offset": 0.0},
        "ankle_pitch": {"amp": 0.20, "phase": 1.57, "offset": 0.0},
    },
    "feedback": {"roll_gain": 0.0, "pitch_gain": 0.0},
}

# punto di partenza "prudente" per l'ottimizzazione: stesse fasi dell'esempio, ampiezze ridotte
# (con le ampiezze piene dell'esempio il robot simulato cade in < 1 s)
GAIT_PARTENZA = {
    "version": 1, "type": "cpg", "name": "partenza",
    "freq": 1.6,
    "joints": {
        "hip_roll": {"amp": 0.08, "phase": 0.0, "offset": 0.02},
        "hip_pitch": {"amp": 0.15, "phase": 1.57, "offset": 0.0},
        "ankle_pitch": {"amp": 0.10, "phase": 1.57, "offset": 0.0},
    },
    "feedback": {"roll_gain": 0.0, "pitch_gain": 0.0},
}

# (chiave, minimo, massimo) dei 12 parametri ottimizzati, nell'ordine del vettore
PARAMETRI = [
    ("freq", 0.8, 2.5),
    ("hip_roll.amp", 0.0, 0.35), ("hip_roll.phase", -np.pi, np.pi), ("hip_roll.offset", -0.10, 0.20),
    ("hip_pitch.amp", 0.0, 0.8), ("hip_pitch.phase", -np.pi, np.pi), ("hip_pitch.offset", -0.3, 0.3),
    ("ankle_pitch.amp", 0.0, 0.8), ("ankle_pitch.phase", -np.pi, np.pi), ("ankle_pitch.offset", -0.3, 0.3),
    ("feedback.roll_gain", -2.0, 2.0), ("feedback.pitch_gain", -2.0, 2.0),
]
N_PARAM = len(PARAMETRI)


def gait_a_vettore(g):
    """gait.json → vettore normalizzato u ∈ [0,1]^12."""
    u = np.zeros(N_PARAM)
    for i, (k, lo, hi) in enumerate(PARAMETRI):
        if k == "freq":
            v = g["freq"]
        elif k.startswith("feedback."):
            v = g["feedback"][k.split(".")[1]]
        else:
            j, p = k.split(".")
            v = g["joints"][j][p]
        u[i] = (v - lo) / (hi - lo)
    return np.clip(u, 0.0, 1.0)


def vettore_a_gait(u, nome="ars-cpg"):
    """Vettore normalizzato → dizionario gait.json (formato SPEC §5.1)."""
    u = np.clip(u, 0.0, 1.0)
    v = {k: lo + (hi - lo) * x for (k, lo, hi), x in zip(PARAMETRI, u)}
    return {
        "version": 1, "type": "cpg", "name": nome,
        "freq": round(float(v["freq"]), 4),
        "joints": {j: {p: round(float(v[f"{j}.{p}"]), 4) for p in ("amp", "phase", "offset")}
                   for j in TIPI},
        "feedback": {"roll_gain": round(float(v["feedback.roll_gain"]), 4),
                     "pitch_gain": round(float(v["feedback.pitch_gain"]), 4)},
    }


class CPG:
    """Valuta un gait.json nel tempo, esattamente come descritto in SPEC §5.1."""

    def __init__(self, gait):
        self.g = gait

    @classmethod
    def da_file(cls, path):
        with open(path) as f:
            return cls(json.load(f))

    def q(self, t, vx=1.0, wz=0.0, speed=1.0, roll_imu=0.0, pitch_imu=0.0):
        g = self.g
        f = g["freq"] * float(np.clip(speed, 0.3, 1.5))
        vx = float(np.clip(vx, -1, 1)); wz = float(np.clip(wz, -1, 1))
        q = np.zeros(6)
        for gamba, base, dphi in (("L", 0, 0.0), ("R", 3, np.pi)):
            for k, tipo in enumerate(TIPI):
                p = g["joints"][tipo]
                amp = p["amp"]
                if tipo in ("hip_pitch", "ankle_pitch"):
                    amp *= vx                                   # vx<0 → camminata all'indietro
                if tipo == "hip_pitch":
                    amp *= (1 - 0.6 * wz) if gamba == "L" else (1 + 0.6 * wz)
                q[base + k] = p["offset"] + amp * np.sin(2 * np.pi * f * t + p["phase"] + dphi)
            q[base + 0] += g["feedback"]["roll_gain"] * roll_imu
            q[base + 2] += g["feedback"]["pitch_gain"] * pitch_imu
        return np.clip(q, Q_MIN, Q_MAX)
