Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
524 changes: 41 additions & 483 deletions dimos/cli/commands/imitation.py

Large diffs are not rendered by default.

240 changes: 80 additions & 160 deletions dimos/cli/commands/test_imitation.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,192 +12,112 @@
# See the License for the specific language governing permissions and
# limitations under the License.

from pathlib import Path
from typing import Any

from pytest_mock import MockerFixture
from textual.widgets import Button, Static
import pytest
from typer.testing import CliRunner

from dimos.cli.commands.imitation import (
CollectionApp,
CollectionSession,
_default_dataset,
_default_recording,
imitation_app,
)
from dimos.constants import STATE_DIR
from dimos.imitation.workflows import get_workflow
from dimos.cli.commands.imitation import imitation_app
from dimos.imitation.collection.recording import RecordingSchema
from dimos.msgs.imitation_msgs.EpisodeStatus import EpisodeStatus
from dimos.robot.manipulators.openyam.collection import OPENYAM_TEACH_COLLECTION


def _status(
state: str = "idle", *, event: str = "init", saved: int = 0, discarded: int = 0
) -> EpisodeStatus:
return EpisodeStatus(
ts=1.0,
state=state, # type: ignore[arg-type]
episodes_saved=saved,
episodes_discarded=discarded,
last_event=event, # type: ignore[arg-type]
task_label="pick up the block",
)


def _session(mocker: MockerFixture) -> tuple[CollectionSession, Any, Any]:
driver = mocker.Mock()
monitor = mocker.Mock()
driver.get_module.return_value = monitor
monitor.get_status.return_value = _status()
monitor.command.side_effect = [
_status("recording", event="start"),
_status("idle", event="discard", discarded=1),
]
return CollectionSession(driver), driver, monitor


def test_cli_exposes_one_workflow_group_and_removes_old_commands() -> None:
from dimos.cli.dimos import main

result = CliRunner().invoke(main, ["--help"])

assert result.exit_code == 0
assert "imitation" in result.output
assert "collect" not in result.output
assert "dataprep" not in result.output


def test_imitation_help_exposes_the_complete_workflow() -> None:
def test_help_exposes_attached_controls_and_no_workflow_launcher():
result = CliRunner().invoke(imitation_app, ["--help"])

assert result.exit_code == 0
for command in ("list", "collect", "prepare", "inspect", "train", "run"):
for command in ("collect", "rollout", "prepare", "inspect", "train"):
assert command in result.output
assert CliRunner().invoke(imitation_app, ["list"]).exit_code == 2
assert CliRunner().invoke(imitation_app, ["run"]).exit_code == 2
assert CliRunner().invoke(imitation_app, ["collect", "--module", "foo"]).exit_code == 2


def test_list_does_not_load_hardware_modules(mocker: MockerFixture) -> None:
imported = mocker.patch("importlib.import_module")

result = CliRunner().invoke(imitation_app, ["list"])
@pytest.mark.parametrize(
("command", "app"),
[
("collect", "CollectionApp"),
("rollout", "RolloutApp"),
],
)
def test_live_commands_only_connect_and_detach(command, app, mocker):
driver = mocker.Mock()
driver.find_module_by_spec.return_value.get_status.return_value = EpisodeStatus(
ts=1.0, state="idle", episodes_saved=0, episodes_discarded=0
)
mocker.patch("dimos.cli.commands.imitation.Dimos.connect", return_value=driver)
ui = mocker.patch(f"dimos.cli.commands.imitation.{app}")
result = CliRunner().invoke(imitation_app, [command])
assert result.exit_code == 0, result.output
ui.return_value.run.assert_called_once_with()
driver.run.assert_not_called()
driver.stop.assert_called_once_with()

assert result.exit_code == 0
assert "openyam-teach" in result.output
assert "openyam-quest" in result.output
imported.assert_not_called()

@pytest.mark.parametrize(
"error",
[
LookupError("No deployed module matches EpisodeControlSpec"),
ValueError("Multiple modules match EpisodeControlSpec"),
],
)
def test_discovery_errors_close_connection_and_explain_failure(error, mocker):
driver = mocker.Mock()
driver.find_module_by_spec.side_effect = error
mocker.patch("dimos.cli.commands.imitation.Dimos.connect", return_value=driver)
result = CliRunner().invoke(imitation_app, ["collect"])
assert result.exit_code == 1
assert str(error) in result.output
driver.stop.assert_called_once_with()

def test_train_forwards_arguments_and_exit_code(mocker: MockerFixture) -> None:
completed = mocker.Mock(returncode=17)
run = mocker.patch("dimos.cli.commands.imitation.subprocess.run", return_value=completed)

result = CliRunner().invoke(
imitation_app,
["train", "--policy.type=act", "--dataset.repo_id=local/test"],
)
@pytest.fixture
def recording(tmp_path):
directory = tmp_path / "session"
directory.mkdir()
(directory / "schema.json").write_text(OPENYAM_TEACH_COLLECTION.to_schema().model_dump_json())
(directory / "recording.mcap").touch()
return directory

