#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
arm_wave_drake_faze4.py · LateAI 站内原创教学代码
在 MIT Drake (pydrake) 中「真实模拟」Faze4 六轴机械臂
（全 3D 打印、关节内置摆线针轮减速箱）：

  1) Parser 加载 faze4_arm.urdf（关节 origin/rpy/axis 逐字段取自官方
     Source-Robotics/Faze4-Robotic-arm 的 URDF_FAZE4/urdf/Final_light_assembly_URDF.urdf；
     官方六个关节均为 continuous，URDF 中的限位为本站按机械结构取的估算值，见文件内注释）
  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_faze4.py                        # 无头数值验证（默认）
  python3 arm_wave_drake_faze4.py --meshcat              # 打开浏览器 3D 视图
  python3 arm_wave_drake_faze4.py --q-traj q_traj.npy    # 回放网页导出的关节轨迹
  python3 arm_wave_drake_faze4.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, "faze4_arm.urdf")
# 官方 URDF 的关节名是克罗地亚语（Faze4 项目来自克罗地亚），依次对应 J1~J6：
#   回转底座 / 大臂 / 肘 / 前臂滚转 / 腕俯仰 / 末端滚转
JOINT_NAMES = ["rotary_base", "nadlaktica", "lakat", "podlaktica", "saka", "hvataljka"]
JOINT_LABELS = ["J1", "J2", "J3", "J4", "J5", "J6"]

# 关节限位（度）：官方 URDF 六个关节全为 continuous（无限位），
# 这里按机械结构取合理限位，仅用于打印提示与目标角度的安全裁剪。
LIMIT_DEG = [(-180, 180), (-120, 120), (-150, 150), (-180, 180), (-120, 120), (-180, 180)]

# 扭矩上限（N·m）：Faze4 六个关节均为 NEMA17 步进电机 + 内置摆线针轮减速箱，
# 由「保持扭矩 × 减速比 × 传动效率」估算（NEMA17 保持扭矩约 0.45 N·m，效率取 0.7）。
# 真机标定后请替换为实测值。
EFFORT = [7.9, 8.5, 4.7, 3.7, 3.5, 6.0]

