#!/usr/bin/env python3
"""
RX1 Robot Brain Server v2 — LEARNED action model in the loop.

Same HTTP contract as scripts/rx1_brain_server.py (the v1 IK server is left
untouched), but the robot is driven by the trained 40-DoF (arms-only) action
model (models/policy_learned.py) instead of scripted IK:

  ZelPi OS → Gemma/Qwen → firewall → brain_server_v2 (this file)
    → Brain.set_task() → _LearnedRX1Policy (high-level planner picks the cube +
      place target; the LEARNED model outputs the arm trajectory + grasp)
    → RX1Env.step() → MuJoCo

The model is queried once per pick-place step and the trajectory is played
open-loop (see models/policy_learned.py) which is why think_freq is paced to the
control rate. If no trained checkpoint exists it transparently falls back to the
v1 IK policy, so the server always runs.

Endpoints (port 8788, same as v1 — run one server at a time):
  POST /task   {"instruction": "stack the cubes"}
  GET  /status · GET /health

Usage:
  python scripts/rx1_brain_server_v2.py --render
  python scripts/rx1_brain_server_v2.py --ckpt models/ckpt/act.pt
"""
from __future__ import annotations

import argparse
import json
import logging
import os
import sys
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Optional

import yaml

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

from brain.brain import Brain
from env.rx1_env import RX1Env
from hal.hal import HAL
from models.diffusion_policy import DiffusionPathPlanner
from models.language import OpenVLALanguageModel
from models.world_model import CosmosWorldModel
from models.policy_learned import make_learned_policy

log = logging.getLogger("rx1_brain_server_v2")
_brain: Optional[Brain] = None


def _perceive_objects():
    """Live object poses from the env, so the planner snapshot can resolve
    concrete pick/place targets. Returns None if unavailable."""
    try:
        env = getattr(_brain.hal, "env", None) if _brain is not None else None
        if env is not None and hasattr(env, "perceive_objects"):
            return env.perceive_objects()
    except Exception as e:
        log.debug("perceive_objects failed: %s", e)
    return None


class _Handler(BaseHTTPRequestHandler):
    def log_message(self, fmt, *a):
        log.debug("HTTP %s", fmt % a)

    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, "running": _brain is not None,
                             "robot": "RX1", "policy": "learned-v2"})
        elif self.path == "/status":
            if _brain is None:
                self._json(503, {"error": "brain not started"})
            else:
                st = _brain.stats.copy()
                pol = _brain.policy
                if hasattr(pol, "plan_snapshot"):
                    snap = pol.plan_snapshot(_perceive_objects())
                    st["plan"] = snap
                    st["active_step"] = snap.get("active")   # JSON-safe (no ndarrays)
                    st["queue"] = snap.get("queue_remaining", 0)
                elif hasattr(pol, "_step"):
                    st["queue"] = len(getattr(pol, "_queue", []))
                self._json(200, st)
        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()
        if not instr:
            self._json(400, {"error": "instruction required"}); return
        if _brain is None:
            self._json(503, {"error": "brain not started"}); return
        _brain.set_task(instr)
        log.info("[rx1-v2] task → '%s'", instr)
        snap = None
        pol = _brain.policy
        if hasattr(pol, "plan_snapshot"):
            try:
                snap = pol.plan_snapshot(_perceive_objects())
            except Exception as e:
                log.debug("plan_snapshot failed: %s", e)
        self._json(200, {"ok": True, "task": instr, "plan": snap})


def main():
    p = argparse.ArgumentParser(description="ZelPi RX1 Brain Server v2 (learned)")
    p.add_argument("--port", type=int, default=int(os.environ.get("BRAIN_PORT", 8788)))
    p.add_argument("--render", action="store_true",
                   default=os.environ.get("BRAIN_RENDER", "") == "1")
    p.add_argument("--duration", type=float,
                   default=float(os.environ.get("BRAIN_DURATION", "inf")))
    p.add_argument("--task", default=os.environ.get("BRAIN_TASK", ""))
    p.add_argument("--ckpt", default=os.environ.get("RX1_ACT_CKPT", "models/ckpt/act.pt"))
    p.add_argument("--device", default=os.environ.get("RX1_DEVICE",
                   "cuda" if os.environ.get("RX1_DEVICE") is None and _cuda() else "cpu"))
    p.add_argument("--robot-config", default="configs/rx1_robot.yaml")
    p.add_argument("--model-config", default="configs/rx1_models_v2.yaml")
    p.add_argument("--log-level", default="INFO")
    args = p.parse_args()

    logging.basicConfig(level=getattr(logging, args.log_level),
                        format="%(asctime)s  %(name)-20s  %(levelname)s  %(message)s",
                        datefmt="%H:%M:%S")

    def _load(path):
        path = Path(path)
        if not path.is_absolute():
            path = Path(__file__).resolve().parents[1] / path
        with open(path, encoding="utf-8") as f:
            return yaml.safe_load(f)

    robot_cfg = _load(args.robot_config)
    model_cfg = _load(args.model_config)
    pol_cfg = dict(model_cfg.get("policy", {}))
    pol_cfg.setdefault("ckpt_path", args.ckpt)
    pol_cfg.setdefault("device", args.device)

    log.info("Building RX1 v2 (learned) brain stack … ckpt=%s device=%s",
             pol_cfg["ckpt_path"], pol_cfg["device"])
    env = RX1Env(robot_cfg)
    world_model = CosmosWorldModel(model_cfg["world_model"])
    policy = make_learned_policy(pol_cfg)            # learned model, or v1 fallback
    language = OpenVLALanguageModel(model_cfg["language"])
    planner = DiffusionPathPlanner(model_cfg["diffusion_planner"])
    hal = HAL(env, robot_cfg["hal"])

    global _brain
    _brain = Brain(hal, world_model, policy, language, model_cfg["brain"],
                   path_planner=planner)

    initial = args.task or model_cfg["brain"].get("default_task", "stand by")
    httpd = ThreadingHTTPServer(("0.0.0.0", args.port), _Handler)

    if args.render:
        threading.Thread(target=httpd.serve_forever, daemon=True,
                         name="HTTPServer").start()
        log.info("RX1 v2 HTTP → http://localhost:%d  (viewer)", args.port)
        try:
            _brain.run_with_viewer(task=initial, duration=args.duration)
        finally:
            httpd.shutdown()
    else:
        threading.Thread(target=lambda: _brain.run(task=initial, duration=args.duration),
                         daemon=True, name="BrainLoop").start()
        log.info("RX1 v2 Brain HTTP → http://localhost:%d  (headless)", args.port)
        log.info("Send tasks:  curl -X POST http://localhost:%d/task "
                 "-H 'Content-Type: application/json' "
                 "-d '{\"instruction\": \"stack the cubes\"}'", args.port)
        try:
            httpd.serve_forever()
        finally:
            httpd.shutdown()


def _cuda():
    try:
        import torch
        return torch.cuda.is_available()
    except Exception:
        return False


if __name__ == "__main__":
    main()
