import numpy as np
import mujoco
import onnxruntime as ort
import imageio

MODEL_XML = "mujoco_menagerie/unitree_go2/scene.xml"
POLICY_ONNX = "/home/ubuntu/isaacsim/IsaacLab/logs/rsl_rl/unitree_go2_flat/2026-09-11_23-24-25/exported/policy.onnx"

# Policy joint order: PhysX articulation order used by this training run (confirmed via
# scripts/sim2sim_transfer/config/newton_to_physx_go2.yaml target_joint_names, and env.yaml
# showing physics backend = isaaclab_physx.physics.physx_manager:PhysxManager).
# PhysX groups joints by TYPE (all hips, then all thighs, then all calves), each in FL,FR,RL,RR order.
POLICY_JOINT_ORDER = [
    "FL_hip_joint", "FR_hip_joint", "RL_hip_joint", "RR_hip_joint",
    "FL_thigh_joint", "FR_thigh_joint", "RL_thigh_joint", "RR_thigh_joint",
    "FL_calf_joint", "FR_calf_joint", "RL_calf_joint", "RR_calf_joint",
]
DEFAULT_JOINT_POS = {
    "FL_hip_joint": 0.1, "FR_hip_joint": -0.1, "RL_hip_joint": 0.1, "RR_hip_joint": -0.1,
    "FL_thigh_joint": 0.8, "FR_thigh_joint": 0.8, "RL_thigh_joint": 1.0, "RR_thigh_joint": 1.0,
    "FL_calf_joint": -1.5, "FR_calf_joint": -1.5, "RL_calf_joint": -1.5, "RR_calf_joint": -1.5,
}
ACTION_SCALE = 0.25
STIFFNESS = 25.0
DAMPING = 0.5
SIM_DT = 0.005
DECIMATION = 4
CONTROL_DT = SIM_DT * DECIMATION
SIM_DURATION = 10.0

CMD = np.array([0.5, 0.0, 0.0], dtype=np.float32)  # lin_vel_x, lin_vel_y, ang_vel_z

model = mujoco.MjModel.from_xml_path(MODEL_XML)
model.opt.timestep = SIM_DT
data = mujoco.MjData(model)

joint_qpos_adr = []
joint_qvel_adr = []
actuator_id = []
default_pos = np.zeros(12, dtype=np.float32)
for i, jn in enumerate(POLICY_JOINT_ORDER):
    jid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_JOINT, jn)
    joint_qpos_adr.append(model.jnt_qposadr[jid])
    joint_qvel_adr.append(model.jnt_dofadr[jid])
    act_name = jn.replace("_joint", "")
    aid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_ACTUATOR, act_name)
    actuator_id.append(aid)
    default_pos[i] = DEFAULT_JOINT_POS[jn]

joint_qpos_adr = np.array(joint_qpos_adr)
joint_qvel_adr = np.array(joint_qvel_adr)
actuator_id = np.array(actuator_id)

mujoco.mj_resetDataKeyframe(model, data, 0) if model.nkey > 0 else mujoco.mj_forward(model, data)
# set default joint pose explicitly (freejoint occupies qpos[0:7])
data.qpos[joint_qpos_adr] = default_pos
mujoco.mj_forward(model, data)

sess = ort.InferenceSession(POLICY_ONNX, providers=["CPUExecutionProvider"])
input_name = sess.get_inputs()[0].name
output_name = sess.get_outputs()[0].name

last_action = np.zeros(12, dtype=np.float32)
q_targets = default_pos.copy()

frames = []
renderer = mujoco.Renderer(model, height=480, width=640)
cam = mujoco.MjvCamera()
mujoco.mjv_defaultCamera(cam)
cam.distance = 1.5
cam.azimuth = 120
cam.elevation = -20
cam.trackbodyid = mujoco.mj_name2id(model, mujoco.mjtObj.mjOBJ_BODY, "base")
cam.type = mujoco.mjtCamera.mjCAMERA_TRACKING

n_steps = int(SIM_DURATION / SIM_DT)
control_every = DECIMATION
base_heights = []
rp_list = []
frame_stride = int((1 / 30) / SIM_DT)
SETTLE_STEPS = int(0.5 / CONTROL_DT)  # 0.5s standing before commanding velocity

for step in range(n_steps):
    if step % control_every == 0:
        qpos = data.qpos[joint_qpos_adr].astype(np.float32)
        qvel = data.qvel[joint_qvel_adr].astype(np.float32)
        joint_pos_rel = qpos - default_pos
        joint_vel_rel = qvel

        quat = data.qpos[3:7]  # w,x,y,z (mujoco convention)
        w, x, y, z = quat
        # gravity vector in body frame
        R = np.array([
            [1 - 2*(y*y+z*z), 2*(x*y - z*w), 2*(x*z + y*w)],
            [2*(x*y + z*w), 1 - 2*(x*x+z*z), 2*(y*z - x*w)],
            [2*(x*z - y*w), 2*(y*z + x*w), 1 - 2*(x*x+y*y)],
        ])
        gravity_world = np.array([0, 0, -1.0])
        projected_gravity = R.T @ gravity_world

        lin_vel_world = data.qvel[0:3]
        ang_vel_world = data.qvel[3:6]
        base_lin_vel = R.T @ lin_vel_world
        base_ang_vel = R.T @ ang_vel_world

        cmd_now = CMD if step >= SETTLE_STEPS else np.zeros(3, dtype=np.float32)
        obs = np.concatenate([
            base_lin_vel, base_ang_vel, projected_gravity, cmd_now,
            joint_pos_rel, joint_vel_rel, last_action,
        ]).astype(np.float32).reshape(1, -1)

        action = sess.run([output_name], {input_name: obs})[0].reshape(-1)
        last_action = action.copy()
        q_targets = default_pos + action * ACTION_SCALE

    q = data.qpos[joint_qpos_adr]
    dq = data.qvel[joint_qvel_adr]
    torque = STIFFNESS * (q_targets - q) - DAMPING * dq
    data.ctrl[actuator_id] = torque

    mujoco.mj_step(model, data)

    base_heights.append(data.qpos[2])
    quat_now = data.qpos[3:7]
    w, x, y, z = quat_now
    roll = np.arctan2(2*(w*x+y*z), 1-2*(x*x+y*y))
    pitch = np.arcsin(np.clip(2*(w*y-z*x), -1, 1))
    rp_list.append((np.degrees(roll), np.degrees(pitch)))
    if step % frame_stride == 0:
        renderer.update_scene(data, camera=cam)
        frames.append(renderer.render().copy())

base_heights = np.array(base_heights)
rp_arr = np.array(rp_list)
print("min height:", base_heights.min(), "max height:", base_heights.max(), "final height:", base_heights[-1])
print("any nan:", np.isnan(data.qpos).any())
print("final base xy:", data.qpos[0], data.qpos[1])
print("roll range:", rp_arr[:,0].min(), rp_arr[:,0].max())
print("pitch range:", rp_arr[:,1].min(), rp_arr[:,1].max())

imageio.mimsave("go2_run.mp4", frames, fps=30)
print("saved video: go2_run.mp4, frames:", len(frames))