# ---------------------------------------------------------------------------
# Faze4 的「惯量鸿沟」与默认步长/整定式的由来（重要，改参数前请先读）
#
# Faze4 的腕部末端（hvataljka，J6）只带动一个约 20 g 的夹爪，关节侧连杆
# 惯量仅约 6.4e-6 kg·m²；而 J2 带着整条大臂+前臂，惯量 0.4 kg·m²，高出
# 四五个量级。把离散 PD 控制的无量纲刚度记为
#       b_i = Δt² · kp_i / I_i        （I_i 为该关节有效惯量）
# 实测（本机 pydrake，--no-gravity 排除干扰，逐个 kp 扫描看 J6 速度）：
#   b = 0.44 → 干净，J6 角速度 ≈ 0.19 rad/s（跟随 0.5 rad/s 正弦的目标值）
#   b = 0.50 → 起振：J6 以 ±256 rad/s 的步频极限环颤振
# 即判别式约为 b < 0.5；据此把整定式收紧到 kp_i ≤ 0.4·I_i/Δt²（20% 余量），
# 见 auto_gain。这个颤振极其阴险：位置摆幅只有 ~0.05 rad，0.5 s 打印一次
# 时「最大跟踪误差 0.036 rad、看着还行」，只有 --debug 打印 qd 才看得见
# 256 rad/s —— 所以 auto_gain 里 kp 的上限必须按这条离散判据卡，而不是按
# 「误差 0.05 rad 不饱和」的静态思路。
#
# 为什么默认 Δt = 0.2 ms：把 kp 上限与 Δt 挂钩后，
#   · 1 kHz（Δt=1 ms）：J6 的 kp 上限只有 6.4e-6/1e-6·0.4 ≈ 2.5，太软，
#     一有负载误差就顶到 6 N·m 而发散成 NaN；
#   · 0.2 ms（5 kHz）：J6 的 kp 上限升到 63.6，配合 kd = I/(2Δt) 的阻尼，
#     六个关节全部稳跟 0.5 rad/s 摆动，扭矩饱和占比 0%。
# 真实 Faze4 固件也跑在数 kHz 控制环上，故 5 kHz 既贴近真机又数值稳定。
# 换更轻/更重的末端时请优先调 --dt（更小更稳、更慢），再考虑 --kp/--kd。
#
# （早期尝试：把 NEMA17 转子惯量 J_rotor·N² 折算进腕部连杆，用减速比平方
#  把 J6 惯量抬到 2e-3 kg·m² 以救活 1 kHz。但 Drake 会对刚体惯量做
#  CouldBePhysicallyValid() 三角不等式校验，手动叠加的对角惯量破坏了惯量
#  张量的物理可行性，被 Drake 替换成更病态的值，反而更不稳定。最终采用
#  「保留官方 URDF 惯量 + 收紧 kp 上限 + 缩小步长」这一更干净的方案。）
# ---------------------------------------------------------------------------


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

    零位取限位区间中点，再叠加相位依次错开的正弦摆动。
    Faze4 六个关节的限位都关于 0 对称，故零位即 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,  KP_B_MAX · I / Δt²)
      前项保证 0.05 rad 误差内电机不饱和（线性工作区）；
      后项是离散时间 PD 的稳定上限：无量纲刚度 b = Δt²·kp/I 必须留裕量
      （实测 b ≳ 0.5 即出现步频极限环，见文件顶部说明），取 KP_B_MAX=0.4。
    kd = min(1.6·sqrt(kp·I),  I/(2Δt))
      前项是临界阻尼略偏大的系数，用于抑制摆动超调；
      后项是离散化的硬约束：微分时间常数 τ = I/kd 必须 ≥ 2Δt，
      否则一步之内速度反馈就能把力矩顶到上限，形成正负交替的高频极限环。

    注意前项其实很少起作用 —— 对 J6 这种惯量极小的腕部轴，I/(2Δt) 总是更小
    的那个，因此 kd 恒等于 I/(2Δt)（即 a = Δt·kd/I = 0.5 恰好落在稳定边界
    内侧）。真正决定「能不能跑」的是 kp 这一项：J2 的 I 是 J6 的 6 万倍，
    同样 Δt 下 J6 的 kp 上限被压到 J2 的六万分之一，这正是 Faze4 必须把
    --dt 取到 0.2 ms 的根本原因。
    """
    I = np.array(inertia)
    # KP_B_MAX：无量纲刚度 b = Δt²·kp/I 的上限（含 20% 稳定裕量）
    KP_B_MAX = 0.4
    kp = np.minimum(np.array(effort) / 0.05, KP_B_MAX * I / dt ** 2)
    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="Faze4 六轴机械臂 · MIT Drake 真实动力学模拟")
    ap.add_argument("--urdf", default=URDF, help="URDF 路径（默认同目录 faze4_arm.urdf）")
    ap.add_argument("--duration", type=float, default=14.0, help="仿真时长(秒)")
    ap.add_argument("--dt", type=float, default=0.0002,
                    help="控制刷新周期(秒)，默认 0.2 ms / 5 kHz —— 原因见文件顶部《惯量鸿沟》说明")
    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, "faze4_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：刷新率与主循环一致（默认 5 kHz），模拟真机
    #「控制器以固定周期下发位置命令 + 驱动器同步更新」的节奏。
    # 步长越小越稳，但也越慢；Faze4 腕部惯量极小，故默认取 0.2 ms，见文件顶部说明。
    if args.meshcat:
        builder = DiagramBuilder()
        plant, scene_graph = AddMultibodyPlantSceneGraph(builder, time_step=args.dt)
    else:
        plant = MultibodyPlant(time_step=args.dt)

    # 直接加载官方 URDF：惯量保持原样（不做转子惯量折算），
    # 由「缩小控制周期」而非「篡改惯量」来保证腕部数值稳定。
    Parser(plant).AddModels(args.urdf)

    # faze4_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_link").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:
        # 关掉「自碰撞」——这一步对 Faze4 是必需的，不是可选优化：
        # Faze4 的 URDF 碰撞体直接用的 3D 打印件实体网格，相邻连杆在关节处
        # 本来就是互相咬合的（减速箱/轴承座嵌进对方腔体里），所以只要接了
        # SceneGraph 就会算出持续的自碰撞接触力：实测初始位姿下就有 1 个接触点，
        # 把 3.7 N·m 的 J4 顶到 95% 饱和并颤振（无头分支没有 SceneGraph、
        # 压根不做接触计算，所以看不到这个现象——这正是两个分支结果不一致的原因）。
        # 本教程场景里没有地面也没有被抓物，自碰撞纯属网格咬合造成的伪力，
        # 用官方推荐的做法整体排除即可（真机做碰撞检测时才会保留它）。
        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("[faze4] plant=%d dof · gravity=%.2f m/s²"
          % (plant.num_positions(), plant.gravity_field().gravity_vector()[2]))
    print("[faze4] 扭矩上限(N·m) = %s" % np.array2string(effort, precision=1))
    print("[faze4] 有效惯量(kg·m²) = %s" % np.array2string(inertia, precision=5))
    print("[faze4] 整定 kp = %s" % np.array2string(kp_arr, precision=1))
    print("[faze4] 整定 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("[faze4] 回放 q_traj: %d 帧 → 线性重采样到 %.1f s（原周期 %.1f s）"
              % (len(Q), args.duration, args.traj_period))

    # 初始状态直接放在命令轨迹的起点上，避免 t=0 出现阶跃激励。
    # 这一步对 Faze4 是必须的：Faze4 六个关节的限位都关于 0 对称，
    # 零位处的 sin 相位并不为 0，若从全零姿态起步，腕部（有效惯量极小）
    # 会在第一个控制周期吃到满幅误差而瞬间发散。
    q_ini = (np.array(Q[0], dtype=float) if Q is not None
             else np.array([cmd_angle(i, 0.0, args.freq, args.amp_deg) for i in range(6)]))
    plant.SetPositions(ctx, q_ini)
    plant.SetVelocities(ctx, np.zeros(plant.num_velocities()))
    print("[faze4] 初始位姿(度) = %s" % np.array2string(np.degrees(q_ini), precision=1))

    u_port = plant.get_actuation_input_port()
    u = np.zeros(u_port.size())
    u_port.FixValue(ctx, u)      # 端口属于 plant，必须用 plant 的 context

    # ---------- 3) 仿真主循环：PD + 速度前馈 + 重力补偿 ----------
    sim.Initialize()
    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 True:
        # 时刻一律读仿真器自己的时钟，而不是自己把 dt 累加几十万次。
        # 默认 0.2 ms 步长、14 s 就是 7 万次累加，浮点漂移会很快超过 1e-9：
        # 一来会漏打印（见下面 next_print 的注释），二来请求的步进时刻会
        # 逐渐偏离 plant 的离散更新网格。读 root_ctx.get_time() 则每次
        # 步进都精确落在仿真器的时钟上。
        t = root_ctx.get_time()
        if t >= args.duration - 1e-12:
            break
        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)      # 同上：始终喂给 plant 的 context
        sim.AdvanceTo(t + args.dt)
        t = root_ctx.get_time()      # 步进后的真实时刻（下轮循环开头会重读）

        if t >= next_print:
            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))
            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)))
            # 用「本次时刻 + 周期」而非累加，避免长时间仿真积累浮点漂移
            # （默认 0.2 ms 步长下 t 的累加误差会超过 1e-9，累加式判据会漏打印）
            next_print = t + args.print_every

    # ---------- 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()
