#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
arm_wave_drake_ar4.py · LateAI 站内原创教学代码
在 MIT Drake (pydrake) 中「真实模拟」Annin AR4 (MK5) 六轴机械臂：

  1) Parser 加载 ar4_arm.urdf（关节 origin/axis/限位逐字段取自官方
     annin_ar4_description/urdf/ar_macro.xacro，见文件内注释）
  2) 6 个关节各配一个 JointActuator，施加 PD 位置控制扭矩
        tau_i = kp_i*(q_des_i - q_i) + kd_i*(qdot_des_i - qdot_i) - tau_g_i
     （受 effort 限位夹取；tau_g 为重力广义力前馈，托住自重）
  3) 重力场中真实积分动力学方程（含重力 / 惯性 / 科氏耦合）
  4) 每 0.5 s 打印「位置传感器」实测角度（即 plant 状态），末端全局坐标写入 CSV
     —— 等价于真机编码器回读 + 轨迹记录

与网页版（arm_traj_3d.html）的关系：
  网页是浏览器端自写 IK 的运动学演示，不含动力学；
  本脚本是同一台臂在真实动力学下的复现，用来回答「这组关节角真机能不能跟得上」。
  用网页「下载 q_traj.npy」后可直接 --q-traj 回放，逐帧校验跟踪误差。

用法:
  python3 arm_wave_drake_ar4.py                        # 无头数值验证（默认）
  python3 arm_wave_drake_ar4.py --meshcat              # 打开浏览器 3D 视图
  python3 arm_wave_drake_ar4.py --q-traj q_traj.npy    # 回放网页导出的关节轨迹
  python3 arm_wave_drake_ar4.py --duration 20 --amp-deg 20
