#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
arm_wave_drake.py · LateAI 站内原创教学代码
在 MIT Drake (pydrake) 中「真实模拟」SmallRobotArm 六轴机械臂：
  1) 用 Parser 加载 smallrobot_arm.urdf（含质量/惯量/关节限位/阻尼）
  2) 6 个关节各配一个 JointActuator，施加 PD 位置控制扭矩
     tau_i = kp*(q_des_i - q_i) - kd*qdot_i   （并受 effort 限位夹取）
  3) 重力场中真实积分动力学方程（含重力/惯性/科氏耦合）
  4) 每 0.5 s 打印位置传感器实测角度（即 plant 状态），并把末端
     全局坐标写入 CSV —— 等价于真机编码器回读 + 轨迹记录
用法:
  python3 arm_wave_drake.py                  # 无头数值验证（默认）
  python3 arm_wave_drake.py --meshcat        # 打开浏览器 3D 视图
  python3 arm_wave_drake.py --duration 30 --kp 60 --kd 8
"""
import argparse
import math
import os

import numpy as np

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

HERE = os.path.dirname(os.path.abspath(__file__))
URDF = os.path.join(HERE, "smallrobot_arm.urdf")
JOINT_NAMES = [f"J{i}" for i in range(1, 7)]


def cmd_angle(i, t, freq=0.5, amp_deg=30.0):
    """第 i 个关节 (0-based) 的目标角度(rad)：与站内 Canvas 演示同一公式。

    freq: 摆动圆频率(rad/s)；amp_deg: 摆动幅值(°)。真实电机有扭矩上限，
    角速度/角加速度太大就会“堵转丢步”，教程里可借此讲明限幅的意义。
    """
    return math.radians(amp_deg) * math.sin(freq * t + i * math.pi / 3.0)


# 分关节 PD 增益：小惯量腕部必须用更小增益（离散 200Hz 数字控制的稳定带宽，
# ω·Δt ≈ sqrt(kp/I)·0.005 < 2；腕部惯量 ~1e-4 所以 kp 只能给 ~15 量级）
# 1kHz 数字控制（Δt=0.001s）：真实微控制器/伺服驱动器常见速率。
# 增益匹配：(a) 线性域 δ=effort/kp 取 ~0.04-0.05 rad（误差在此范围内
# 电机在扭矩上限内线性输出，不落入 bang-bang 开关饱和）；
# (b) 带宽稳定判据 ω·Δt ≈ sqrt(kp/I)·0.001 < 1.5（宽松满足）。
# 惯性耦合是真实存在的（快速运动时关节互相牵动），采样越快、
# 线性域越宽，耦合扰动越难累积成误差。
KP_DEFAULT = [200.0, 300.0, 200.0, 120.0, 80.0, 30.0]
KD_DEFAULT = [12.0, 10.0, 4.5, 1.3, 0.4, 0.1]


def main():
    ap = argparse.ArgumentParser(description="SmallRobotArm · MIT Drake 真实动力学模拟")
    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=30.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("--print-every", type=float, default=0.5, help="打印周期(秒)")
    ap.add_argument("--csv", default=os.path.join(HERE, "trail.csv"), help="末端轨迹 CSV 输出")
    ap.add_argument("--meshcat", action="store_true", help="启动浏览器 3D 可视化")
    ap.add_argument("--debug", action="store_true", help="打印命令/误差/重力矩调试行")
    ap.add_argument("--no-ff", action="store_true", help="关闭重力补偿前馈（诊断用）")
    ap.add_argument("--no-gravity", action="store_true", help="关闭重力场（诊断用）")
    ap.add_argument("--static", action="store_true", help="目标恒为 0（诊断用）")
    ap.add_argument("--test-joint", type=int, default=0,
                    help="仅让第 N 个关节(1-6)做 0.3 rad 阶跃，其余恒 0（诊断用）")
    args = ap.parse_args()

    # ---------- 1) 建世界：加载 URDF + 手动挂 6 个执行器 ----------
    # Drake 的 URDF 不自动创建执行器：revolute 关节默认是被动关节，
    # 必须用 AddJointActuator 把它变成“可控关节”（并声明扭矩上限 effort）。
    # 离散时间 plant：与主循环控制刷新率保持一致(1kHz/Δt=1ms)，
    # 模拟真机“微控制器每秒发 1000 次位置命令 + 电机驱动器 1000Hz 更新”的节奏，
    # 比连续积分器更符合真实数字控制，也天然稳定。
    if args.meshcat:
        builder = DiagramBuilder()
        plant, scene_graph = AddMultibodyPlantSceneGraph(builder, time_step=args.dt)
    else:
        plant = MultibodyPlant(time_step=args.dt)
    Parser(plant).AddModels(URDF)
    # 扭矩上限按 NEMA17 步进电机 + 减速箱的真实量级取值：
    # 肩/肘带 ~15:1 减速箱(输出 ~8-12 N·m)，腕部带 ~5:1 减速箱(1.5-5 N·m)
    EFFORT = [8.0, 12.0, 8.0, 5.0, 3.0, 1.5]   # N·m
    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:
        meshcat = Meshcat()
        vis = MeshcatVisualizer(meshcat=meshcat)
        builder.AddSystem(vis)
        builder.Connect(
            scene_graph.get_query_output_port(),
            vis.get_geometry_query_input_port(),
        )
        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()          # plant 独立 => root 即 plant
        root_ctx = ctx

    joints = [plant.GetJointByName(n) for n in JOINT_NAMES]
    actuators = [plant.GetJointActuatorByName(n) for n in JOINT_NAMES]
    effort = [act.effort_limit() for act in actuators]
    tool = plant.GetBodyByName("tool_link")
    kp_arr = np.full(6, args.kp) if args.kp else np.array(KP_DEFAULT)
    kd_arr = np.full(6, args.kd) if args.kd else np.array(KD_DEFAULT)
    u_port = plant.get_actuation_input_port()
    u = np.zeros(u_port.size())
    u_port.FixValue(ctx, u)

    # ---------- 2) 仿真主循环：PD 位置控制 ----------
    # 关节角度/速度用关节 API 直接读（plant 的完整状态向量含 world 自由体，
    # 直接切片不可靠）；每步把 PD 扭矩写进执行器输入，再推进积分。
    vel_idx = [joints[i].velocity_start() for i in range(6)]
    sim.Initialize()
    t = 0.0
    next_print = 0.0
    err_max = np.zeros(6)
    sat_steps = np.zeros(6)   # 各关节扭矩饱和(堵转/丢步风险)步数
    rows = []
    print("[arm_wave] Drake controller started · plant=%d dof · gravity=%.2f"
          % (plant.num_positions(), plant.gravity_field().gravity_vector()[2]))
    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 args.static:
            cmd = np.zeros(6)
        elif args.test_joint > 0:
            cmd = np.zeros(6)
            cmd[args.test_joint - 1] = 0.3
        else:
            cmd = np.array([cmd_angle(i, t, args.freq, args.amp_deg)
                            for i in range(6)])
        cmd_dot = np.zeros(6) if (args.static or args.test_joint) else np.array(
            [math.radians(args.amp_deg) * args.freq
             * math.cos(args.freq * t + i * math.pi / 3.0) for i in range(6)])
        # 控制器 = PD(误差) − 重力广义力前馈。
        # CalcGravityGeneralizedForces 返回“重力把关节拉向下垂方向”的广义力
        # （沿 +q 为正，见 q=0 时 J2≈+7.5 N·m），因此前馈要取负号才能托住自重；
        # 真实机械臂上位机的重力补偿同理：先抵消重力项，PD 只纠正剩余误差。
        tau_g = plant.CalcGravityGeneralizedForces(ctx)[vel_idx]
        # PD + 速度前馈（减小正弦跟踪相位滞后）+ 重力补偿
        u[:] = kp_arr * (cmd - q) + kd_arr * (cmd_dot - qd)
        if not args.no_ff:
            u -= tau_g
        u[:] = np.clip(u, -np.array(effort), np.array(effort))
        sat_steps += (np.abs(u) >= np.array(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:
                line = ("t=%.1fs 实测(rad): %s  末端XYZ(m): %.3f %.3f %.3f"
                        % (t, " ".join("%8.3f" % v for v in q), *xyz))
                print(line)
                if args.debug:
                    print("   cmd : %s\n   err : %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 tau_g),
                             " ".join("%7.3f" % v for v in u)))

    # ---------- 3) 汇总 ----------
    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))
    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()
