#!/usr/bin/env python3
"""
RX1 Robot Brain Server — VLA-driven pick-and-place via Jacobian IK.

Pipeline (per Gemma intent):
  ZelPi OS → Gemma 4 → Firewall → brain_server (this file)
  → Brain.set_task() → _RX1IKPolicy.set_task_text()
  → PickExecutor → Jacobian IK → joint targets
  → RX1Env.step() → MuJoCo (with weld constraints for cube pickup)

Endpoints (port 8788):
  POST /task   {"instruction": "..."} — set new task
  GET  /status                        — brain stats + executor phase
  GET  /health                        — liveness probe

Usage:
  cd F:\\world_model_test
  python scripts/rx1_brain_server.py --render     # with MuJoCo viewer
  python scripts/rx1_brain_server.py              # headless
  python scripts/rx1_brain_server.py --task "pick red cube"
"""
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.policy import CosmosPolicy
from models.world_model import CosmosWorldModel

log = logging.getLogger("rx1_brain_server")

_brain: Optional[Brain] = None


# ──────────────────────────────────────────────────── HTTP handler

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

    def _send_json(self, code: int, obj: dict):
        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._send_json(200, {"ok": True, "running": _brain is not None,
                                  "robot": "RX1"})
        elif self.path == "/status":
            if _brain is None:
                self._send_json(503, {"error": "brain not started"})
            else:
                stats = _brain.stats.copy()
                # Add executor phase if available
                policy = _brain.policy
                if hasattr(policy, '_executor'):
                    stats["executor_phase"] = policy._executor.phase.name
                    stats["weld_active"] = bool(policy._executor.weld_active)
                self._send_json(200, stats)
        else:
            self._send_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._send_json(404, {"error": "not found"})
            return
        length = int(self.headers.get("Content-Length", 0))
        try:
            body = json.loads(self.rfile.read(length) or b"{}")
        except json.JSONDecodeError:
            self._send_json(400, {"error": "invalid JSON"})
            return

        instruction = str(body.get("instruction", "")).strip()
        if not instruction:
            self._send_json(400, {"error": "instruction required"})
            return
        if _brain is None:
            self._send_json(503, {"error": "brain not started"})
            return

        _brain.set_task(instruction)
        log.info("[rx1] task → '%s'", instruction)
        phase = "IDLE"
        if hasattr(_brain.policy, '_executor'):
            phase = _brain.policy._executor.phase.name
        self._send_json(200, {"ok": True, "task": instruction, "phase": phase})


# ──────────────────────────────────────────────────── main

def main():
    parser = argparse.ArgumentParser(description="ZelPi RX1 Brain Server")
    parser.add_argument("--port", type=int,
                        default=int(os.environ.get("BRAIN_PORT", 8788)))
    parser.add_argument("--render", action="store_true",
                        default=os.environ.get("BRAIN_RENDER", "") == "1",
                        help="Open MuJoCo viewer (requires display)")
    parser.add_argument("--duration", type=float,
                        default=float(os.environ.get("BRAIN_DURATION", "inf")))
    parser.add_argument("--task", default=os.environ.get("BRAIN_TASK", ""))
    parser.add_argument("--robot-config",  default="configs/rx1_robot.yaml")
    parser.add_argument("--model-config",  default="configs/rx1_models.yaml")
    parser.add_argument("--log-level", default="INFO",
                        choices=["DEBUG", "INFO", "WARNING", "ERROR"])
    args = parser.parse_args()

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

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

    robot_cfg = _load(args.robot_config)
    model_cfg  = _load(args.model_config)

    # Always stub mode (no GPU required for RX1 IK policy)
    model_cfg["world_model"]["stub"]      = True
    model_cfg["policy"]["stub"]           = True
    model_cfg["policy"]["use_rx1_ik"]     = True
    model_cfg["language"]["stub"]         = True
    model_cfg["diffusion_planner"]["stub"] = True

    log.info("Building RX1 brain stack …")
    env          = RX1Env(robot_cfg)
    world_model  = CosmosWorldModel(model_cfg["world_model"])
    policy       = CosmosPolicy(model_cfg["policy"])
    language_model = OpenVLALanguageModel(model_cfg["language"])
    path_planner = DiffusionPathPlanner(model_cfg["diffusion_planner"])
    hal          = HAL(env, robot_cfg["hal"])

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

    initial_task = args.task or model_cfg["brain"].get("default_task", "stand by")

    httpd = ThreadingHTTPServer(("0.0.0.0", args.port), _Handler)

    if args.render:
        # MuJoCo viewer must own the main thread (OpenGL)
        http_thread = threading.Thread(target=httpd.serve_forever,
                                       daemon=True, name="HTTPServer")
        http_thread.start()
        log.info("RX1 HTTP → http://localhost:%d  (background thread)", args.port)
        log.info("Launching MuJoCo viewer — close window to exit.")
        try:
            _brain.run_with_viewer(task=initial_task, duration=args.duration)
        except KeyboardInterrupt:
            log.info("Viewer closed by user.")
        finally:
            httpd.shutdown()
    else:
        brain_thread = threading.Thread(
            target=lambda: _brain.run(task=initial_task, duration=args.duration),
            daemon=True, name="BrainLoop",
        )
        brain_thread.start()
        log.info("RX1 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\": \"pick red cube\"}'", args.port)
        try:
            httpd.serve_forever()
        except KeyboardInterrupt:
            log.info("Shutting down.")
        finally:
            httpd.shutdown()


if __name__ == "__main__":
    main()
