"""Run the supplied LeRobot controller with local weights, live depth, and JSONL telemetry.

Defaults to the LeRobot simulation and stock SONIC standing mode. --real enables
onboard DDS; --dry-run performs inference with synthetic state and sends no commands.
Space switches SONIC/Dodge; q quits. h/m/f mark an observed hit/miss/fall in the log.
"""
import argparse
import logging
import json
from pathlib import Path
import queue
import math
import os
import shlex
import subprocess
import sys
import threading
import time

import cv2
import numpy as np

from lerobot.robots.unitree_g1.config_unitree_g1 import UnitreeG1Config
from lerobot.robots.unitree_g1.controllers.depth_dodge import DEFAULT_REPO_ID, DEPTH_SHAPE, FILES, DepthDodgeController
from lerobot.robots.unitree_g1.controllers.dodge_assets import resolve_weight
from lerobot.robots.unitree_g1.g1_utils import G1_29_JointIndex
from lerobot.robots.unitree_g1.unitree_g1 import UnitreeG1


EVENTS = queue.SimpleQueue()
args = None
use_percept = False

def parse_args(argv=None):
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("--seconds", type=float, default=600.0)
    ap.add_argument("--real", action="store_true", help="run onboard: direct DDS lowcmd/lowstate, shm perception")
    ap.add_argument("--network_interface", default=None, help="DDS interface for --real (default: auto)")
    ap.add_argument("--onscreen", action="store_true")
    ap.add_argument("--end_effector", default="dummy", choices=["dummy", "dex1", "dex3"])
    ap.add_argument("--band_seconds", type=float, default=1.0, help="sim: keep the elastic band on this long")
    ap.add_argument(
        "--camera",
        default=None,
        help="'shm' (default with --real), 'zmq://host:port' for the sim on a laptop, or 'none' for all-black",
    )
    ap.add_argument(
        "--rs_python",
        default=os.path.expanduser("~/rs_pub_venv/bin/python"),
        help="interpreter with pyrealsense2 (+ onnxruntime-gpu for the fast path); spawns the shm writer",
    )
    ap.add_argument("--rs_res", default="480x270x60", help="RealSense depth WxHxFPS for the writer")
    ap.add_argument("--camera_hz", type=float, default=25.0, help="policy camera rate (train 25); writer frames are sampled")
    ap.add_argument("--mode", default="sonic", choices=["sonic", "dodge"], help="starting mode; SPACE toggles, q quits")
    ap.add_argument(
        "--cpu_theia",
        action="store_true",
        help="run Theia in this process on the CPU (~40 ms) instead of in the writer on the GPU (~8 ms)",
    )
    ap.add_argument(
        "--ball_only",
        action="store_true",
        help="writer detects the orange ball in colour and publishes only its depth pixels (rest = far)",
    )
    ap.add_argument(
        "--view_port",
        type=int,
        default=8080,
        help="with --ball_only: live MJPEG of colour+detection / published depth / mask / Theia input; 0 = off",
    )
    ap.add_argument(
        "--writer_args",
        default="",
        help='extra args for depth_shm.py, e.g. "--ball-max-range 2.8 --min-circularity 0.3 --trace"',
    )
    ap.add_argument("--weights-dir", default=str(Path(__file__).resolve().parents[1] / "weights/released"))
    ap.add_argument("--dry-run", action="store_true", help="Run inference with synthetic lowstate; never construct/connect a robot")
    ap.add_argument("--external-writer", action="store_true", help="Read an already running shared-memory writer")
    ap.add_argument("--log", default="results/robot/session.jsonl", help="JSONL telemetry and manually marked trial outcomes")
    return ap.parse_args(argv)

