#!/usr/bin/env python3
"""
Shadow Dexterous Hand demo server — ZelPi /demo2 (5-finger gripper).

A fully self-contained stack (its own 3-D model under demo2/model/, its own
controller, its own server) — independent of the RX1 sim. Drives the 20-DoF
Shadow Hand to natural-language gestures via simple position control.

HTTP (port 8790):
  POST /task   {"instruction": "grasp"}     # open·fist·grasp·point·pinch·peace·
                                            #   ok·spread·wave·count·rest …
  GET  /status · GET /health

Usage:
  python demo2/dexhand_server.py --render        # MuJoCo viewer
  python demo2/dexhand_server.py --gesture wave  # headless
  (or via the OS:  zelpi demo2 launch --render  ·  zelpi demo2 task "fist")
"""
from __future__ import annotations

import argparse
import json
import logging
import os
import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path

import numpy as np
import mujoco

from dexhand_controller import DexHandController
import grasp_planner as gp

log = logging.getLogger("dexhand")
_HERE = Path(__file__).resolve().parent
_SCENE = _HERE / "model" / "scene_right.xml"

# slew-rate limit (rad/control-step) so gestures move smoothly, not instantly
_CTRL_DT = 0.002 * 1          # timestep; substeps applied in loop
_MAX_DCTRL = 2.5 * 0.01       # ~2.5 rad/s over a 10 ms control tick

# instruction keywords that trigger the perceive→plan→grasp pipeline (instead of
# a fixed gesture). "fist"/"close" stay pure gestures.
_GRASP_KW = ("grasp", "pick", "grab", "take")

# grasp execution phases and their timeouts (seconds)
_PRESHAPE_S, _CLOSE_S, _SETTLE_S = 0.8, 3.0, 0.4

_state = {
    "gesture": "open", "controller": None, "data": None, "t0": None,
    # grasp pipeline
    "mode": "gesture", "maps": None, "percept": None, "plan": None,
    "phase": None, "phase_t0": 0.0, "touched": set(), "grasp_spec": None,
}


