diff --git a/src/fastcs/demo/controllers.py b/src/fastcs/demo/controllers.py index 5926fc8ce..b39937eee 100755 --- a/src/fastcs/demo/controllers.py +++ b/src/fastcs/demo/controllers.py @@ -8,7 +8,7 @@ from fastcs.attributes import AttributeIO, AttributeIORef, AttrR, AttrRW, AttrW from fastcs.connections import IPConnection, IPConnectionSettings -from fastcs.controllers import Controller +from fastcs.controllers import Controller, ControllerVector from fastcs.datatypes import Enum, Float, Int, Waveform from fastcs.logging import logger from fastcs.methods import command, scan @@ -80,15 +80,16 @@ def __init__(self, settings: TemperatureControllerSettings) -> None: self._settings = settings - self._ramp_controllers: list[TemperatureRampController] = [] - for index in range(1, settings.num_ramp_controllers + 1): - controller = TemperatureRampController(index, self.connection) - self._ramp_controllers.append(controller) - self.add_sub_controller(f"R{index}", controller) + self.ramps = ControllerVector( + { + index: TemperatureRampController(index, self.connection) + for index in range(1, settings.num_ramp_controllers + 1) + } + ) @command() async def cancel_all(self) -> None: - for rc in self._ramp_controllers: + for rc in self.ramps.values(): await rc.enabled.put(OnOffEnum.Off, sync_setpoint=True) # TODO: The requests all get concatenated and the sim doesn't handle it await asyncio.sleep(0.1) @@ -118,14 +119,14 @@ async def update_voltages(self): await self.voltages.update(voltages) - for index, controller in enumerate(self._ramp_controllers): + for index, controller in self.ramps.items(): self.log_event( "Update voltages", topic=controller.voltage, query=query, response=voltages, ) - await controller.voltage.update(float(voltages[index])) + await controller.voltage.update(float(voltages[index - 1])) class TemperatureRampController(Controller): diff --git a/tests/demo/test_controllers.py b/tests/demo/test_controllers.py new file mode 100644 index 000000000..bd7775bdf --- /dev/null +++ b/tests/demo/test_controllers.py @@ -0,0 +1,60 @@ +from unittest.mock import AsyncMock + +import numpy as np +import pytest + +from fastcs.connections import IPConnectionSettings +from fastcs.controllers import ControllerVector +from fastcs.demo.controllers import ( + OnOffEnum, + TemperatureController, + TemperatureControllerSettings, + TemperatureRampController, +) + + +@pytest.fixture +def controller() -> TemperatureController: + settings = TemperatureControllerSettings( + num_ramp_controllers=4, + ip_settings=IPConnectionSettings(ip="localhost", port=25565), + ) + controller = TemperatureController(settings) + controller.post_initialise() + return controller + + +def test_ramps_is_controller_vector(controller: TemperatureController): + assert isinstance(controller.ramps, ControllerVector) + assert list(controller.ramps) == [1, 2, 3, 4] + for index, ramp in controller.ramps.items(): + assert isinstance(ramp, TemperatureRampController) + assert controller.ramps[index] is ramp + + +@pytest.mark.asyncio +async def test_cancel_all_disables_every_ramp(controller: TemperatureController): + puts = {} + for index, ramp in controller.ramps.items(): + puts[index] = AsyncMock() + ramp.enabled.put = puts[index] # type: ignore[method-assign] + + await controller.cancel_all() + + for put in puts.values(): + put.assert_awaited_once_with(OnOffEnum.Off, sync_setpoint=True) + + +@pytest.mark.asyncio +async def test_update_voltages_updates_waveform_and_each_ramp( + controller: TemperatureController, +): + controller.connection.send_query = AsyncMock(return_value="[1, 2, 3, 4]\r\n") + + await controller.update_voltages() + + np.testing.assert_array_equal( + controller.voltages.get(), np.array([1, 2, 3, 4], dtype=np.int32) + ) + for index, ramp in controller.ramps.items(): + assert ramp.voltage.get() == pytest.approx(float(index))