#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
arm_meshcat_real.py · LateAI 站内原创教学代码
用【真实 3D 模型】（SolidWorks 导出的 URDF + OBJ/STL 网格）在 MIT Drake 中做真实模拟，
并用 Meshcat 输出可在浏览器里回放的 3D 动画。

与 smallrobot_arm.urdf（长方体示意模型）的区别：
  · 这里是真机 STL/OBJ 网格 → 看到的就是真机外观
  · 执行器由 URDF 自带（motor1..motor5，effort 各不同），不需要手动 AddJointActuator
  · 增益不再手调：脚本先用「动能法」测出每个关节的有效惯量，再按
    kp = min(effort/0.05, 0.5·(1.5/Δt)²·I)、kd = 1.6·sqrt(kp·I) 自动整定

用法:
  python3 arm_meshcat_real.py --html arm_real_3d.html     # 跑并导出 3D 动画（默认）
  python3 arm_meshcat_real.py --no-meshcat --duration 20  # 只出数值与 CSV
"""
import argparse
import math
import os

import numpy as np
from pydrake.multibody.tree import JointActuatorIndex
from pydrake.all import (
    AddMultibodyPlantSceneGraph,
    DiagramBuilder,
    Meshcat,
    MeshcatVisualizer,
    MultibodyPlant,
    Parser,
    Simulator,
)

DEFAULT_URDF = "/Users/brucezhao/mit/smallrobotarm/urdf/smallrobotarm_with_actuator.urdf"
HERE = os.path.dirname(os.path.abspath(__file__))


def cmd_angle(i, t, freq=0.5, amp_deg=25.0):
    """第 i 个关节的目标角度(rad)，相位依次错开。"""
    return math.radians(amp_deg) * math.sin(freq * t + i * math.pi / 3.0)


def effective_inertia(plant, ctx, vel_idx):
    """动能法测各关节有效转动惯量：令 v = e_i，则 T = 0.5·I_i。"""
    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 = np.minimum(np.array(effort) / 0.05,
                    0.5 * (1.5 / dt) ** 2 * np.array(inertia))
    kd = 1.6 * np.sqrt(kp * np.array(inertia))
    return kp, kd


def main():
    ap = argparse.ArgumentParser(description="SmallRobotArm 真实 3D 模型 · Drake + Meshcat 模拟")
    ap.add_argument("--urdf", default=DEFAULT_URDF, help="真实 URDF 路径（含 OBJ/STL 网格）")
    ap.add_argument("--duration", type=float, default=12.0, help="仿真时长(秒)")
    ap.add_argument("--dt", type=float, default=0.001, help="控制刷新周期(秒，1kHz)")
    ap.add_argument("--time-step", type=float, default=0.0,
                    help="plant 时间步；0=连续。真实模型部分连杆惯量极小(1e-6 量级)，"
                         "离散半隐欧拉会发散，故默认连续积分器")
    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("--publish-period", type=float, default=0.02, help="3D 可视化刷新周期(秒)")
    ap.add_argument("--csv", default=os.path.join(HERE, "real_arm_trail.csv"))
    ap.add_argument("--html", default=os.path.join(HERE, "arm_real_3d.html"),
                    help="导出的 meshcat 3D 动画 HTML")
    ap.add_argument("--no-meshcat", action="store_true", help="不启动 3D，只出数值")
    args = ap.parse_args()

    use_vis = not args.no_meshcat
    if use_vis:
        builder = DiagramBuilder()
        plant, scene_graph = AddMultibodyPlantSceneGraph(builder, time_step=args.time_step)
    else:
        plant = MultibodyPlant(time_step=args.time_step)

    Parser(plant).AddModels(args.urdf)
    # 真实 URDF 的 base 是根连杆：必须焊到世界，否则整臂自由落体
    plant.WeldFrames(plant.world_frame(), plant.GetBodyByName("base").body_frame())
    plant.Finalize()

    n = plant.num_actuators()
    acts = [plant.get_joint_actuator(JointActuatorIndex(i)) for i in range(n)]
    joints = [a.joint() for a in acts]
    effort = np.array([a.effort_limit() for a in acts])
    vel_idx = [j.velocity_start() for j in joints]
    tool = plant.GetBodyByName("link5")   # 末端连杆（末端执行器所在）

    if use_vis:
        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

    inertia = effective_inertia(plant, ctx, vel_idx)
    kp_arr, kd_arr = auto_gain(effort, inertia, args.dt)
    print("[real_arm] joints=%d · effort=%s" % (n, np.array2string(effort, precision=1)))
    print("[real_arm] 有效惯量=%s" % np.array2string(inertia, precision=5))
    print("[real_arm] 自动整定 kp=%s" % np.array2string(kp_arr, precision=1))
    print("[real_arm] 自动整定 kd=%s" % np.array2string(kd_arr, precision=2))

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

    sim.Initialize()
    if use_vis:
        meshcat.StartRecording()
    t = 0.0
    next_print = 0.0
    err_max = np.zeros(n)
    sat_steps = np.zeros(n)
    rows = []
    while t < args.duration - 1e-12:
        q = np.array([joints[i].get_angle(ctx) for i in range(n)])
        qd = plant.GetVelocities(ctx)[vel_idx]
        cmd = np.array([cmd_angle(i, t, args.freq, args.amp_deg) for i in range(n)])
        cmd_dot = np.array([math.radians(args.amp_deg) * args.freq
                            * math.cos(args.freq * t + i * math.pi / 3.0)
                            for i in range(n)])
        tau_g = plant.CalcGravityGeneralizedForces(ctx)[vel_idx]
        u[:] = kp_arr * (cmd - q) + kd_arr * (cmd_dot - qd) - 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，必须用 plant 的 context
        sim.AdvanceTo(t + args.dt)
        t += args.dt
        if t >= next_print:
            next_print += 0.5
            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:
                print("t=%.1fs 实测(rad): %s  末端XYZ(m): %.3f %.3f %.3f"
                      % (t, " ".join("%8.3f" % v for v in q), *xyz))

    if use_vis:
        meshcat.StopRecording()
        meshcat.PublishRecording()
        html = meshcat.StaticHtml()
        with open(args.html, "w") as f:
            f.write(html)
        print("[real_arm] 3D 动画已导出 -> %s (%.1f MB)" % (args.html, len(html) / 1e6))

    with open(args.csv, "w") as f:
        f.write("t," + ",".join("q%d" % (i + 1) for i in range(n)) + ",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[1 + n:] 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()
