"""
Collect pick-and-place demonstrations for the demo2 arm via the scripted teacher.

Each episode randomizes the object and place-pad positions, runs PickPlaceTeacher
to a successful place, and records one goal-conditioned action-chunk sample:
    obs   = [ initial arm qpos (5) , pick xyz (3) , place xyz (3) ]   (11)
    chunk = teacher action trajectory resampled to CHUNK steps, each
            [ 5 arm ctrl , 1 grasp bit ]                              (CHUNK, 6)

This is the BC dataset for models/act_model.py (goal-conditioned action chunking),
mirroring the RX1 v2 pipeline. The runtime plays the predicted chunk open-loop.

Run:  python demo2/collect_arm_demos.py --episodes 120  -> model/demos_arm.npz
"""
from __future__ import annotations

import argparse
from pathlib import Path
import numpy as np
import mujoco

from build_arm_scene import build_model
from arm_pick_place import ArmIK, PickPlaceTeacher, ARM_ACT
import grasp_planner as gp

HERE = Path(__file__).resolve().parent
OUT = HERE / "model" / "demos_arm.npz"
CHUNK = 100
ACT_DIM = 6          # 5 arm ctrl + grasp
OBS_DIM = 11         # 5 arm qpos + pick(3) + place(3)
_GRASP_BIT = len(ARM_ACT)   # index 5 in the action row


def _settle(m, d, ik, steps=150):
    d.ctrl[:] = ik.rest()
    for _ in range(steps):
        mujoco.mj_step(m, d)


_OBJECTS = ("object", "object_red", "object_blue")


def _place_objects(m, d, rng):
    """Scatter the 3 objects at non-overlapping spots in the reachable band."""
    spots = []
    for body in _OBJECTS:
        bid = mujoco.mj_name2id(m, mujoco.mjtObj.mjOBJ_BODY, body)
        qa = int(m.jnt_qposadr[m.body_jntadr[bid]])
        for _ in range(50):
            x = rng.uniform(0.24, 0.30); y = rng.uniform(-0.16, -0.02)
            if all((x - sx) ** 2 + (y - sy) ** 2 > 0.07 ** 2 for sx, sy in spots):
                break
        spots.append((x, y))
        d.qpos[qa:qa + 3] = [x, y, 0.55]
        d.qpos[qa + 3:qa + 7] = [1, 0, 0, 0]


def run_episode(m, d, ik, rng, pp_gid):
    """Run one teacher episode picking a random object; return (obs, chunk)."""
    _place_objects(m, d, rng)
    qx = rng.uniform(0.24, 0.30); qy = rng.uniform(0.06, 0.15)
    m.geom_pos[pp_gid] = [qx, qy, 0.501]
    _settle(m, d, ik)

    target = _OBJECTS[rng.integers(len(_OBJECTS))]
    oid = mujoco.mj_name2id(m, mujoco.mjtObj.mjOBJ_BODY, target)
    sid = ik._site
    pick = d.xpos[oid].copy()
    place = np.array([qx, qy, 0.526])
    teach = PickPlaceTeacher(ik, pick, place)
    obs0 = np.concatenate([[d.ctrl[a] for a in ARM_ACT], pick, place]).astype(np.float32)

    MAXD = 3.0 * 0.002 * 5
    traj = []
    for _ in range(4000):
        ctrl, grasp = teach.act(d.ctrl[:m.nu].copy(), d.xpos[oid].copy(),
                                d.site_xpos[sid].copy())
        prev = d.ctrl[:m.nu]
        d.ctrl[:m.nu] = prev + np.clip(ctrl - prev, -MAXD, MAXD)
        gp.set_weld(m, d, eq_name="grasp_weld", active=grasp, body2=target)
        for _ in range(5):
            mujoco.mj_step(m, d)
        row = np.concatenate([[d.ctrl[a] for a in ARM_ACT], [1.0 if grasp else 0.0]])
        traj.append(row.astype(np.float32))
        if teach.done:
            break

    err = float(np.linalg.norm(d.xpos[oid][:2] - place[:2]))
    if err > 0.06 or len(traj) < CHUNK // 2:
        return None, err
    traj = np.asarray(traj, np.float32)
    idx = np.linspace(0, len(traj) - 1, CHUNK).round().astype(int)
    return (obs0, traj[idx]), err


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--episodes", type=int, default=120)
    ap.add_argument("--seed", type=int, default=0)
    args = ap.parse_args()

    m = build_model(); d = mujoco.MjData(m); ik = ArmIK(m)
    pp_gid = mujoco.mj_name2id(m, mujoco.mjtObj.mjOBJ_GEOM, "place_pad")
    rng = np.random.default_rng(args.seed)

    obs_list, chunk_list = [], []
    ok = 0
    for ep in range(args.episodes):
        res, err = run_episode(m, d, ik, rng, pp_gid)
        if res is not None:
            obs_list.append(res[0]); chunk_list.append(res[1]); ok += 1
        if (ep + 1) % 20 == 0:
            print(f"  {ep+1}/{args.episodes}  kept={ok}  last_err={err:.3f}")

    obs = np.asarray(obs_list, np.float32)
    chunk = np.asarray(chunk_list, np.float32)
    np.savez(OUT, obs=obs, chunk=chunk, chunk_len=CHUNK, act_dim=ACT_DIM,
             obs_dim=OBS_DIM)
    print(f"[collect] {ok}/{args.episodes} successful  obs{obs.shape} chunk{chunk.shape} -> {OUT.name}")


if __name__ == "__main__":
    main()