def keyboard_loop(controller, stop: threading.Event):
    """SPACE toggles sonic <-> dodge, q quits. Raw tty so it works over ssh."""
    import select
    import termios
    import tty

    fd = sys.stdin.fileno()
    old = termios.tcgetattr(fd)
    try:
        tty.setcbreak(fd)
        while not stop.is_set():
            if select.select([sys.stdin], [], [], 0.1)[0]:
                ch = sys.stdin.read(1)
                if ch == " ":
                    controller.mode = "dodge" if controller.mode == "sonic" else "sonic"
                    print(f"\n>>> MODE: {controller.mode.upper()}\n", flush=True)
                elif ch in ("h", "m", "f"):
                    EVENTS.put({"event":"manual_outcome", "outcome":{"h":"hit", "m":"miss", "f":"fall"}[ch]})
                elif ch in ("q", "\x03"):
                    stop.set()
    finally:
        termios.tcsetattr(fd, termios.TCSADRAIN, old)


# Policy camera contract: 64x96 pinhole, 45 deg vertical FOV, 0 deg pitch.
POLICY_VFOV = math.radians(45.0)
POLICY_HFOV = 2 * math.atan(DEPTH_SHAPE[1] / DEPTH_SHAPE[0] * math.tan(POLICY_VFOV / 2))  # ~63.7 deg


def make_cropper(info: dict | None, h: int, w: int):
    """Return f(depth_mm uint16 HxW) -> (64,96) float32 metres, matching the policy's FOV."""
    if info:
        fx, fy, cx, cy = info["fx"], info["fy"], info["cx"], info["cy"]
    else:  # D435i nominal (58 deg vFOV)
        fy = fx = (h / 2) / math.tan(math.radians(58.0) / 2)
        cx, cy = w / 2, h / 2
    half_h = fy * math.tan(POLICY_VFOV / 2)
    half_w = fx * math.tan(POLICY_HFOV / 2)
    y0, y1 = int(round(cy - half_h)), int(round(cy + half_h))
    x0, x1 = int(round(cx - half_w)), int(round(cx + half_w))
    y0, x0 = max(y0, 0), max(x0, 0)
    y1, x1 = min(y1, h), min(x1, w)
    print(f"depth crop rows {y0}:{y1} cols {x0}:{x1} of {h}x{w} -> {DEPTH_SHAPE}", flush=True)

    def crop(depth_mm: np.ndarray) -> np.ndarray:
        roi = depth_mm[y0:y1, x0:x1]
        # Nearest keeps holes as holes (0 = no return), like the sim's rendered depth.
        small = cv2.resize(roi, (DEPTH_SHAPE[1], DEPTH_SHAPE[0]), interpolation=cv2.INTER_NEAREST)
        return small.astype(np.float32) * 1e-3

    return crop


stats = {"cam_hz": 0.0, "valid": 0.0, "near": 0.0, "frames": 0, "theia_ms": 0.0, "lat_ms": 0.0, "stalls": 0}


def open_camera():
    """Raw depth. Return (read() -> (depth_mm HxW uint16, capture_time|None), intrinsics|None, close())."""
    if args.camera == "shm":
        from lerobot.cameras.depth_shm import DepthShmReader

        reader = DepthShmReader(wait_s=15.0)

        def read():
            depth, meta = reader.read(timeout_s=2.0)
            return depth, meta["t_capture"]

        return read, reader.intrinsics, reader.close

    host, _, port = args.camera.removeprefix("zmq://").partition(":")
    from lerobot.cameras.zmq import ZMQCamera, ZMQCameraConfig

    cam = ZMQCamera(
        ZMQCameraConfig(
            server_address=host, port=int(port or 5555), camera_name="head_camera",
            use_rgb=False, use_depth=True, warmup_s=3,
        )
    )
    cam.connect()

    def read():
        return cam.async_read_depth(timeout_ms=2000)[:, :, 0], cam.latest_server_timestamp

    return read, cam.depth_info, cam.disconnect


def open_percept():
    """Tokens from the writer. Return (read() -> (tokens, depth_m, t_capture, theia_ms), close())."""
    from lerobot.cameras.depth_shm import PerceptShmReader

    reader = PerceptShmReader(wait_s=30.0)  # the writer warms up CUDA first

    def read():
        tokens, depth_m, meta = reader.read(timeout_s=2.0)
        return tokens, depth_m, meta["t_capture"], meta["theia_ms"]

    return read, reader.close