assert result.exit_code == 17
command = run.call_args.args[0]
assert command[-3:] == [
"lerobot-train",
"--policy.type=act",
"--dataset.repo_id=local/test",
]
assert run.call_args.kwargs == {"check": False}


def test_prepare_rejects_an_existing_output(tmp_path: Path) -> None:
recording = tmp_path / "session.mcap"
def test_prepare_uses_saved_schema_not_robot_lookup(recording, tmp_path, mocker):
output = tmp_path / "dataset"
output.mkdir()
prepare = mocker.patch("dimos.cli.commands.imitation.run_lerobot_dataprep", return_value=output)
result = CliRunner().invoke(imitation_app, ["prepare", str(recording), "--output", str(output)])
assert result.exit_code == 0, result.output
config = prepare.call_args.args[0]
assert config.source == str(recording / "recording.mcap")
assert config.observation == RecordingSchema.read(recording).observation
assert config.output.path == output


def test_prepare_rejects_existing_output(recording, tmp_path):
result = CliRunner().invoke(
imitation_app,
["prepare", "openyam-teach", str(recording), "--output", str(output)],
imitation_app, ["prepare", str(recording), "--output", str(tmp_path)]
)

assert result.exit_code == 2
assert "already exists" in result.output


def test_default_artifacts_live_in_state_and_are_unique() -> None:
workflow = get_workflow("openyam-teach")

first = _default_recording(workflow)
second = _default_recording(workflow)

assert first.parent == STATE_DIR / "recordings"
assert first.suffix == ".mcap"
assert first != second
assert _default_dataset(first) == STATE_DIR / "datasets" / first.stem


def test_session_routes_commands_and_stops_owned_stack(mocker: MockerFixture) -> None:
session, driver, monitor = _session(mocker)

assert session.command("toggle").state == "recording"
session.close()

monitor.command.assert_called_once_with("toggle")
driver.stop.assert_called_once_with()
def test_inspect_reads_recording_schema(recording, mocker):
inspect = mocker.patch(
"dimos.cli.commands.imitation.inspect_recording", return_value={"episodes": 2}
)
result = CliRunner().invoke(imitation_app, ["inspect", str(recording)])
assert result.exit_code == 0, result.output
assert '"episodes": 2' in result.output
assert inspect.call_args.args == (recording / "recording.mcap",)


def test_collect_stops_driver_when_stack_start_fails(
tmp_path: Path,
mocker: MockerFixture,
) -> None:
workflow = mocker.Mock(name="workflow")
workflow.name = "openyam-teach"
workflow.load_collection_builder.return_value = mocker.Mock(return_value="blueprint")
driver = mocker.Mock()
driver.run.side_effect = RuntimeError("hardware failed")
mocker.patch("dimos.cli.commands.imitation.get_workflow", return_value=workflow)
mocker.patch("dimos.cli.commands.imitation._require_idle_coordinator")
mocker.patch("dimos.cli.commands.imitation.Dimos", return_value=driver)

def test_train_forwards_arguments_and_exit_code(mocker):
run = mocker.patch(
"dimos.cli.commands.imitation.subprocess.run", return_value=mocker.Mock(returncode=17)
)
result = CliRunner().invoke(
imitation_app,
[
"collect",
"openyam-teach",
"--task",
"pick up block",
"--recording",
str(tmp_path / "session.mcap"),
],
imitation_app, ["train", "--policy.type=act", "--dataset.repo_id=local/test"]
)

assert result.exit_code == 1
assert "hardware failed" in result.output
driver.stop.assert_called_once_with()


def test_collection_app_guards_normal_quit_while_recording(mocker: MockerFixture) -> None:
session, _, monitor = _session(mocker)
app = CollectionApp(session, "openyam-teach")
mocker.patch.object(app, "_refresh")
exit_mock = mocker.patch.object(app, "exit")

app.action_toggle_recording()
app.action_quit()
app.action_discard()
app.action_quit()
app.action_quit()

assert monitor.command.call_args_list == [mocker.call("toggle"), mocker.call("discard")]
exit_mock.assert_called_once_with()


async def test_collection_dashboard_tracks_episode_state(mocker: MockerFixture) -> None:
session, _, monitor = _session(mocker)
app = CollectionApp(session, "openyam-teach")
mocker.patch.object(app, "set_interval")

async with app.run_test(size=(80, 24)) as pilot:
assert str(app.query_one("#state", Static).render()) == "READY"
await pilot.click("#toggle")
assert "RECORDING" in str(app.query_one("#state", Static).render())
assert app.query_one("#stop", Button).disabled
await pilot.click("#discard")
assert str(app.query_one("#state", Static).render()) == "READY"

assert monitor.command.call_args_list == [mocker.call("toggle"), mocker.call("discard")]
assert result.exit_code == 17
assert run.call_args.args[0][-3:] == [
"lerobot-train",
"--policy.type=act",
"--dataset.repo_id=local/test",
]
Loading
Loading