"""
import argparse
import math
import os

import numpy as np

from pydrake.all import (
    AddMultibodyPlantSceneGraph,
    CollisionFilterDeclaration,
    DiagramBuilder,
    GeometrySet,
    Meshcat,
    MeshcatVisualizer,
    MultibodyPlant,
    Parser,
    Simulator,
)

HERE = os.path.dirname(os.path.abspath(__file__))
URDF = os.path.join(HERE, "ar4_arm.urdf")
JOINT_NAMES = ["J1", "J2", "J3", "J4", "J5", "J6"]

# 官方 xacro 限位（度），仅用于打印提示与做目标角度的安全裁剪
LIMIT_DEG = [(-170, 170), (-42, 90), (-89, 52), (-180, 180), (-105, 105), (-180, 180)]

# 扭矩上限（N·m）：按 AR4 真实配置估算 —— J1~J3 为 NEMA23 + 减速箱，
# J4~J6 为 NEMA17 + 减速箱。真机标定后请替换为实测值。
EFFORT = [20.0, 20.0, 15.0, 8.0, 8.0, 5.0]


def cmd_angle(i, t, freq=0.5, amp_deg=25.0):
    """第 i 个关节 (0-based) 的目标角度(rad)。

    零位取官方限位区间内的安全中点，再叠加相位依次错开的正弦摆动 ——
    AR4 的 J2/J3 限位不对称（-42°~90° / -89°~52°），不能简单地以 0 为中心大摆。
    """
    lo, hi = np.radians(LIMIT_DEG[i])
    mid = 0.5 * (lo + hi)
    span = 0.5 * (hi - lo)
    amp = min(math.radians(amp_deg), 0.6 * span)   # 留 40% 余量，避免顶到限位
    return mid + amp * math.sin(freq * t + i * math.pi / 3.0)


def cmd_angle_dot(i, t, freq=0.5, amp_deg=25.0):
    """目标角速度(rad/s)，即上式的解析导数（用于速度前馈，减小跟踪相位滞后）。"""
    lo, hi = np.radians(LIMIT_DEG[i])
    span = 0.5 * (hi - lo)
    amp = min(math.radians(amp_deg), 0.6 * span)
    return amp * freq * math.cos(freq * t + i * math.pi / 3.0)


def effective_inertia(plant, ctx, vel_idx):
    """动能法测各关节有效转动惯量：令 v = e_i，则 T = 0.5·I_i ⇒ I_i = 2T。"""
    g_backup = plant.gravity_field().gravity_vector().copy()
    plant.mutable_gravity_field().set_gravity_vector([0, 0, 0])
    out = []
    for s in vel_idx:
        v = np.zeros(plant.num_velocities())
        v[s] = 1.0
        plant.SetVelocities(ctx, v)
        out.append(2.0 * plant.CalcKineticEnergy(ctx))
    plant.SetVelocities(ctx, np.zeros(plant.num_velocities()))
    plant.mutable_gravity_field().set_gravity_vector(g_backup)
    return np.array(out)


def auto_gain(effort, inertia, dt):
    """按「线性域宽度 + 离散化稳定判据」自动整定 PD。

    kp = min(effort/0.05,  0.5·(1.5/Δt)²·I)
      前项保证 0.05 rad 误差内电机不饱和（线性工作区），
      后项保证 1kHz 数字控制的稳定带宽（ω·Δt < 1.5）。
    kd = min(1.6·sqrt(kp·I),  I/(2Δt))
      前项是临界阻尼略偏大的系数，用于抑制摆动超调；
      后项是离散化的硬约束：微分时间常数 τ = I/kd 必须 ≥ 2Δt，
      否则一步之内速度反馈就能把力矩顶到上限，形成正负交替的高频
      极限环。腕部轴惯量比 J2 小 2~3 个量级（J6 仅 9e-5 kg·m²），
      不加这条约束时 J6 的 kd=0.15 对应 τ=0.6ms < Δt=1ms，
      实测角速度飙到 ±47 rad/s（目标仅 0.22 rad/s），
      表现为「饱和占比 74% 但跟踪误差却只有 0.02 rad」的假象。
    """
    I = np.array(inertia)
    kp = np.minimum(np.array(effort) / 0.05, 0.5 * (1.5 / dt) ** 2 * I)
    kd = np.minimum(1.6 * np.sqrt(kp * I), I / (2.0 * dt))
    return kp, kd


def load_traj(path):
    """读取网页导出的 q_traj.npy（N×6 关节角，rad）。"""
    q = np.load(path)
    q = np.atleast_2d(q)
    if q.shape[1] != 6 and q.shape[0] == 6:
        q = q.T
    if q.shape[1] != 6:
        raise ValueError("q_traj 形状应为 N×6，实际为 %s" % (q.shape,))
    return q


def main():
    ap = argparse.ArgumentParser(description="Annin AR4 (MK5) · MIT Drake 真实动力学模拟")
    ap.add_argument("--urdf", default=URDF, help="URDF 路径（默认同目录 ar4_arm.urdf）")
    ap.add_argument("--duration", type=float, default=14.0, help="仿真时长(秒)")
    ap.add_argument("--dt", type=float, default=0.001, help="控制刷新周期(秒，1kHz)")
    ap.add_argument("--freq", type=float, default=0.5, help="摆动圆频率(rad/s)")
    ap.add_argument("--amp-deg", type=float, default=25.0, help="摆动幅值(°)")
    ap.add_argument("--kp", type=float, default=None, help="统一 PD 位置增益（默认自动整定）")
    ap.add_argument("--kd", type=float, default=None, help="统一 PD 速度阻尼（默认自动整定）")
    ap.add_argument("--q-traj", default=None, help="回放网页导出的 q_traj.npy（N×6）")
    ap.add_argument("--traj-period", type=float, default=5.0,
                    help="q_traj 对应的原始周期(秒)，用于线性重采样")
    ap.add_argument("--print-every", type=float, default=0.5, help="打印周期(秒)")
    ap.add_argument("--csv", default=os.path.join(HERE, "ar4_trail.csv"), help="末端轨迹 CSV 输出")
    ap.add_argument("--meshcat", action="store_true", help="启动浏览器 3D 可视化")
    ap.add_argument("--no-ff", action="store_true", help="关闭重力补偿前馈（诊断用）")
    ap.add_argument("--no-gravity", action="store_true", help="关闭重力场（诊断用）")
    ap.add_argument("--debug", action="store_true", help="打印命令/误差/重力矩调试行")
    args = ap.parse_args()

    # ---------- 1) 建世界：加载 URDF ----------
    # 离散时间 plant：与主循环控制刷新率一致(1 kHz)，模拟真机
    #「微控制器每秒发 1000 次位置命令 + 驱动器 1 kHz 更新」的节奏。
    if args.meshcat:
        builder = DiagramBuilder()
        plant, scene_graph = AddMultibodyPlantSceneGraph(builder, time_step=args.dt)
    else:
        plant = MultibodyPlant(time_step=args.dt)

    Parser(plant).AddModels(args.urdf)

    # ar4_arm.urdf 自带 <link name="world"> + world_joint，Drake 直接把它当世界刚体；
    # 若换成不含 world 的 URDF（如官方原始 xacro 展开件），必须手动把 base 焊到世界，
    # 否则整臂会自由落体。
    # 注意：不能用 GetBodyByName("world") 判断 —— Drake 的 world body 恒名为 "world"，
    # 永远查得到，会漏掉焊接这一步。
    try:
        plant.GetJointByName("world_joint")
    except RuntimeError:
        plant.WeldFrames(plant.world_frame(), plant.GetBodyByName("base").body_frame())

    # URDF 不自动创建执行器（revolute 关节默认被动），需显式 AddJointActuator
    # 并声明扭矩上限 effort，关节才变成“可控关节”。
    for n, ef in zip(JOINT_NAMES, EFFORT):
        plant.AddJointActuator(n, plant.GetJointByName(n), effort_limit=ef)
    plant.Finalize()

    if args.no_gravity:
        plant.mutable_gravity_field().set_gravity_vector([0.0, 0.0, 0.0])

    if args.meshcat:
        # 关掉「自碰撞」：URDF 的碰撞体直接用的 3D 打印件实体网格，相邻连杆
        # 在关节处本来就是互相咬合的，接了 SceneGraph 后 Drake 会算出持续的
        # 自碰撞接触力（实测把 J4 顶到 76% 饱和并颤振）。无头分支没有
        # SceneGraph、压根不做接触计算，所以只有 --meshcat 才看得到这个伪力。
        # 本教程场景没有地面也没有被抓物，自碰撞纯属网格咬合，整体排除即可。
        mi = plant.GetBodyByName("tool_link").model_instance()
        geom_set = GeometrySet()
        for b_idx in plant.GetBodyIndices(mi):
            for g_id in plant.GetCollisionGeometriesForBody(plant.get_body(b_idx)):
                geom_set.Add(g_id)
        decl = CollisionFilterDeclaration()
        decl.ExcludeWithin(geom_set)
        scene_graph.collision_filter_manager().Apply(decl)

        meshcat = Meshcat()
        MeshcatVisualizer.AddToBuilder(builder, scene_graph, meshcat)
        diagram = builder.Build()
        sim = Simulator(diagram)
        root_ctx = sim.get_mutable_context()
        ctx = plant.GetMyContextFromRoot(root_ctx)
    else:
        sim = Simulator(plant)
        ctx = sim.get_mutable_context()
        root_ctx = ctx

    joints = [plant.GetJointByName(n) for n in JOINT_NAMES]
    actuators = [plant.GetJointActuatorByName(n) for n in JOINT_NAMES]
    effort = np.array([a.effort_limit() for a in actuators])
    vel_idx = [j.velocity_start() for j in joints]
    tool = plant.GetBodyByName("tool_link")

    # ---------- 2) 自动整定 PD 增益 ----------
    inertia = effective_inertia(plant, ctx, vel_idx)
    kp_auto, kd_auto = auto_gain(effort, inertia, args.dt)
    kp_arr = np.full(6, args.kp) if args.kp else kp_auto
    kd_arr = np.full(6, args.kd) if args.kd else kd_auto
    print("[ar4] plant=%d dof · gravity=%.2f m/s²"
          % (plant.num_positions(), plant.gravity_field().gravity_vector()[2]))
    print("[ar4] 扭矩上限(N·m) = %s" % np.array2string(effort, precision=1))
    print("[ar4] 有效惯量(kg·m²) = %s" % np.array2string(inertia, precision=5))
    print("[ar4] 整定 kp = %s" % np.array2string(kp_arr, precision=1))
    print("[ar4] 整定 kd = %s" % np.array2string(kd_arr, precision=2))

    Q = load_traj(args.q_traj) if args.q_traj else None
    if Q is not None:
        print("[ar4] 回放 q_traj: %d 帧 → 线性重采样到 %.1f s（原周期 %.1f s）"
              % (len(Q), args.duration, args.traj_period))

    u_port = plant.get_actuation_input_port()
    u = np.zeros(u_port.size())
    u_port.FixValue(ctx, u)

    # ---------- 3) 仿真主循环：PD + 速度前馈 + 重力补偿 ----------
    sim.Initialize()
    t = 0.0
    next_print = 0.0
    err_max = np.zeros(6)
    sat_steps = np.zeros(6)      # 各关节扭矩饱和步数（堵转/丢步风险）
    rows = []
    if args.debug:
        print("t  | q1..q6 | cmd1..6 | err1..6 | tau_g1..6 | u1..6")
    while t < args.duration - 1e-12:
        q = np.array([joints[i].get_angle(ctx) for i in range(6)])   # 实测角度
        qd = plant.GetVelocities(ctx)[vel_idx]                       # 实测角速度
        if Q is not None:
            x = np.clip(t / args.traj_period * (len(Q) - 1), 0, len(Q) - 1)
            i0 = int(np.floor(x))
            i1 = min(i0 + 1, len(Q) - 1)
            f = x - i0
            cmd = Q[i0] * (1.0 - f) + Q[i1] * f
            seg_dt = args.traj_period / max(len(Q) - 1, 1)   # 相邻帧的时间间隔
            cmd_dot = (Q[i1] - Q[i0]) / seg_dt               # 帧间差分速度前馈
        else:
            cmd = np.array([cmd_angle(i, t, args.freq, args.amp_deg) for i in range(6)])
            cmd_dot = np.array([cmd_angle_dot(i, t, args.freq, args.amp_deg)
                                for i in range(6)])

        # 控制器 = PD(误差) + 速度前馈 − 重力广义力前馈。
        # CalcGravityGeneralizedForces 返回重力把关节拉向下垂方向的广义力，
        # 前馈取负号才能托住自重；真实机械臂上位机的重力补偿同理。
        tau_g = plant.CalcGravityGeneralizedForces(ctx)[vel_idx]
        u[:] = kp_arr * (cmd - q) + kd_arr * (cmd_dot - qd)
        if not args.no_ff:
            u -= tau_g
        if not np.isfinite(u).all():
            print("NaN! t=%.4f\n q=%s\n qd=%s\n cmd=%s\n tau_g=%s" % (t, q, qd, cmd, tau_g))
            break
        u[:] = np.clip(u, -effort, effort)
        sat_steps += (np.abs(u) >= effort - 1e-9).astype(float)
        u_port.FixValue(ctx, u)
        sim.AdvanceTo(t + args.dt)
        t += args.dt

        if t >= next_print:
            next_print += args.print_every
            if t >= 2.0:                 # 跳过启动瞬态，统计稳态跟踪误差
                err_max = np.maximum(err_max, np.abs(cmd - q))
            xyz = plant.EvalBodyPoseInWorld(ctx, tool).translation()
            rows.append([t] + list(q) + list(xyz))
            if abs(t - round(t)) < 1e-9 or t - next_print + args.print_every < 1e-9:
                print("t=%.1fs 实测(rad): %s  末端XYZ(m): %.3f %.3f %.3f"
                      % (t, " ".join("%8.3f" % v for v in q), *xyz))
                if args.debug:
                    print("   cmd : %s\n   err : %s\n   qd  : %s\n   tau_g: %s\n   u   : %s"
                          % (" ".join("%7.3f" % v for v in cmd),
                             " ".join("%7.3f" % v for v in cmd - q),
                             " ".join("%7.3f" % v for v in qd),
                             " ".join("%7.3f" % v for v in tau_g),
                             " ".join("%7.3f" % v for v in u)))

    # ---------- 4) 汇总 ----------
    with open(args.csv, "w") as f:
        f.write("t,q1,q2,q3,q4,q5,q6,x,y,z\n")
        for r in rows:
            f.write(",".join("%.6f" % v for v in r) + "\n")
    n_steps = max(int(args.duration / args.dt), 1)
    print("[summary] 最大跟踪误差(rad): %s" % " ".join("%.3f" % v for v in err_max))
    print("[summary] 扭矩饱和占比(%%):  %s"
          % " ".join("%.0f" % (100 * s / n_steps) for s in sat_steps))
    if rows:
        xyz_a = np.array([r[7:] for r in rows])
        print("[summary] 末端轨迹范围(m): x[%.3f, %.3f] y[%.3f, %.3f] z[%.3f, %.3f]"
              % (xyz_a[:, 0].min(), xyz_a[:, 0].max(), xyz_a[:, 1].min(),
                 xyz_a[:, 1].max(), xyz_a[:, 2].min(), xyz_a[:, 2].max()))
    print("[done] 轨迹已保存 -> %s" % args.csv)


if __name__ == "__main__":
    main()
