"""Train a policy for a robot spec.

  python train.py --spec ../robots/wheeled_biped.json --steps 5_000_000

Curriculum: the velocity command ramps from zero to full over the first
`--cmd-ramp` fraction of training, so the robot learns to stand before it is
asked to go anywhere.
"""

import argparse
import json
import time
from pathlib import Path

import numpy as np

from env import RobotEnv
from ppo import PPO, Policy, gae
from rollout import record
from vecenv import make_vec


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--spec", default=str(Path(__file__).parent.parent / "robots/wheeled_biped.json"))
    ap.add_argument("--task", default=None,
                    help="incentive name (stand, drive, spin, goto, climb, recover)")
    ap.add_argument("--run", default=None, help="run name under runs/")
    ap.add_argument("--steps", type=int, default=5_000_000)
    ap.add_argument("--envs", type=int, default=16)
    ap.add_argument("--workers", type=int, default=3)
    ap.add_argument("--rollout", type=int, default=128, help="steps per env per iteration")
    ap.add_argument("--lr", type=float, default=3e-4)
    ap.add_argument("--gamma", type=float, default=0.99)
    ap.add_argument("--lam", type=float, default=0.95)
    ap.add_argument("--clip", type=float, default=0.2)
    ap.add_argument("--epochs", type=int, default=6)
    ap.add_argument("--minibatches", type=int, default=4)
    ap.add_argument("--ent-coef", type=float, default=0.004)
    ap.add_argument("--hidden", type=int, nargs="+", default=[64, 64])
    ap.add_argument("--cmd-ramp", type=float, default=0.35)
    ap.add_argument("--no-dr", action="store_true", help="disable domain randomisation")
    ap.add_argument("--seed", type=int, default=0)
    ap.add_argument("--resume", default=None)
    ap.add_argument("--layout", default="union", choices=["union", "legacy"],
                    help="goal-channel layout; union lets a policy carry between stages")
    ap.add_argument("--snapshot-every", type=int, default=500_000,
                    help="record a watchable replay every N steps (0 to disable)")
    args = ap.parse_args()

    spec_path = str(Path(args.spec).resolve())
    run_name = args.run or "_".join(filter(None, [
        Path(spec_path).stem, args.task, time.strftime("%Y%m%d_%H%M%S")]))
    run_dir = Path(__file__).parent.parent / "runs" / run_name
    run_dir.mkdir(parents=True, exist_ok=True)

    rng = np.random.default_rng(args.seed)
    venv = make_vec(spec_path, args.envs, args.workers, seed=args.seed,
                    randomise=(False if args.no_dr else None), cmd_scale=0.0,
                    task=args.task, layout=args.layout)

    if args.resume:
        policy = Policy.load(args.resume)
        print(f"resumed from {args.resume}")
    else:
        policy = Policy(venv.obs_dim, venv.act_dim, hidden=tuple(args.hidden), seed=args.seed)

    algo = PPO(policy, lr=args.lr, clip=args.clip, epochs=args.epochs,
               minibatches=args.minibatches, ent_coef=args.ent_coef)

    T, N = args.rollout, venv.n
    batch = T * N
    iters = max(1, args.steps // batch)

    # A dedicated single env for snapshots. Watching the same seed at 0.5M,
    # 1M, 2M steps is the clearest possible answer to "what has it learned",
    # and it costs one episode of simulation per snapshot.
    snap_env = None
    if args.snapshot_every > 0:
        snap_env = RobotEnv(spec_path, seed=4242, task=args.task,
                            layout=args.layout, randomise=False)
        (run_dir / "snapshots").mkdir(exist_ok=True)
    next_snap = 0

    obs = venv.reset()
    ep_ret = np.zeros(N)
    ep_len = np.zeros(N, dtype=int)
    ret_hist, len_hist = [], []
    best = -1e18
    log_path = run_dir / "log.jsonl"
    t_start = time.time()

    print(f"run {run_name}  incentive={venv.task_name}  obs={venv.obs_dim} "
          f"act={venv.act_dim} envs={N} batch={batch} iters={iters}", flush=True)

    for it in range(1, iters + 1):
        frac = (it - 1) / max(1, iters - 1)
        cmd_scale = min(1.0, frac / args.cmd_ramp) if args.cmd_ramp > 0 else 1.0
        venv.set_cmd_scale(cmd_scale)

        b_obs = np.empty((T, N, venv.obs_dim))
        b_act = np.empty((T, N, venv.act_dim))
        b_logp = np.empty((T, N))
        b_val = np.empty((T, N))
        b_rew = np.empty((T, N))
        b_term = np.zeros((T, N), dtype=bool)
        b_trunc = np.zeros((T, N), dtype=bool)

        for t in range(T):
            a, v, lp = policy.act(obs, rng)
            b_obs[t], b_act[t], b_logp[t], b_val[t] = obs, a, lp, v
            obs, rew, term, trunc, final = venv.step(np.clip(a, -1, 1))
            b_rew[t], b_term[t], b_trunc[t] = rew, term, trunc

            # A time-limit cut is not a real terminal state: fold the value of
            # the observation we were about to see back into the reward.
            if final:
                fidx = sorted(final)
                fv = policy.value(np.stack([final[i] for i in fidx]))
                for j, i in enumerate(fidx):
                    b_rew[t, i] += args.gamma * fv[j]

            ep_ret += rew
            ep_len += 1
            done = term | trunc
            if done.any():
                ret_hist += list(ep_ret[done])
                len_hist += list(ep_len[done])
                ep_ret[done] = 0.0
                ep_len[done] = 0

        last_val = policy.value(obs)
        adv, ret = gae(b_rew, b_val, b_term, b_trunc, last_val, args.gamma, args.lam)

        flat_obs = b_obs.reshape(-1, venv.obs_dim)
        policy.norm.update(flat_obs)
        stats = algo.update(flat_obs, b_act.reshape(-1, venv.act_dim),
                            b_logp.reshape(-1), adv.reshape(-1), ret.reshape(-1), rng)

        ret_hist, len_hist = ret_hist[-200:], len_hist[-200:]
        mean_ret = float(np.mean(ret_hist)) if ret_hist else float("nan")
        mean_len = float(np.mean(len_hist)) if len_hist else float("nan")
        steps_done = it * batch
        elapsed = time.time() - t_start
        row = {
            "iter": it, "task": venv.task_name, "steps": steps_done, "cmd_scale": round(cmd_scale, 3),
            "ep_return": round(mean_ret, 2), "ep_len": round(mean_len, 1),
            "ep_seconds": round(mean_len * 0.01, 2),
            "sps": int(steps_done / max(1e-6, elapsed)),
            "elapsed_s": int(elapsed),
            "log_std": round(float(policy.log_std.mean()), 3),
            **{k: round(v, 4) for k, v in stats.items()},
        }
        with open(log_path, "a") as f:
            f.write(json.dumps(row) + "\n")

        if it % 5 == 0 or it == 1:
            print(f"it {it:4d} | {steps_done/1e6:5.2f}M | ret {mean_ret:8.1f} | "
                  f"eplen {mean_len:6.1f} ({mean_len*0.01:4.1f}s) | cmd {cmd_scale:.2f} | "
                  f"kl {stats['kl']:.4f} | {row['sps']} sps | {elapsed/60:.1f} min",
                  flush=True)

        if snap_env is not None and steps_done >= next_snap:
            snap_env.cmd_scale = cmd_scale
            snap = record(snap_env, policy, snap_env.reset(),
                          seconds=snap_env.max_steps * snap_env.control_dt,
                          seed=4242, steps_trained=steps_done)
            with open(run_dir / "snapshots" / f"{steps_done:09d}.json", "w") as f:
                json.dump(snap, f)
            next_snap = steps_done + args.snapshot_every

        policy.save(run_dir / "policy_latest.json",
                    meta={"spec": spec_path, "task": venv.task_name, "iter": it,
                          "steps": steps_done, "ep_return": mean_ret,
                          "ep_len": mean_len, "obs_layout": args.layout})
        if ret_hist and mean_ret > best:
            best = mean_ret
            policy.save(run_dir / "policy_best.json",
                        meta={"spec": spec_path, "task": venv.task_name,
                              "iter": it, "steps": steps_done,
                              "ep_return": mean_ret, "ep_len": mean_len,
                              "obs_layout": args.layout})

    venv.close()
    print(f"done: {run_dir}  best mean return {best:.1f}")


if __name__ == "__main__":
    main()