class Cadence:
    """Pick frames at a fixed rate from a faster stream, by capture time (not arrival time), so the
    ring contains 16 frames at 25 Hz like in training even though the writer runs at 60 Hz."""

    def __init__(self, hz: float):
        self.period = 1.0 / hz
        self.next: float | None = None

    def take(self, t_capture: float) -> bool:
        if self.next is None:
            self.next = t_capture + self.period
            return True
        if t_capture + 1e-4 < self.next:
            return False
        # Advance to the slot containing this frame; skip slots if we fell behind (stall).
        self.next += self.period * max(1, math.floor((t_capture - self.next) / self.period) + 1)
        return True


def camera_loop(controller, stop: threading.Event):
    cadence = Cadence(args.camera_hz)
    if use_percept:
        read_percept, close = open_percept()
        tokens, depth_m, t_capture, _ = read_percept()
        print("camera connected: tokens from the writer (Theia onboard GPU)", flush=True)
        crop = None
    else:
        read, intrinsics, close = open_camera()
        first, _ = read()
        crop = make_cropper(intrinsics, *first.shape)
        print(f"camera connected: {first.shape}, intrinsics {intrinsics}", flush=True)
    last = time.monotonic()
    try:
        while not stop.is_set():
            try:
                if use_percept:
                    tokens, depth_m, t_capture, theia_ms = read_percept()
                else:
                    depth_mm, t_capture = read()
            except TimeoutError as e:
                # Camera stalled (USB hiccup / writer restarting its pipeline). Don't act on stale
                # tokens: show the policy an empty scene until frames come back.
                stats["stalls"] += 1
                print(f"\n!!! camera stall #{stats['stalls']}: {e}; feeding black frames until it recovers\n", flush=True)
                for _ in range(controller.depth_ring.maxlen):
                    controller.observe_tokens(controller.black_tokens)
                cadence.next = None
                continue
            if t_capture is None:
                t_capture = time.time()
            if not cadence.take(t_capture):
                continue
            now = time.monotonic()
            if use_percept:
                controller.observe_tokens(tokens)
                stats["theia_ms"] = theia_ms
            else:
                depth_m = crop(depth_mm)
                t_theia = time.perf_counter()
                controller.observe_depth(depth_m)
                stats["theia_ms"] = 0.8 * stats["theia_ms"] + 0.2 * (time.perf_counter() - t_theia) * 1e3
            # capture (sensor global time) -> tokens in the ring; same clock onboard
            stats["lat_ms"] = 0.8 * stats["lat_ms"] + 0.2 * (time.time() - t_capture) * 1e3
            stats["cam_hz"] = 0.9 * stats["cam_hz"] + 0.1 / max(now - last, 1e-6)
            last = now
            valid = depth_m > 0
            stats["valid"] = valid.mean()
            stats["near"] = float(depth_m[valid].min()) if valid.any() else float("nan")
            stats["frames"] += 1
    finally:
        close()


def start_writer():
    if args.camera != "shm" or args.external_writer:
        return None
    from lerobot.cameras import depth_shm
    w, h, fps = args.rs_res.split("x")
    command = [str(Path(args.rs_python).expanduser()), depth_shm.__file__,
               "--width", w, "--height", h, "--fps", fps]
    if args.ball_only:
        command += ["--ball-only"]
        if args.view_port:
            command += ["--view-port", str(args.view_port)]
    if use_percept:
        command += ["--theia", resolve_weight(FILES["theia"])]
    command += shlex.split(args.writer_args)
    process = subprocess.Popen(command, stdin=subprocess.DEVNULL)
    print(f"Started camera writer PID {process.pid}", flush=True)
    return process


