#!/usr/bin/env python

# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from __future__ import annotations

import importlib
import logging
import threading
import time
from dataclasses import dataclass, field
from functools import cached_property
from typing import TYPE_CHECKING, Protocol, runtime_checkable

import numpy as np

from lerobot.cameras import DepthCamera, make_cameras_from_configs
from lerobot.lerobot_types import RobotAction, RobotObservation
from lerobot.utils.import_utils import _unitree_sdk_available, require_package

from ..robot import Robot
from .config_unitree_g1 import UnitreeG1Config
from .g1_kinematics import G1_29_ArmIK
from .g1_utils import (
    NUM_MOTORS,
    REMOTE_AXES,
    G1_29_JointArmIndex,
    G1_29_JointIndex,
    default_remote_input,
)

if TYPE_CHECKING or _unitree_sdk_available:
    from unitree_sdk2py.core.channel import (
        ChannelFactoryInitialize as _SDKChannelFactoryInitialize,
        ChannelPublisher as _SDKChannelPublisher,
        ChannelSubscriber as _SDKChannelSubscriber,
    )
    from unitree_sdk2py.idl.default import unitree_hg_msg_dds__LowCmd_
    from unitree_sdk2py.idl.unitree_hg.msg.dds_ import (
        LowCmd_ as hg_LowCmd,
        LowState_ as hg_LowState,
    )
    from unitree_sdk2py.utils.crc import CRC
else:
    _SDKChannelFactoryInitialize = None
    _SDKChannelPublisher = None
    _SDKChannelSubscriber = None
    unitree_hg_msg_dds__LowCmd_ = None
    hg_LowCmd = None
    hg_LowState = None
    CRC = None

if TYPE_CHECKING:
    import gymnasium as gym

logger = logging.getLogger(__name__)


@runtime_checkable
class RobotController(Protocol):
    """Interface for the controllers driving ``UnitreeG1``'s background control thread.

    Covers both locomotion controllers (GR00T, Holosoma) and whole-body ones (SONIC).

    Each tick the robot hands the controller the latest lowstate plus a snapshot of the
    incoming action, and publishes the absolute joint targets it returns, keyed
    ``<joint>.q``. It lives here rather than in ``controllers/`` so that importing the
    robot does not pull in the controller implementations and their onnxruntime dependency.

    Controllers may also expose any of the following, which the robot picks up when present:

    - ``kp`` / ``kd``: ``(29,)`` PD gains published with the targets, overriding the config.
    - ``default_angles``: ``(29,)`` home pose that residual actions are applied onto.
    - ``action_ft`` / ``observation_ft``: feature dicts that take over the robot's default
      29-DoF action space and proprioceptive state (SONIC's 64-D latent token).
    - ``observation_state()``: current values for the keys advertised in ``observation_ft``.
    """

    control_dt: float
    """Control period in seconds; sets the rate of the robot's controller thread."""

    def run_step(self, action: dict, lowstate) -> dict:
        """Map one lowstate plus action into absolute joint targets keyed ``<joint>.q``."""
        ...

    def reset(self) -> None:
        """Drop per-episode state, e.g. history buffers and held commands."""
        ...


def make_robot_controller(name: str | None) -> RobotController | None:
    """Instantiate a robot controller by class name. Returns None if name is None."""
    if name is None:
        return None
    controllers = {
        "GrootLocomotionController": "lerobot.robots.unitree_g1.controllers.gr00t_locomotion",
        "HolosomaLocomotionController": "lerobot.robots.unitree_g1.controllers.holosoma_locomotion",
        "SonicWholeBodyController": "lerobot.robots.unitree_g1.controllers.sonic_whole_body",
    }
    module_path = controllers.get(name)
    if module_path is None:
        raise ValueError(f"Unknown controller: {name!r}. Available: {list(controllers)}")
    module = importlib.import_module(module_path)
    return getattr(module, name)()


# DDS topic names follow Unitree SDK naming conventions
# ruff: noqa: N816
kTopicLowCommand_Debug = "rt/lowcmd"
kTopicLowState = "rt/lowstate"


@dataclass
class MotorState:
    q: float | None = None  # position
    dq: float | None = None  # velocity
    tau_est: float | None = None  # estimated torque
    temperature: float | None = None  # motor temperature


