| """Integration tests for the REST API, via FastAPI's TestClient (httpx).""" |
|
|
| from __future__ import annotations |
|
|
| import numpy as np |
| import pytest |
|
|
| pytest.importorskip("fastapi") |
|
|
| from fastapi.testclient import TestClient |
|
|
| from toaster.api import create_app, decode_array |
|
|
|
|
| @pytest.fixture |
| def client(tmp_path): |
| |
| rng = np.random.default_rng(0) |
| a = rng.normal((0, 0, 0), 0.08, (50, 3)) |
| b = rng.normal((10, 0, 0), 0.08, (50, 3)) |
| pts = np.hstack([np.vstack([a, b]), rng.random((100, 1))]).astype(np.float32) |
| path = tmp_path / "scan.bin" |
| pts.tofile(path) |
|
|
| app = create_app() |
| c = TestClient(app) |
| c.post("/api/open", json={"path": str(path)}) |
| return c |
|
|
|
|
| def test_state_before_open_is_conflict(): |
| client = TestClient(create_app()) |
| assert client.get("/api/state").status_code == 409 |
|
|
|
|
| def test_cloud_roundtrips_geometry(client): |
| payload = client.get("/api/cloud").json() |
| xyz = decode_array(payload["xyz"]) |
| assert xyz.shape == (100, 3) |
| assert xyz.dtype == np.float32 |
| assert "intensity" in payload["features"] |
|
|
|
|
| def test_state_shape(client): |
| state = client.get("/api/state").json() |
| assert decode_array(state["labels"]).shape == (100,) |
| assert state["grouping"] is None |
| assert state["snapshot"]["classes"][0]["name"] == "unlabeled" |
|
|
|
|
| def test_segment_then_assign_group(client): |
| state = client.post( |
| "/api/segment", |
| json={"name": "dbscan", "params": {"eps": 0.5, "min_samples": 5}}, |
| ).json() |
| snap = state["snapshot"] |
| assert len(snap["segments"]) == 2 |
| assert state["grouping"] is not None |
|
|
| gid = snap["segments"][0]["id"] |
| |
| sel = client.post("/api/group/select", json={"group_id": gid}).json() |
| assert decode_array(sel["selection"]).size == snap["segments"][0]["count"] |
|
|
| |
| after = client.post("/api/group/assign", json={"group_id": gid, "class_id": 4}).json() |
| delta = after["labels_delta"] |
| assert decode_array(delta["indices"]).size == snap["segments"][0]["count"] |
| assert (decode_array(delta["values"]) == 4).all() |
| |
| labels = decode_array(client.get("/api/state").json()["labels"]) |
| assert int((labels == 4).sum()) == snap["segments"][0]["count"] |
|
|
|
|
| def test_group_visibility_in_state(client): |
| state = client.post("/api/segment", json={"name": "dbscan", "params": {"eps": 0.5}}).json() |
| gid = state["snapshot"]["segments"][0]["id"] |
| hidden = client.post("/api/group/visibility", json={"group_id": gid, "visible": False}).json() |
| seg = next(s for s in hidden["snapshot"]["segments"] if s["id"] == gid) |
| assert seg["visible"] is False |
| allhidden = client.post("/api/groups/hide_all").json() |
| assert all(s["visible"] is False for s in allhidden["snapshot"]["segments"]) |
| shown = client.post("/api/groups/show_all").json() |
| assert all(s["visible"] for s in shown["snapshot"]["segments"]) |
|
|
|
|
| def test_assign_visible_groups_endpoint(client): |
| state = client.post("/api/segment", json={"name": "dbscan", "params": {"eps": 0.5}}).json() |
| segs = state["snapshot"]["segments"] |
| assert len(segs) == 2 |
| hide = segs[0] |
| |
| client.post("/api/group/visibility", json={"group_id": hide["id"], "visible": False}) |
| after = client.post("/api/groups/assign_visible", json={"class_id": 3}).json() |
| delta = after["labels_delta"] |
| |
| visible_count = sum(s["count"] for s in segs if s["id"] != hide["id"]) |
| assert decode_array(delta["indices"]).size == visible_count |
| assert (decode_array(delta["values"]) == 3).all() |
| labels = decode_array(client.get("/api/state").json()["labels"]) |
| assert int((labels == 3).sum()) == sum(s["count"] for s in segs if s["id"] != hide["id"]) |
|
|
|
|
| def test_clear_grouping_endpoint(client): |
| state = client.post("/api/segment", json={"name": "dbscan", "params": {"eps": 0.5}}).json() |
| assert state["grouping"] is not None |
| cleared = client.post("/api/grouping/clear").json() |
| assert cleared["grouping"] is None |
| assert cleared["snapshot"]["active_grouping"] is None |
| assert cleared["snapshot"]["display_mode"] == "labels" |
|
|
|
|
| def test_class_editing_endpoints(client): |
| s = client.post("/api/class/add", json={"name": "tree", "color": "#0a0b0c"}).json() |
| new = max(c["id"] for c in s["snapshot"]["classes"]) |
| assert s["snapshot"]["active_class"] == new |
|
|
| s = client.post("/api/class/rename", json={"class_id": new, "name": "shrub"}).json() |
| assert next(c for c in s["snapshot"]["classes"] if c["id"] == new)["name"] == "shrub" |
|
|
| s = client.post("/api/class/color", json={"class_id": new, "color": [1, 2, 3]}).json() |
| assert next(c for c in s["snapshot"]["classes"] if c["id"] == new)["color"] == [1, 2, 3] |
|
|
| s = client.post("/api/class/remove", json={"class_id": new}).json() |
| assert new not in [c["id"] for c in s["snapshot"]["classes"]] |
|
|
|
|
| def test_bad_segment_params_return_400(client): |
| r = client.post("/api/segment", json={"name": "dbscan", "params": {"min_samples": -22}}) |
| assert r.status_code == 400 |
|
|
|
|
| def test_pick_assign_undo(client): |
| client.post("/api/active_class", json={"class_id": 2}) |
| client.post("/api/pick", json={"index": 5}) |
| after = client.post("/api/assign", json={}).json() |
| assert decode_array(after["labels_delta"]["indices"]).tolist() == [5] |
| assert decode_array(after["labels_delta"]["values"]).tolist() == [2] |
| back = client.post("/api/undo").json() |
| assert decode_array(back["labels_delta"]["indices"]).tolist() == [5] |
| assert decode_array(back["labels_delta"]["values"]).tolist() == [0] |
| assert decode_array(client.get("/api/state").json()["labels"])[5] == 0 |
|
|
|
|
| def test_browse_lists_dir_and_flags_openable(client, tmp_path): |
| (tmp_path / "notes.txt").write_text("hi") |
| (tmp_path / "sub").mkdir() |
| data = client.get("/api/browse", params={"path": str(tmp_path)}).json() |
| by = {e["name"]: e for e in data["entries"]} |
| assert by["scan.bin"]["openable"] is True |
| assert by["notes.txt"]["openable"] is False |
| assert by["sub"]["is_dir"] and by["sub"]["openable"] |
| assert ".bin" in data["extensions"] |
|
|
|
|
| def test_browse_defaults_to_open_cloud_folder(client): |
| data = client.get("/api/browse").json() |
| assert any(e["name"] == "scan.bin" for e in data["entries"]) |
|
|
|
|
| def _open_scan(tmp_path): |
| pts = np.random.default_rng(0).random((20, 4)).astype(np.float32) |
| path = tmp_path / "scan.bin" |
| pts.tofile(path) |
| c = TestClient(create_app()) |
| c.post("/api/open", json={"path": str(path)}) |
| return c, path |
|
|
|
|
| def test_save_writes_labels_and_schema_sidecars(tmp_path): |
| c, path = _open_scan(tmp_path) |
| r = c.post("/api/save", json={}).json() |
| |
| assert (tmp_path / "scan_toaster.npy").exists() |
| schema = tmp_path / "scan_toaster_schema.yaml" |
| assert schema.exists() |
| assert r["saved"].endswith("scan_toaster.npy") |
| assert r["schema"].endswith("scan_toaster_schema.yaml") |
| |
| assert "cloud: scan.bin" in schema.read_text() |
| assert str(path) not in schema.read_text() |
|
|
|
|
| def test_save_to_a_chosen_base_path(tmp_path): |
| c, _ = _open_scan(tmp_path) |
| c.post("/api/save", json={"path": str(tmp_path / "myscan")}) |
| assert (tmp_path / "myscan_toaster.npy").exists() |
| assert (tmp_path / "myscan_toaster_schema.yaml").exists() |
|
|
|
|
| def test_mkdir_creates_folder_then_saves_into_it(tmp_path): |
| c, _ = _open_scan(tmp_path) |
| target = tmp_path / "runs" / "labels_v2" |
| r = c.post("/api/mkdir", json={"path": str(target)}).json() |
| assert target.is_dir() |
| assert r["path"].endswith("labels_v2") |
| |
| c.post("/api/save", json={"path": str(target / "scan")}) |
| assert (target / "scan_toaster.npy").exists() |
| assert (target / "scan_toaster_schema.yaml").exists() |
|
|
|
|
| def test_reopen_restores_labels_and_class_names(tmp_path): |
| c, path = _open_scan(tmp_path) |
| c.post("/api/class/rename", json={"class_id": 1, "name": "rock"}) |
| c.post("/api/box", json={"indices": [0, 1, 2]}) |
| c.post("/api/assign", json={"class_id": 1}) |
| c.post("/api/save", json={}) |
|
|
| |
| c2 = TestClient(create_app()) |
| c2.post("/api/open", json={"path": str(path)}) |
| state = c2.get("/api/state").json() |
| assert (decode_array(state["labels"])[:3] == 1).all() |
| names = {cl["id"]: cl["name"] for cl in state["snapshot"]["classes"]} |
| assert names[1] == "rock" |
|
|
|
|
| def test_browse_bad_dir_is_400(client): |
| assert client.get("/api/browse", params={"path": "/no/such/dir/zzz"}).status_code == 400 |
|
|