diff --git a/pslab/external/motor.py b/pslab/external/motor.py index 1fba065..d77e839 100644 --- a/pslab/external/motor.py +++ b/pslab/external/motor.py @@ -125,21 +125,32 @@ def import_timeline_from_csv(self, filepath: str) -> List[List[int]]: Returns ------- List[List[int]] - A timeline consisting of servo angle values per timestep. + A timeline consisting of servo angle values per timestep, with one + angle per servo of this arm. """ timeline = [] with open(filepath, mode="r", newline="") as csvfile: reader = csv.DictReader(csvfile) + columns = [f"Servo{i}" for i in range(1, RoboticArm.MAX_SERVOS + 1)] + if reader.fieldnames is None or any( + column not in reader.fieldnames for column in columns + ): + raise ValueError("CSV must contain the Servo1-Servo4 columns") for row in reader: angles = [] - for key in ["Servo1", "Servo2", "Servo3", "Servo4"]: - value = row[key] - if value == "null": + for column in columns: + value = row.get(column) + # Short rows from older exports leave trailing servos unset. + if value in (None, "", "null"): angles.append(None) else: angles.append(int(value)) - timeline.append(angles) + if any(angle is not None for angle in angles[len(self.servos) :]): + raise ValueError( + f"Timeline sets angles for more than {len(self.servos)} servos" + ) + timeline.append(angles[: len(self.servos)]) return timeline @@ -157,6 +168,12 @@ def export_timeline_to_csv( Directory path where the CSV file will be saved. The filename will include a timestamp to ensure uniqueness. """ + for i, row in enumerate(timeline): + if len(row) > RoboticArm.MAX_SERVOS: + raise ValueError( + f"Timestep {i} has more than {RoboticArm.MAX_SERVOS} angles" + ) + timestamp = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") filename = f"Robotic_Arm{timestamp}.csv" filepath = os.path.join(folderpath, filename) @@ -165,5 +182,7 @@ def export_timeline_to_csv( writer = csv.writer(csvfile) writer.writerow(["Timestep", "Servo1", "Servo2", "Servo3", "Servo4"]) for i, row in enumerate(timeline): + # Pad to four servos so every row matches the header. + row = list(row) + [None] * (RoboticArm.MAX_SERVOS - len(row)) pos = ["null" if val is None else val for val in row] writer.writerow([i] + pos) diff --git a/tests/test_robotic_arm.py b/tests/test_robotic_arm.py new file mode 100644 index 0000000..815440e --- /dev/null +++ b/tests/test_robotic_arm.py @@ -0,0 +1,79 @@ +"""Tests for pslab.external.motor.RoboticArm CSV import and export. + +These tests do not require a connected PSLab. +""" + +from unittest.mock import MagicMock + +import pytest + +from pslab.external.motor import RoboticArm, Servo + + +def make_arm(servo_count: int) -> RoboticArm: + pwm = MagicMock() + return RoboticArm([Servo(f"SQ{i + 1}", pwm) for i in range(servo_count)]) + + +def write_csv(path, rows): + header = "Timestep,Servo1,Servo2,Servo3,Servo4\n" + path.write_text(header + "".join(f"{row}\n" for row in rows)) + return str(path) + + +@pytest.mark.parametrize("servo_count", [1, 2, 3, 4]) +def test_export_import_round_trip(tmp_path, servo_count): + arm = make_arm(servo_count) + timeline = [ + [10 * (i + 1) for i in range(servo_count)], + [None] + [90] * (servo_count - 1), + ] + + arm.export_timeline_to_csv(timeline, str(tmp_path)) + (exported,) = tmp_path.glob("*.csv") + + assert arm.import_timeline_from_csv(str(exported)) == timeline + + +def test_export_pads_rows_to_four_servos(tmp_path): + make_arm(2).export_timeline_to_csv([[10, 20]], str(tmp_path)) + (exported,) = tmp_path.glob("*.csv") + + assert exported.read_text().splitlines()[1] == "0,10,20,null,null" + + +def test_imported_timeline_runs_on_a_smaller_arm(tmp_path): + arm = make_arm(2) + path = write_csv(tmp_path / "t.csv", ["0,10,20,null,null"]) + + arm.run_schedule(arm.import_timeline_from_csv(path), time_step=0) + + assert [servo.angle for servo in arm.servos] == [10, 20] + + +def test_import_accepts_short_rows(tmp_path): + path = write_csv(tmp_path / "t.csv", ["0,10,20"]) + + assert make_arm(2).import_timeline_from_csv(path) == [[10, 20]] + + +def test_import_rejects_angles_for_missing_servos(tmp_path): + path = write_csv(tmp_path / "t.csv", ["0,10,20,30,null"]) + + with pytest.raises(ValueError, match="more than 2 servos"): + make_arm(2).import_timeline_from_csv(path) + + +def test_import_rejects_missing_servo_columns(tmp_path): + path = tmp_path / "t.csv" + path.write_text("Timestep,Servo1,Servo2\n0,10,20\n") + + with pytest.raises(ValueError, match="Servo1-Servo4"): + make_arm(2).import_timeline_from_csv(str(path)) + + +def test_export_rejects_more_than_four_angles(tmp_path): + with pytest.raises(ValueError, match="more than 4 angles"): + make_arm(4).export_timeline_to_csv([[1, 2, 3, 4, 5]], str(tmp_path)) + + assert list(tmp_path.iterdir()) == []