@dataclass
class IMUState:
    quaternion: np.ndarray | None = None  # [w, x, y, z]
    gyroscope: np.ndarray | None = None  # [x, y, z] angular velocity (rad/s)
    accelerometer: np.ndarray | None = None  # [x, y, z] linear acceleration (m/s²)
    rpy: np.ndarray | None = None  # [roll, pitch, yaw] (rad)
    temperature: float | None = None  # IMU temperature


# g1 observation class
@dataclass
class G1_29_LowState:  # noqa: N801
    motor_state: list[MotorState] = field(default_factory=lambda: [MotorState() for _ in G1_29_JointIndex])
    imu_state: IMUState = field(default_factory=IMUState)
    wireless_remote: bytes | None = None  # Raw wireless remote data
    mode_machine: int = 0  # Robot mode


class UnitreeG1(Robot):
    config_class = UnitreeG1Config
    name = "unitree_g1"

    def __init__(self, config: UnitreeG1Config) -> None:
        require_package("unitree-sdk2py", extra="unitree_g1", import_name="unitree_sdk2py")
        super().__init__(config)

        logger.info("Initialize UnitreeG1...")

        self.config = config
        self.control_dt = config.control_dt

        # Initialize cameras config (ZMQ-based) - actual connection in connect()
        self._cameras = make_cameras_from_configs(config.cameras)

        # Import channel classes based on mode
        if config.is_simulation:
            self._ChannelFactoryInitialize = _SDKChannelFactoryInitialize
            self._ChannelPublisher = _SDKChannelPublisher
            self._ChannelSubscriber = _SDKChannelSubscriber
        else:
            from .unitree_sdk2_socket import (
                ChannelFactoryInitialize,
                ChannelPublisher,
                ChannelSubscriber,
            )

            self._ChannelFactoryInitialize = ChannelFactoryInitialize
            self._ChannelPublisher = ChannelPublisher
            self._ChannelSubscriber = ChannelSubscriber

        # Initialize state variables
        self.sim_env = None
        self._env_wrapper: dict[str, dict[int, gym.vector.VectorEnv]] | None = None
        self._lowstate = None
        self._lowstate_lock = threading.Lock()
        # Guards the shared lowcmd message: the controller thread, send_action(), reset() and
        # the shutdown path all publish through it, and a torn update still carries a valid CRC.
        self._lowcmd_lock = threading.Lock()
        # Decides who may drive the joints over a span of time: one controller tick, or a whole
        # reset sweep. Coarser than _lowcmd_lock, which only makes a single command atomic.
        self._control_lock = threading.Lock()
        self._shutdown_event = threading.Event()
        self.subscribe_thread: threading.Thread | None = None

        self.arm_ik = G1_29_ArmIK() if config.gravity_compensation else None

        # Controller loaded dynamically
        self.controller: RobotController | None = make_robot_controller(config.controller)
        # Controller thread state
        self._controller_thread: threading.Thread | None = None
        self._controller_action_lock = threading.Lock()
        self.controller_input = default_remote_input()
        self.controller_output: RobotAction = {}
        # PD gains, resolved in connect() from the controller when it provides them, else from the config.
        self.kp: np.ndarray
        self.kd: np.ndarray

    def _subscribe_lowstate(self):  # polls robot state @ 250Hz
        while not self._shutdown_event.is_set():
            start_time = time.time()

            # Step simulation if in simulation mode
            if self.config.is_simulation and self.sim_env is not None:
                self.sim_env.step()

            msg = self.lowstate_subscriber.Read()
            if msg is not None:
                lowstate = G1_29_LowState()

                # Capture motor states using jointindex
                for joint in G1_29_JointIndex:
                    lowstate.motor_state[joint].q = msg.motor_state[joint].q
                    lowstate.motor_state[joint].dq = msg.motor_state[joint].dq
                    lowstate.motor_state[joint].tau_est = msg.motor_state[joint].tau_est
                    lowstate.motor_state[joint].temperature = msg.motor_state[joint].temperature

                # Capture IMU state
                lowstate.imu_state.quaternion = list(msg.imu_state.quaternion)
                lowstate.imu_state.gyroscope = list(msg.imu_state.gyroscope)
                lowstate.imu_state.accelerometer = list(msg.imu_state.accelerometer)
                lowstate.imu_state.rpy = list(msg.imu_state.rpy)
                lowstate.imu_state.temperature = msg.imu_state.temperature

                # Capture wireless remote data
                lowstate.wireless_remote = msg.wireless_remote

                # Capture mode_machine
                lowstate.mode_machine = msg.mode_machine

                with self._lowstate_lock:
                    self._lowstate = lowstate

            current_time = time.time()
            all_t_elapsed = current_time - start_time
            sleep_time = max(0, (self.control_dt - all_t_elapsed))  # maintain constant control dt
            time.sleep(sleep_time)

    def publish_lowcmd(
        self,
        action: RobotAction,
        kp: np.ndarray | list[float] | None = None,
        kd: np.ndarray | list[float] | None = None,
        tau: np.ndarray | list[float] | None = None,
    ) -> None:  # writes robot command whenever requested
        with self._lowcmd_lock:
            for motor in G1_29_JointIndex:
                key = f"{motor.name}.q"
                if key in action:
                    self.msg.motor_cmd[motor.value].q = action[key]
                    self.msg.motor_cmd[motor.value].qd = 0
                    self.msg.motor_cmd[motor.value].kp = (
                        kp[motor.value] if kp is not None else self.kp[motor.value]
                    )
                    self.msg.motor_cmd[motor.value].kd = (
                        kd[motor.value] if kd is not None else self.kd[motor.value]
                    )
                    self.msg.motor_cmd[motor.value].tau = tau[motor.value] if tau is not None else 0.0

            self.msg.crc = self.crc.Crc(self.msg)
            self.lowcmd_publisher.Write(self.msg)

    @property
    def _cameras_ft(self) -> dict[str, tuple]:
        features: dict[str, tuple] = {}
        for cam in self.cameras:
            cfg = self.config.cameras[cam]
            if getattr(cfg, "use_rgb", True):
                features[cam] = (cfg.height, cfg.width, 3)
            if getattr(cfg, "use_depth", False):
                features[f"{cam}_depth"] = (cfg.height, cfg.width, 1)
        return features

    @cached_property
    def observation_features(self) -> dict[str, type | tuple]:
        # A controller advertising its own proprio state (SONIC's 64-D token echo) replaces the
        # raw joint positions rather than extending them, the way action_features hands the
        # action space over to the controller.
        controller_ft = getattr(self.controller, "observation_ft", None)
        proprio_ft = self._motors_ft if controller_ft is None else dict(controller_ft)
        return {**proprio_ft, **self._cameras_ft}

    @cached_property
    def action_features(self) -> dict[str, type]:
        # No controller configured at all: raw 29-DoF joint teleop.
        if self.controller is None:
            return {f"{G1_29_JointIndex(motor).name}.q": float for motor in G1_29_JointIndex}

        # Whole-body controllers (SONIC): 64-D latent token.
        controller_ft = getattr(self.controller, "action_ft", None)
        if controller_ft is not None:
            return dict(controller_ft)

        # Locomotion controllers (GR00T / Holosoma): arm joint targets + joystick axes.
        # TODO: have GR00T/Holosoma advertise their own action_features too, so every
        # controller declares its action space and this fallthrough can be dropped.
        arm_features = {f"{G1_29_JointArmIndex(motor).name}.q": float for motor in G1_29_JointArmIndex}
        remote_features = dict.fromkeys(REMOTE_AXES, float)
        return {**arm_features, **remote_features}

    def _controller_loop(self):
        """Background thread that runs controller at policy's control_dt."""
        control_dt = self.controller.control_dt
        logger.info(f"Controller loop starting with control_dt={control_dt} ({1.0 / control_dt:.1f}Hz)")

        loop_count = 0
        last_log_time = time.time()

        while not self._shutdown_event.is_set():
            start_time = time.time()

            with self._lowstate_lock:
                lowstate = self._lowstate

            if lowstate is not None and self.controller is not None:
                loop_count += 1
                if time.time() - last_log_time >= 5.0:  # Log every 5 seconds
                    actual_hz = loop_count / (time.time() - last_log_time)
                    logger.info(
                        f"Controller actual rate: {actual_hz:.1f}Hz (target: {1.0 / control_dt:.1f}Hz)"
                    )
                    loop_count = 0
                    last_log_time = time.time()
                # Read controller input snapshot
                with self._controller_action_lock:
                    controller_input = dict(self.controller_input)

                # Run controller step and publish it as one turn of control, so a reset sweep
                # cannot interleave its own targets with this tick's.
                with self._control_lock:
                    controller_action = self.controller.run_step(controller_input, lowstate)

                    # Write controller output snapshot
                    with self._controller_action_lock:
                        self.controller_output = dict(controller_action)

                    ctrl_kp = self.controller.kp if hasattr(self.controller, "kp") else None
                    ctrl_kd = self.controller.kd if hasattr(self.controller, "kd") else None
                    self.publish_lowcmd(controller_action, kp=ctrl_kp, kd=ctrl_kd)

            elapsed = time.time() - start_time
            sleep_time = max(0, control_dt - elapsed)
            time.sleep(sleep_time)

    def calibrate(self) -> None:
        # TODO: implement g1_29 calibration
        pass

    def configure(self) -> None:
        pass

    def connect(self, calibrate: bool = True) -> None:  # connect to DDS
        # Initialize DDS channel and simulation environment
        if self.config.is_simulation:
            from lerobot.envs import make_env

            self._ChannelFactoryInitialize(0, "lo")
            self._env_wrapper = make_env(self.config.sim_env, trust_remote_code=True)
            # Extract the actual gym env from the dict structure
            self.sim_env = self._env_wrapper["hub_env"][0].envs[0]
        else:
            self._ChannelFactoryInitialize(0, config=self.config)

        # Initialize direct motor control interface
        self.lowcmd_publisher = self._ChannelPublisher(kTopicLowCommand_Debug, hg_LowCmd)
        self.lowcmd_publisher.Init()
        self.lowstate_subscriber = self._ChannelSubscriber(kTopicLowState, hg_LowState)
        self.lowstate_subscriber.Init()

        # Start subscribe thread to read robot state
        self.subscribe_thread = threading.Thread(target=self._subscribe_lowstate)
        self.subscribe_thread.start()

        # Connect cameras
        for cam in self._cameras.values():
            if not cam.is_connected:
                cam.connect()

        logger.info(f"Connected {len(self._cameras)} camera(s).")

        # Initialize lowcmd message
        self.crc = CRC()
        self.msg = unitree_hg_msg_dds__LowCmd_()
        self.msg.mode_pr = 0

        # Wait for first state message to arrive
        lowstate = None
        deadline = time.time() + 10.0
        while lowstate is None:
            with self._lowstate_lock:
                lowstate = self._lowstate
            if lowstate is None:
                if time.time() > deadline:
                    raise TimeoutError("Timed out waiting for robot state (10s)")
                logger.warning("[UnitreeG1] Waiting for robot state...")
                time.sleep(0.01)
        logger.info("[UnitreeG1] Connected to robot.")
        self.msg.mode_machine = lowstate.mode_machine

        # Prefer the active controller's gains (e.g. SONIC loads kp/kd from its ONNX);
        # otherwise fall back to the config defaults.
        if self.controller is not None and hasattr(self.controller, "kp"):
            self.kp = np.array(self.controller.kp, dtype=np.float32)
            self.kd = np.array(self.controller.kd, dtype=np.float32)
        else:
            self.kp = np.array(self.config.kp, dtype=np.float32)
            self.kd = np.array(self.config.kd, dtype=np.float32)

        for joint in G1_29_JointIndex:
            self.msg.motor_cmd[joint].mode = 1
            self.msg.motor_cmd[joint].kp = self.kp[joint.value]
            self.msg.motor_cmd[joint].kd = self.kd[joint.value]
            self.msg.motor_cmd[joint].q = lowstate.motor_state[joint.value].q

        # Ease into the controller's home pose before it takes over, so the first commands
        # don't snap from the connect-time pose. reset() picks that pose up on its own.
        if self.controller is not None and hasattr(self.controller, "default_angles"):
            self.reset()

        # Start controller thread if enabled
        if self.controller is not None:
            self._controller_thread = threading.Thread(target=self._controller_loop, daemon=True)
            self._controller_thread.start()
            fps = int(1.0 / self.controller.control_dt)
            logger.info(f"Controller thread started ({fps}Hz)")

    def _send_zero_torque(self) -> None:
        """Send a zero-gain command to make joints passive before shutting down."""
        try:
            with self._lowstate_lock:
                lowstate = self._lowstate
            if lowstate is None:
                return
            action = {f"{motor.name}.q": lowstate.motor_state[motor.value].q for motor in G1_29_JointIndex}
            zero_gains = np.zeros(29, dtype=np.float32)
            self.publish_lowcmd(action, kp=zero_gains, kd=zero_gains, tau=zero_gains)
            logger.info("Sent zero-torque command for safe shutdown")
        except Exception as e:
            logger.warning(f"Failed to send zero-torque on disconnect: {e}")

    def disconnect(self):
        # Signal threads to stop and unblock any waits
        self._shutdown_event.set()

        # Wait for controller thread to finish. It has to be stopped before going passive,
        # otherwise a tick already in flight re-stiffens the joints and zero torque is not the
        # robot's last command.
        if self._controller_thread is not None:
            self._controller_thread.join(timeout=2.0)
            if self._controller_thread.is_alive():
                logger.warning("Controller thread did not stop cleanly")

        # Put robot in passive mode
        if not self.config.is_simulation:
            self._send_zero_torque()

        # Wait for subscribe thread to finish
        if self.subscribe_thread is not None:
            self.subscribe_thread.join(timeout=2.0)
            if self.subscribe_thread.is_alive():
                logger.warning("Subscribe thread did not stop cleanly")

        # Close simulation environment
        if self.config.is_simulation and self.sim_env is not None:
            try:
                # Force-kill the image publish subprocess first to avoid long waits
                if hasattr(self.sim_env, "simulator") and hasattr(self.sim_env.simulator, "sim_env"):
                    sim_env_inner = self.sim_env.simulator.sim_env
                    if hasattr(sim_env_inner, "image_publish_process"):
                        proc = sim_env_inner.image_publish_process
                        if proc.process and proc.process.is_alive():
                            logger.info("Force-terminating image publish subprocess...")
                            proc.stop_event.set()
                            proc.process.terminate()
                            proc.process.join(timeout=1)
                            if proc.process.is_alive():
                                proc.process.kill()
                self.sim_env.close()
            except Exception as e:
                logger.warning(f"Error closing sim_env: {e}")
            self.sim_env = None
            self._env_wrapper = None

        # Disconnect cameras
        for cam in self._cameras.values():
            cam.disconnect()

    def get_observation(self) -> RobotObservation:
        with self._lowstate_lock:
            lowstate = self._lowstate
        if lowstate is None:
            return {}

        obs = {}

        # Motors - q, dq, tau for all joints
        for motor in G1_29_JointIndex:
            name = motor.name
            idx = motor.value
            obs[f"{name}.q"] = lowstate.motor_state[idx].q
            obs[f"{name}.dq"] = lowstate.motor_state[idx].dq
            obs[f"{name}.tau"] = lowstate.motor_state[idx].tau_est

        # IMU - gyroscope
        if lowstate.imu_state.gyroscope:
            obs["imu.gyro.x"] = lowstate.imu_state.gyroscope[0]
            obs["imu.gyro.y"] = lowstate.imu_state.gyroscope[1]
            obs["imu.gyro.z"] = lowstate.imu_state.gyroscope[2]

        # IMU - accelerometer
        if lowstate.imu_state.accelerometer:
            obs["imu.accel.x"] = lowstate.imu_state.accelerometer[0]
            obs["imu.accel.y"] = lowstate.imu_state.accelerometer[1]
            obs["imu.accel.z"] = lowstate.imu_state.accelerometer[2]

        # IMU - quaternion
        if lowstate.imu_state.quaternion:
            obs["imu.quat.w"] = lowstate.imu_state.quaternion[0]
            obs["imu.quat.x"] = lowstate.imu_state.quaternion[1]
            obs["imu.quat.y"] = lowstate.imu_state.quaternion[2]
            obs["imu.quat.z"] = lowstate.imu_state.quaternion[3]

        # IMU - rpy
        if lowstate.imu_state.rpy:
            obs["imu.rpy.roll"] = lowstate.imu_state.rpy[0]
            obs["imu.rpy.pitch"] = lowstate.imu_state.rpy[1]
            obs["imu.rpy.yaw"] = lowstate.imu_state.rpy[2]

        # Wireless remote (raw bytes for teleoperator)
        if lowstate.wireless_remote:
            obs["wireless_remote"] = lowstate.wireless_remote

        # Controller-contributed observation (e.g. SONIC echoes its last decoded token as
        # observation.state so a token-output VLA closes the loop on its own previous token).
        if self.controller is not None and hasattr(self.controller, "observation_state"):
            obs.update(self.controller.observation_state())

        # Cameras - read images from ZMQ cameras
        for cam_name, cam in self._cameras.items():
            if getattr(cam, "use_rgb", True):
                obs[cam_name] = cam.read_latest()
            if isinstance(cam, DepthCamera) and cam.use_depth:
                obs[f"{cam_name}_depth"] = cam.read_latest_depth()

        return obs

    def send_action(self, action: RobotAction) -> RobotAction:
        action_to_publish = action
        if self.controller is not None:
            # Controller thread owns legs/waist. Here we only update joystick inputs
            # and publish arm targets from the teleoperator.
            self._update_controller_action(action)
            arm_prefixes = tuple(j.name for j in G1_29_JointArmIndex)
            action_to_publish = {
                key: value
                for key, value in action.items()
                if key.endswith(".q") and key.startswith(arm_prefixes)
            }
            if not action_to_publish:
                # Nothing here for the arms, so publishing would only re-send the controller
                # thread's own last command with a fresh CRC, at the caller's rate on top of the
                # controller's. Token-only actions (a SONIC policy) hit this on every step.
                return action

        tau: np.ndarray | None = None
        if self.config.gravity_compensation and self.arm_ik is not None:
            tau = np.zeros(29, dtype=np.float32)
            action_np = np.array(
                [
                    action_to_publish.get(f"{joint.name}.q", self.msg.motor_cmd[joint.value].q)
                    for joint in G1_29_JointArmIndex
                ],
                dtype=np.float32,
            )
            arm_tau = self.arm_ik.solve_tau(action_np)
            arm_start_idx = G1_29_JointArmIndex.kLeftShoulderPitch.value
            for joint in G1_29_JointArmIndex:
                local_idx = joint.value - arm_start_idx
                tau[joint.value] = arm_tau[local_idx]

        self.publish_lowcmd(action_to_publish, tau=tau)
        return action

    def _update_controller_action(self, action: RobotAction) -> None:
        """Forward incoming teleop action values into ``controller_input``; each controller
        reads only the keys it understands."""
        with self._controller_action_lock:
            for key, value in action.items():
                if isinstance(key, str) and value is not None:
                    self.controller_input[key] = value

    @property
    def is_calibrated(self) -> bool:
        return True

    @property
    def is_connected(self) -> bool:
        with self._lowstate_lock:
            return self._lowstate is not None

    @property
    def _motors_ft(self) -> dict[str, type]:
        """Joint positions for all 29 joints."""
        return {f"{G1_29_JointIndex(motor).name}.q": float for motor in G1_29_JointIndex}

    @property
    def cameras(self) -> dict:
        return self._cameras

    def reset(
        self,
        control_dt: float | None = None,
        default_positions: list[float] | None = None,
    ) -> None:  # move robot to default position
        if control_dt is None:
            control_dt = self.config.control_dt
        if default_positions is None:
            # Home to the controller's own pose when it has one: that is what its policy was
            # trained around (SONIC reads it from the ONNX metadata), whereas the config default
            # is a generic fallback for raw joint teleop.
            controller_home = getattr(self.controller, "default_angles", None)
            source = self.config.default_positions if controller_home is None else controller_home
            default_positions = np.array(source, dtype=np.float32)

        # Hold control authority for the whole sweep. Otherwise the controller thread keeps
        # publishing its own targets throughout, and the robot is driven by two writers at once.
        with self._control_lock:
            if self.config.is_simulation and self.sim_env is not None:
                self.sim_env.reset()
                self.publish_lowcmd(
                    {f"{motor.name}.q": float(default_positions[motor.value]) for motor in G1_29_JointIndex}
                )
            else:
                total_time = 3.0
                num_steps = int(total_time / control_dt)

                # get current state
                obs = self.get_observation()

                # record current positions
                init_dof_pos: np.ndarray = np.zeros(NUM_MOTORS, dtype=np.float32)
                for motor in G1_29_JointIndex:
                    init_dof_pos[motor.value] = obs[f"{motor.name}.q"]

                # Interpolate to default position
                for step in range(num_steps):
                    start_time = time.time()

                    alpha = step / num_steps
                    action_dict = {}
                    for motor in G1_29_JointIndex:
                        target_pos = default_positions[motor.value]
                        interp_pos = init_dof_pos[motor.value] * (1 - alpha) + target_pos * alpha
                        action_dict[f"{motor.name}.q"] = float(interp_pos)

                    self.publish_lowcmd(action_dict)

                    # Maintain constant control rate
                    elapsed = time.time() - start_time
                    sleep_time = max(0, control_dt - elapsed)
                    time.sleep(sleep_time)

            # Reset controller internal state (gait phase, obs history, etc.) while the thread
            # is still held off, so it cannot half-refill the buffers we are clearing.
            if self.controller is not None and hasattr(self.controller, "reset"):
                self.controller.reset()

        logger.info("Reached default position")