def main(argv=None):
    global args, use_percept
    args = parse_args(argv)
    if args.seconds <= 0 or not 0 < args.camera_hz <= 25:
        raise ValueError("seconds must be positive; camera_hz must be in (0,25]")
    if args.camera is None:
        args.camera = "shm" if args.real else "none"
    use_percept = args.camera == "shm" and not args.cpu_theia
    os.environ["G1_DODGE_WEIGHTS"] = str(Path(args.weights_dir).expanduser().resolve())
    # Check every asset before starting the camera or constructing the robot.
    manifest = json.loads((Path(os.environ["G1_DODGE_WEIGHTS"]) / "manifest.json").read_text())
    for name in manifest["files"]:
        resolve_weight(name)
    logpath = Path(args.log); logpath.parent.mkdir(parents=True, exist_ok=True)
    logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s")
    stop = threading.Event(); writer = robot = None; worker = None
    started = time.monotonic()
    with logpath.open("a", buffering=1) as log:
        def emit(row):
            row = {"unix_time": time.time(), "elapsed_s": time.monotonic()-started, **row}
            log.write(json.dumps(row, allow_nan=False)+"\n")
        emit({"event":"start", "dry_run":args.dry_run, "real":args.real,
              "actor_sha256":manifest["actor"]["sha256"], "training_iteration":manifest["training_iteration"],
              "arguments":vars(args)})
        try:
            writer = start_writer()
            if args.dry_run:
                from types import SimpleNamespace
                controller = DepthDodgeController(mode=args.mode)
                lowstate = SimpleNamespace(
                    motor_state=[SimpleNamespace(q=float(q), dq=0.) for q in controller.default_angles],
                    imu_state=SimpleNamespace(quaternion=[1.,0.,0.,0.], gyroscope=[0.,0.,0.]))
                print("DRY RUN: synthetic lowstate; no robot connection or joint commands.", flush=True)
            else:
                cfg = UnitreeG1Config(is_simulation=not args.real, onboard=args.real,
                    network_interface=args.network_interface, controller="DepthDodgeController",
                    sim_onscreen=args.onscreen if not args.real else None,
                    sim_publish_images=False, end_effector=args.end_effector)
                robot = UnitreeG1(cfg)
                controller = robot.controller; controller.mode = args.mode
                robot.connect()
            if sys.stdin.isatty():
                threading.Thread(target=keyboard_loop, args=(controller, stop), daemon=True).start()
            if args.camera != "none":
                def read_camera():
                    try:
                        camera_loop(controller, stop)
                    except Exception as exc:
                        controller.mode = "sonic"
                        EVENTS.put({"event":"camera_error", "message":str(exc)})
                worker = threading.Thread(target=read_camera, daemon=True)
                worker.start()
            band = getattr(getattr(getattr(robot, "sim_env", None), "sim_env", None), "elastic_band", None)
            next_log = time.monotonic()
            while time.monotonic()-started < args.seconds and not stop.is_set():
                tick = time.monotonic()
                if args.dry_run:
                    controller.run_step({}, lowstate)
                if band is not None and tick-started >= args.band_seconds:
                    band.enable = False; band = None
                if writer is not None and writer.poll() is not None:
                    controller.mode = "sonic"
                    emit({"event":"writer_exited", "returncode":writer.returncode})
                    stop.set()
                while not EVENTS.empty():
                    emit({**EVENTS.get(), "mode":controller.mode})
                if tick >= next_log:
                    tm = controller.timing
                    row = {"event":"telemetry", "mode":controller.mode,
                           "control_hz":1000/tm["period_ms"] if tm["period_ms"]>0 else None, "compute_ms":tm["tick_ms"],
                           "ring_span_ms":1000*controller.ring_span_s,
                           "max_abs_action":float(np.abs(controller.last_action).max()),
                           "camera":{k:(float(v) if np.isfinite(v) else None) for k,v in stats.items()}}
                    emit(row)
                    print(f"{controller.mode:5s} control {row['control_hz'] or 0:.1f}Hz / {row['compute_ms']:.1f}ms "
                          f"camera {stats['cam_hz']:.1f}Hz ring {row['ring_span_ms']:.0f}ms (target 600)", flush=True)
                    next_log = tick+1
                time.sleep(max(0,.02-(time.monotonic()-tick)))
        finally:
            stop.set()
            if robot is not None:
                robot.disconnect()
            if worker is not None:
                worker.join(timeout=2.5)
            if writer is not None and writer.poll() is None:
                writer.terminate()
                try:
                    writer.wait(timeout=3)
                except subprocess.TimeoutExpired:
                    writer.kill(); writer.wait()
            emit({"event":"stop"})


if __name__ == "__main__":
    main()