class _Handler(BaseHTTPRequestHandler):
    def log_message(self, *a):
        pass

    def _json(self, code, obj):
        body = json.dumps(obj).encode()
        self.send_response(code)
        self.send_header("Content-Type", "application/json")
        self.send_header("Content-Length", str(len(body)))
        self.send_header("Access-Control-Allow-Origin", "*")
        self.end_headers()
        self.wfile.write(body)

    def do_GET(self):
        if self.path == "/health":
            self._json(200, {"ok": True, "robot": "ShadowHand", "demo": "demo2"})
        elif self.path == "/status":
            c = _state["controller"]
            d = _state["data"]
            flex = {}
            if d is not None and c is not None:
                for f in ("FF", "MF", "RF", "LF"):
                    nm = f"rh_{f}J3"
                    jid = mujoco.mj_name2id(c.model, mujoco.mjtObj.mjOBJ_JOINT, nm)
                    if jid >= 0:
                        flex[f] = round(float(d.qpos[c.model.jnt_qposadr[jid]]), 3)
            resp = {"robot": "ShadowHand", "gesture": _state["gesture"],
                    "mode": _state["mode"], "finger_flex_rad": flex,
                    "gestures": c.known() if c else []}
            if _state["mode"] == "grasp" and _state["plan"] is not None:
                resp["grasp"] = gp.snapshot(_state["percept"], _state["plan"],
                                            phase=_state["phase"],
                                            touched=_state["touched"])
            self._json(200, resp)
        else:
            self._json(404, {"error": "not found"})

    def do_OPTIONS(self):
        self.send_response(200)
        self.send_header("Access-Control-Allow-Origin", "*")
        self.send_header("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
        self.send_header("Access-Control-Allow-Headers", "Content-Type")
        self.end_headers()

    def do_POST(self):
        if self.path != "/task":
            self._json(404, {"error": "not found"}); return
        n = int(self.headers.get("Content-Length", 0))
        try:
            body = json.loads(self.rfile.read(n) or b"{}")
        except json.JSONDecodeError:
            self._json(400, {"error": "invalid JSON"}); return
        instr = str(body.get("instruction", "")).strip()
        c = _state["controller"]
        d = _state["data"]
        if not instr or c is None:
            self._json(400, {"error": "instruction required / not started"}); return

        # ── grasp pipeline: perceive → plan → (control loop executes) ─────────
        if any(k in instr.lower() for k in _GRASP_KW):
            percept = gp.perceive(c.model, d)
            plan = gp.plan_grasp(c.model, d, percept)
            _state.update(mode="grasp", percept=percept, plan=plan,
                          phase="preshape", phase_t0=d.time,
                          touched=set(), grasp_spec=dict(plan["preshape"]),
                          gesture="grasp", welded=False)
            log.info("[demo2] grasp '%s' → obj@%s dist=%.3f width=%.3f "
                     "strength=%.2f quality=%.2f", instr, percept["center"],
                     percept["distance"], percept["width"],
                     plan["strength"], plan["quality"])
            self._json(200, {"ok": True, "instruction": instr, "mode": "grasp",
                             "grasp": gp.snapshot(percept, plan, phase="preshape",
                                                  touched=set())})
            return

        # ── gesture path ─────────────────────────────────────────────────────
        gp.set_weld(c.model, d, active=False)   # release any grasped object
        _state["welded"] = False
        g = c.set_gesture(instr)
        _state.update(mode="gesture", gesture=g, t0=None, plan=None, phase=None)
        log.info("[demo2] task '%s' -> gesture=%s", instr, g)
        self._json(200, {"ok": True, "instruction": instr, "gesture": g})


def _clamp_step(model, data, target):
    prev = data.ctrl[:model.nu]
    delta = np.clip(target - prev, -_MAX_DCTRL, _MAX_DCTRL)
    data.ctrl[:model.nu] = prev + delta


def _grasp_step(model, data, ctrl):
    """Advance the grasp phase machine and return the actuator target vector.

    preshape → orient wrist + pre-oppose thumb, fingers open;
    close    → drive un-touched fingers shut, freeze each finger the moment its
               tip contacts the object (close-until-contact ≈ adapts to shape);
    hold     → maintain the contact pose.
    """
    now = data.time          # sim seconds — robust to real-time vs fast headless
    plan, spec, maps = _state["plan"], _state["grasp_spec"], _state["maps"]
    phase = _state["phase"]

    if phase == "preshape":
        if now - _state["phase_t0"] >= _PRESHAPE_S:
            _state["phase"] = "close"
            _state["phase_t0"] = now
            _state["grasp_spec"] = spec = dict(plan["preshape"])

    elif phase == "close":
        touched = maps.fingertip_contacts(model, data)
        # Weld only once the grasp is actually secure — ≥3 fingers, or the thumb
        # plus one opposing finger. Welding on the FIRST contact froze weak
        # 2-finger grasps (e.g. the object rolled toward the pinky and LF/RF
        # touched first, locking before the thumb/index arrived).
        secure = len(touched) >= 3 or ("TH" in touched and len(touched) >= 2)
        if secure and not _state.get("welded"):
            gp.set_weld(model, data, active=True)
            _state["welded"] = True
            log.info("[demo2] grasp secured at %d contact(s) %s",
                     len(touched), sorted(touched))
        _state["touched"] = touched
        for f in gp.FINGERS:
            j3, j0 = f"rh_A_{f}J3", f"rh_A_{f}J0"
            if f in touched:                                   # freeze on contact
                spec[j3] = float(data.ctrl[ctrl._idx[j3]])
                spec[j0] = float(data.ctrl[ctrl._idx[j0]])
            else:
                spec[j3] = plan["closed"][j3]
                spec[j0] = plan["closed"][j0]
        for jn in ("rh_A_THJ4", "rh_A_THJ2", "rh_A_THJ1"):
            if "TH" in touched and jn != "rh_A_THJ4":
                spec[jn] = float(data.ctrl[ctrl._idx[jn]])
            else:
                spec[jn] = plan["closed"][jn]
        if len(touched) >= 5 or (now - _state["phase_t0"] >= _CLOSE_S):
            _state["phase"] = "hold"
            _state["phase_t0"] = now
            log.info("[demo2] grasp closed — %d/5 contacts %s",
                     len(touched), sorted(touched))

    return ctrl._pose(spec)


def main():
    p = argparse.ArgumentParser(description="ZelPi /demo2 Shadow Hand server")
    p.add_argument("--port", type=int, default=int(os.environ.get("DEMO2_PORT", 8790)))
    p.add_argument("--render", action="store_true",
                   default=os.environ.get("DEMO2_RENDER", "") == "1")
    p.add_argument("--gesture", default=os.environ.get("DEMO2_GESTURE", "open"))
    p.add_argument("--log-level", default="INFO")
    args = p.parse_args()
    logging.basicConfig(level=getattr(logging, args.log_level), format="%(message)s")

    model = mujoco.MjModel.from_xml_path(str(_SCENE))
    data = mujoco.MjData(model)
    ctrl = DexHandController(model)
    ctrl.set_gesture(args.gesture)
    _state.update(controller=ctrl, data=data, gesture=ctrl.gesture,
                  maps=gp.GraspMaps(model))
    mujoco.mj_forward(model, data)
    n_sub = max(1, int(round((1 / 100.0) / model.opt.timestep)))   # ~100 Hz control

    httpd = ThreadingHTTPServer(("0.0.0.0", args.port), _Handler)
    threading.Thread(target=httpd.serve_forever, daemon=True, name="HTTP").start()
    log.info("Shadow Hand /demo2 → http://localhost:%d  (gestures: %s)",
             args.port, ", ".join(ctrl.known()))
    log.info("send:  curl -X POST http://localhost:%d/task -H 'Content-Type: application/json'"
             " -d '{\"instruction\": \"grasp\"}'", args.port)

    t0 = time.monotonic()

    def control_tick():
        if _state["mode"] == "grasp":
            target = _grasp_step(model, data, ctrl)
        else:
            if _state["t0"] is None:
                _state["t0"] = time.monotonic()
            target = ctrl.target(time.monotonic() - _state["t0"])
        _clamp_step(model, data, target)
        for _ in range(n_sub):
            mujoco.mj_step(model, data)

    if args.render:
        import mujoco.viewer as _mjv     # aliased so it doesn't shadow `mujoco`
        log.info("Launching MuJoCo viewer — close the window to exit.")
        with _mjv.launch_passive(model, data) as viewer:
            while viewer.is_running():
                ts = time.monotonic()
                with viewer.lock():
                    control_tick()
                viewer.sync()
                time.sleep(max(0.0, 0.01 - (time.monotonic() - ts)))
        httpd.shutdown()
    else:
        log.info("Headless — Ctrl-C to stop.")
        try:
            while True:
                ts = time.monotonic()
                control_tick()
                time.sleep(max(0.0, 0.01 - (time.monotonic() - ts)))
        except KeyboardInterrupt:
            httpd.shutdown()


if __name__ == "__main__":
    main()
