"""Score a policy over many randomised episodes.

Mean training return is a poor headline number: it mixes survival with reward
shaping and it is measured on the same randomisation the policy trained on.
What you actually want to know before touching hardware is "out of 100 robots
with different masses, frictions, motor gains and control latency, how many
stay up, and how well do they follow the command".

  python eval.py --policy ../runs/biped_v1/policy_best.json --episodes 100
"""

import argparse
from pathlib import Path

import numpy as np

from env import RobotEnv
from ppo import Policy


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--policy", required=True)
    ap.add_argument("--spec", default=None)
    ap.add_argument("--task", default=None)
    ap.add_argument("--episodes", type=int, default=100)
    ap.add_argument("--seconds", type=float, default=None)
    ap.add_argument("--no-dr", action="store_true")
    ap.add_argument("--seed", type=int, default=9000)
    args = ap.parse_args()

    policy = Policy.load(args.policy)
    spec_path = args.spec or policy.meta.get("spec") or str(
        Path(__file__).parent.parent / "robots/wheeled_biped.json")

    survived, returns, vx_err, wz_err, heights, wins = [], [], [], [], [], []

    for ep in range(args.episodes):
        env = RobotEnv(spec_path, seed=args.seed + ep,
                       task=args.task or policy.meta.get("task"),
                       layout=policy.meta.get("obs_layout", "legacy"),
                       randomise=(False if args.no_dr else None))
        if args.seconds:
            env.max_steps = int(args.seconds / env.control_dt)
        rng = np.random.default_rng(args.seed + ep)
        obs = env.reset()
        total, fell = 0.0, False
        for _ in range(env.max_steps):
            a, _, _ = policy.act(obs[None], rng, deterministic=True)
            obs, r, fell, trunc, info = env.step(np.clip(a[0], -1, 1))
            total += r
            # Tracking error only means something for incentives that issue a
            # velocity command; the rest are scored on survival and success.
            cmd = getattr(env.task_obj, "cmd", None)
            if cmd is not None:
                vx_err.append(abs(float(env.vel_local[0]) - cmd[0]))
                wz_err.append(abs(float(env.gyro_true[2]) - cmd[1]))
            heights.append(info["height"])
            if fell:
                break
        survived.append(env.t)
        returns.append(total)
        wins.append(bool(env.task_obj.success(env)))
        if (ep + 1) % 20 == 0:
            print(f"  {ep+1}/{args.episodes} episodes")

    full = env.max_steps * env.control_dt
    surv = np.array(survived)
    rate = float(np.mean(surv >= full - 1e-6))

    print(f"\npolicy   {args.policy}")
    print(f"spec     {Path(spec_path).name}   incentive {env.task_name}")
    print(f"episodes {args.episodes}  x {full:.0f}s  "
          f"domain randomisation {'off' if args.no_dr else 'on'}")
    print(f"trained  {policy.meta.get('steps', '?')} steps")
    print("-" * 52)
    print(f"survival rate       {rate*100:5.1f} %   (full {full:.0f}s episode, no fall)")
    if any(wins):
        print(f"task success rate   {np.mean(wins)*100:5.1f} %   "
              f"(reached the goal the incentive defines)")
    print(f"time upright        {surv.mean():5.2f} s  median {np.median(surv):.2f}  "
          f"worst {surv.min():.2f}")
    print(f"mean return         {np.mean(returns):7.1f}")
    print(f"base height         {np.mean(heights):5.3f} m  "
          f"(target {env.h_target:.3f})")
    if vx_err:
        print(f"vx tracking error   {np.mean(vx_err):5.3f} m/s")
        print(f"yaw tracking error  {np.mean(wz_err):5.3f} rad/s")


if __name__ == "__main__":
    main()
