Spaces:
Paused
Paused
Download tests/test_export.py from chandrakiran06/rosdiff: direct link, hf CLI and curl.
- Browser
- Download file 2.18 kB
-
https://huggingface.co/spaces/chandrakiran06/rosdiff/resolve/main/tests/test_export.py
- Command line
-
hf download hf://spaces/chandrakiran06/rosdiff/tests/test_export.py
-
curl -L -o test_export.py https://huggingface.co/spaces/chandrakiran06/rosdiff/resolve/main/tests/test_export.py
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 | |
| ) | |