"""MuJoCo humanoid environment with vision and touch sensors."""
from __future__ import annotations

import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Dict, Optional

import mujoco
import numpy as np

_XML_PATH = Path(__file__).parent / "assets" / "humanoid_touch.xml"

_TOUCH_SENSOR_NAMES = [
    "touch_right_hand_palm",
    "touch_right_hand_back",
    "touch_right_fingertip",
    "touch_left_hand_palm",
    "touch_left_hand_back",
    "touch_left_fingertip",
    "touch_right_foot",
    "touch_left_foot",
    "touch_torso_front",
    "touch_torso_back",
]
_IMU_SENSOR_NAMES = ["accel_torso", "gyro_torso"]


@dataclass
class RobotObs:
    vision: Optional[np.ndarray]              # (H, W, 3) uint8 RGB primary ego view
    touch: np.ndarray                         # (N_touch,)  float32 forces [N]
    imu: np.ndarray                           # (6,) float32  [accel xyz | gyro xyz]
    proprio: np.ndarray                       # (qpos+qvel,) float32
    sim_time: float
    multi_view: Optional[Dict[str, np.ndarray]] = None  # {role: (H, W, 3) uint8}
    wall_time: float = field(default_factory=time.time)


class HumanoidTouchEnv:
    """
    MuJoCo humanoid with:
      - 4 cameras: ego, side_camera (overview), right_wrist_camera, left_wrist_camera
      - 10 touch sensors (hands, feet, torso)
      - 6-DOF IMU on torso
      - 21 torque actuators

    All rendering is done inside the HAL control thread that calls step().
    Never render from a different thread — EGL contexts are thread-local.
    """

    N_TOUCH = len(_TOUCH_SENSOR_NAMES)
    N_IMU = 6
    N_ACTUATORS = 21
    CAMERA_W = 224
    CAMERA_H = 224

    def __init__(self, config: dict):
        self.config = config
        vis_cfg = config.get("vision", {})
        self.CAMERA_W = vis_cfg.get("width", 224)
        self.CAMERA_H = vis_cfg.get("height", 224)
        self._vision_camera = vis_cfg.get("camera", "ego_camera")

        mv = vis_cfg.get("multi_view", {})
        self._multi_view_cameras: Dict[str, str] = {
            "overview":    mv.get("overview",    "side_camera"),
            "right_wrist": mv.get("right_wrist", "right_wrist_camera"),
            "left_wrist":  mv.get("left_wrist",  "left_wrist_camera"),
        }

        self.model = mujoco.MjModel.from_xml_path(str(_XML_PATH))
        self.data = mujoco.MjData(self.model)

        sim_cfg = config.get("simulation", {})
        self.model.opt.timestep = sim_cfg.get("physics_timestep", 0.005)
        self._n_substeps = sim_cfg.get("n_substeps", 4)

        self._renderer: Optional[mujoco.Renderer] = None
        self._sensor_adr = self._resolve_sensors()
        self.proprio_dim = self.model.nq + self.model.nv

    # ------------------------------------------------------------------ setup

    def _resolve_sensors(self) -> Dict[str, tuple]:
        adr_map: Dict[str, tuple] = {}
        for name in _TOUCH_SENSOR_NAMES + _IMU_SENSOR_NAMES:
            sid = mujoco.mj_name2id(self.model, mujoco.mjtObj.mjOBJ_SENSOR, name)
            if sid < 0:
                raise ValueError(f"Sensor '{name}' not found in XML")
            adr_map[name] = (int(self.model.sensor_adr[sid]),
                             int(self.model.sensor_dim[sid]))
        return adr_map

    def _ensure_renderer(self):
        """Create renderer on first use — must be called from the owning thread."""
        if self._renderer is None:
            self._renderer = mujoco.Renderer(self.model, self.CAMERA_H, self.CAMERA_W)

    # ------------------------------------------------------------------ public

    def reset(self, render_vision: bool = False) -> RobotObs:
        mujoco.mj_resetData(self.model, self.data)
        mujoco.mj_forward(self.model, self.data)
        return self._make_obs(render_vision=render_vision)

    def step(self, action: np.ndarray,
             render_vision: bool = False,
             render_multi_view: bool = False) -> RobotObs:
        """Apply action, step physics, return observation."""
        ctrl = np.clip(action,
                       -self.config["actuator"]["control_clip"],
                       self.config["actuator"]["control_clip"])
        self.data.ctrl[:] = ctrl
        for _ in range(self._n_substeps):
            mujoco.mj_step(self.model, self.data)
        return self._make_obs(render_vision=render_vision,
                              render_multi_view=render_multi_view)

    # ------------------------------------------------------------------ obs

    def _make_obs(self, render_vision: bool,
                  render_multi_view: bool = False) -> RobotObs:
        vision = self._get_vision() if render_vision else None
        mv = self._get_multi_view() if render_multi_view else None
        return RobotObs(
            vision=vision,
            touch=self._get_touch(),
            imu=self._get_imu(),
            proprio=self._get_proprio(),
            sim_time=float(self.data.time),
            multi_view=mv,
        )

    def _get_vision(self) -> np.ndarray:
        self._ensure_renderer()
        self._renderer.update_scene(self.data, camera=self._vision_camera)
        return self._renderer.render().copy()

    def _get_multi_view(self) -> Dict[str, np.ndarray]:
        """Render all three policy-facing cameras."""
        self._ensure_renderer()
        views: Dict[str, np.ndarray] = {}
        for role, cam_name in self._multi_view_cameras.items():
            self._renderer.update_scene(self.data, camera=cam_name)
            views[role] = self._renderer.render().copy()
        return views

    def _get_touch(self) -> np.ndarray:
        out = np.empty(self.N_TOUCH, dtype=np.float32)
        for i, name in enumerate(_TOUCH_SENSOR_NAMES):
            adr, _ = self._sensor_adr[name]
            out[i] = self.data.sensordata[adr]
        return out

    def _get_imu(self) -> np.ndarray:
        parts = []
        for name in _IMU_SENSOR_NAMES:
            adr, dim = self._sensor_adr[name]
            parts.append(self.data.sensordata[adr: adr + dim])
        return np.concatenate(parts).astype(np.float32)

    def _get_proprio(self) -> np.ndarray:
        return np.concatenate([self.data.qpos.copy(),
                               self.data.qvel.copy()]).astype(np.float32)

    # ------------------------------------------------------------------ utils

    @property
    def dt(self) -> float:
        return self.model.opt.timestep * self._n_substeps

    def close(self):
        if self._renderer is not None:
            self._renderer.close()
            self._renderer = None
