rosdiff / tests /test_export.py
Chandra Kiran
Export runs as LeRobot datasets
b357d7d unverified
Raw History Blame Contribute Delete
2.18 kB
from conftest import complete
from worker.lerobot_export import features
def finished_run(client, backend, h, gen_id, n):
run_id = client.post("/simulate/run", json={"generation_id": gen_id}, headers=h).json()["run_id"]
complete(backend, f"job-{n}", run_id)
backend.jobs[f"job-{n}"]["output"]["files"]["trajectory"] = f"runs/{run_id}/trajectory.npz"
return run_id
def test_export_endpoint(run_env):
client, keys_, backend, services, gen_id = run_env
h = {"X-API-Key": keys_["alice"]}
runs = [finished_run(client, backend, h, gen_id, n) for n in (1, 2)]
r = client.post("/datasets/export", json={"run_ids": runs, "name": "g1-crates"}, headers=h)
assert r.status_code == 200, r.text
assert r.json()["kind"] == "export" and backend.functions[-1] == "export_lerobot"
sent = backend.submitted[-1]
assert sent["repo_id"] == "rosdiff/rosdiff-g1-crates" and sent["robot"] == "unitree_g1" and sent["private"]
assert [e["run_id"] for e in sent["episodes"]] == runs
assert sent["episodes"][0]["task"] == "G1 near a slope and a crate"
assert sent["episodes"][0]["trajectory_url"].endswith("trajectory.npz")
def test_export_rejects_unfinished_or_foreign_runs(run_env):
client, keys_, backend, services, gen_id = run_env
h = {"X-API-Key": keys_["alice"]}
pending = client.post("/simulate/run", json={"generation_id": gen_id}, headers=h).json()["run_id"]
r = client.post("/datasets/export", json={"run_ids": [pending], "name": "x1"}, headers=h)
assert r.status_code == 422 and "not finished" in r.text
r = client.post("/datasets/export", json={"run_ids": [pending], "name": "x1"}, headers={"X-API-Key": keys_["bob"]})
assert r.status_code == 404
assert (
client.post("/datasets/export", json={"run_ids": [pending], "name": "Bad Name"}, headers=h).status_code == 422
)
def test_features():
f = features(["a", "b"], 2, ["m"], with_video=True)
assert f["observation.state"]["shape"] == (2,) and f["action"]["names"] == ["m"]
assert f["observation.images.main"]["dtype"] == "video" and "observation.images.main" not in features(
["a"], 1, [], with_video=False
)