#!/usr/bin/env python3
"""
demo2 arm pick-and-place server — Shadow Hand on a torso+arm.

Drives the learned pick-and-place policy (arm_policy.LearnedArmPolicy) over the
torso+arm+hand scene (build_arm_scene). HTTP:
  POST /task {"instruction":"pick and place"}  -> perceive obj+pad, plan, execute
  GET  /status · /health

Usage:
  python demo2/dexarm_server.py --render
  (or via OS:  zelpi demo2 pickplace)
"""
from __future__ import annotations

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

import numpy as np
import mujoco

from build_arm_scene import build_model
from arm_policy import make_arm_policy
from arm_pick_place import ArmIK
import arm_language

log = logging.getLogger("dexarm")
_state = {"model": None, "data": None, "policy": None, "ik": None,
          "task": None, "running_task": False}


def _perceive(m, d):
    oid = mujoco.mj_name2id(m, mujoco.mjtObj.mjOBJ_BODY, "object")
    pp = mujoco.mj_name2id(m, mujoco.mjtObj.mjOBJ_GEOM, "place_pad")
    pick = d.xpos[oid].copy()
    place = m.geom_pos[pp].copy(); place[2] = 0.526
    return pick, place


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": "ShadowArm", "demo": "demo2-arm",
                             "learned": _state["policy"] is not None})
        elif self.path == "/status":
            pol = _state["policy"]
            self._json(200, {"robot": "ShadowArm",
                             "phase": pol.phase if pol else "no-policy",
                             "plan": pol.plan if pol else None})
        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
        pol = _state["policy"]
        if pol is None:
            self._json(503, {"error": "no learned model (train arm_act.pt first)"}); return
        m, d = _state["model"], _state["data"]
        instr = str(body.get("instruction", "")).strip()
        # language → which object + where (the policy is goal-conditioned on the
        # pick/place coords, so language just resolves colour → object position).
        pick_body, pick, place, color = arm_language.resolve(instr, m, d)
        plan = pol.plan_pickplace(d, pick, place, pick_body=pick_body, color=color)
        _state["running_task"] = True
        log.info("[dexarm] '%s' -> pick %s (%s) -> place %s (chunk %d)",
                 instr, np.round(pick, 3), color, np.round(place, 3), plan["chunk_len"])
        self._json(200, {"ok": True, "mode": "pickplace", "instruction": instr,
                         "plan": plan})


def _control_tick():
    m, d, pol = _state["model"], _state["data"], _state["policy"]
    if _state["running_task"] and pol is not None:
        done, _ = pol.step(m, d)
        if done:
            _state["running_task"] = False
    for _ in range(_state["n_sub"]):
        mujoco.mj_step(m, d)


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--port", type=int, default=int(os.environ.get("DEXARM_PORT", 8791)))
    p.add_argument("--render", action="store_true")
    p.add_argument("--device", default="cpu")
    p.add_argument("--log-level", default="INFO")
    args = p.parse_args()
    logging.basicConfig(level=getattr(logging, args.log_level), format="%(message)s")

    m = build_model(); d = mujoco.MjData(m); ik = ArmIK(m)
    d.ctrl[:] = ik.rest()
    mujoco.mj_forward(m, d)
    pol = make_arm_policy(m, args.device)
    _state.update(model=m, data=d, ik=ik, policy=pol,
                  n_sub=max(1, int(round((1 / 100.0) / m.opt.timestep))))
    if pol is None:
        log.warning("[dexarm] no arm_act.pt — train it: python demo2/train_arm.py")

    httpd = ThreadingHTTPServer(("0.0.0.0", args.port), _Handler)
    threading.Thread(target=httpd.serve_forever, daemon=True, name="HTTP").start()
    log.info("demo2 arm pick-place -> http://localhost:%d  (learned=%s)",
             args.port, pol is not None)

    # let the object settle on the table first
    for _ in range(150):
        mujoco.mj_step(m, d)

    if args.render:
        import mujoco.viewer as _mjv
        with _mjv.launch_passive(m, d) as v:
            while v.is_running():
                t = time.monotonic()
                with v.lock():
                    _control_tick()
                v.sync()
                time.sleep(max(0.0, 0.01 - (time.monotonic() - t)))
        httpd.shutdown()
    else:
        try:
            while True:
                t = time.monotonic(); _control_tick()
                time.sleep(max(0.0, 0.01 - (time.monotonic() - t)))
        except KeyboardInterrupt:
            httpd.shutdown()


if __name__ == "__main__":
    main()
