#!/usr/bin/env python3
"""rosbag -> episodes.json for VisionLibra Robot Evals (visionlibra.com/evals).

Best-effort skeleton extractor: reads a ROS1 .bag or ROS2 bag directory
(sqlite3/mcap) with the pure-python `rosbags` library (no ROS install
needed), splits the recording into episodes on gaps in message flow, and
emits an episodes.json you can upload at https://visionlibra.com/evals.

    pip install rosbags
    python3 rosbag_to_episodes.py my_run.bag -o episodes.json
    python3 rosbag_to_episodes.py my_ros2_bag_dir/ --gap 3.0 --task "Pick bottle"

What it fills in automatically: episode ids, start/end, duration, and a
phase timeline built from activity on marker-ish topics (anything whose
name contains: phase, state, event, status, grasp, gripper). What you
should edit afterwards: task names and result (success/fail) per episode
if your bags don't publish them — or add a /eval/result topic upstream.
"""
import argparse
import json
import sys
from pathlib import Path

try:
    from rosbags.highlevel import AnyReader
except ImportError:
    sys.exit("missing dependency — run:  pip install rosbags")

MARKER_HINTS = ("phase", "state", "event", "status", "grasp", "gripper", "result")


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("bag", help=".bag file or ROS2 bag directory")
    ap.add_argument("-o", "--out", default="episodes.json")
    ap.add_argument("--gap", type=float, default=2.5,
                    help="seconds of silence that starts a new episode (default 2.5)")
    ap.add_argument("--task", default="unlabeled task",
                    help="task name applied to every episode")
    args = ap.parse_args()

    events = []   # (t_seconds, label)
    t0 = None
    with AnyReader([Path(args.bag)]) as reader:
        marker_conns = [c for c in reader.connections
                        if any(h in c.topic.lower() for h in MARKER_HINTS)]
        conns = marker_conns or reader.connections
        for conn, timestamp, raw in reader.messages(connections=conns):
            t = timestamp / 1e9
            if t0 is None:
                t0 = t
            label = conn.topic.strip("/").split("/")[-1]
            if marker_conns:
                try:  # string-ish messages carry the actual phase name
                    msg = reader.deserialize(raw, conn.msgtype)
                    data = getattr(msg, "data", None)
                    if isinstance(data, str) and 0 < len(data) < 60:
                        label = data
                except Exception:
                    pass
            events.append((t - t0, label))

    if not events:
        sys.exit("no messages found in bag")

    episodes, current = [], None
    last_t = None
    for t, label in events:
        if current is None or (last_t is not None and t - last_t > args.gap):
            if current:
                episodes.append(current)
            current = {"id": len(episodes) + 1, "task": args.task,
                       "result": "success",  # <-- EDIT per episode (or publish /eval/result)
                       "phases": []}
        ph = current["phases"]
        if not ph or ph[-1]["label"] != label:
            ph.append({"t": round(t, 2), "label": label})
        last_t = t
    if current:
        episodes.append(current)

    # re-base each episode's phase times to start at 0 and cap phases
    for e in episodes:
        if e["phases"]:
            base = e["phases"][0]["t"]
            e["phases"] = [{"t": round(p["t"] - base, 2), "label": p["label"]}
                           for p in e["phases"][:24]]
            e["duration"] = e["phases"][-1]["t"]

    Path(args.out).write_text(json.dumps({"episodes": episodes}, indent=1))
    print(f"wrote {args.out}: {len(episodes)} episodes "
          f"(now set task/result per episode, then upload at visionlibra.com/evals)")


if __name__ == "__main__":
    main()
