diff --git a/qtoggleserver/core/api/__init__.py b/qtoggleserver/core/api/__init__.py index d3734844..4077359c 100644 --- a/qtoggleserver/core/api/__init__.py +++ b/qtoggleserver/core/api/__init__.py @@ -46,11 +46,6 @@ def to_json(self) -> GenericJSONDict: return dict(error=self.code, **self.params) -class APIAccepted(Exception): - def __init__(self, response: Any = None) -> None: - self.response: Any = response - - class APIRequest: def __init__(self, handler: APIHandler) -> None: self.handler: APIHandler = handler diff --git a/qtoggleserver/core/api/funcs/ports.py b/qtoggleserver/core/api/funcs/ports.py index d6c764d5..b8d19649 100644 --- a/qtoggleserver/core/api/funcs/ports.py +++ b/qtoggleserver/core/api/funcs/ports.py @@ -375,10 +375,8 @@ async def patch_port_value(request: core_api.APIRequest, port_id: str, params: P if not await port.is_writable(): raise core_api.APIError(400, "read-only-port") - old_value = port.get_last_read_value() - try: - await port.transform_and_write_value(value) + await port.push_write_and_wait(value) except core_ports.PortTimeout as e: raise core_api.APIError(504, "port-timeout") from e except core_ports.PortError as e: @@ -389,12 +387,6 @@ async def patch_port_value(request: core_api.APIRequest, port_id: str, params: P # Transform any unhandled exception into APIError(500) raise core_api.APIError(500, "unexpected-error", message=str(e)) from e - # If port value hasn't really changed, use 202 Accepted to inform the consumer - current_value = port.get_last_read_value() - if (old_value == current_value) and (old_value != value): - port.debug("API supplied value hasn't been applied right away") - raise core_api.APIAccepted() - @core_api.api_call(core_api.ACCESS_LEVEL_NORMAL) async def patch_port_sequence(request: core_api.APIRequest, port_id: str, params: GenericJSONDict) -> None: diff --git a/qtoggleserver/core/ports.py b/qtoggleserver/core/ports.py index dfd41e39..b2f09b31 100644 --- a/qtoggleserver/core/ports.py +++ b/qtoggleserver/core/ports.py @@ -8,14 +8,13 @@ from collections import deque from collections.abc import AsyncIterator, Callable, ValuesView -from typing import Any +from typing import Any, NamedTuple from qtoggleserver import persist from qtoggleserver.conf import settings from qtoggleserver.core import events as core_events from qtoggleserver.core import expressions as core_expressions from qtoggleserver.core import history as core_history -from qtoggleserver.core import sequences as core_sequences from qtoggleserver.core.expressions import EvalContext from qtoggleserver.core.expressions import exceptions as expressions_exceptions from qtoggleserver.core.typing import ( @@ -26,10 +25,10 @@ NullablePortValue, PortValue, ) -from qtoggleserver.utils import asyncio as asyncio_utils from qtoggleserver.utils import dynload as dynload_utils from qtoggleserver.utils import json as json_utils from qtoggleserver.utils import logging as logging_utils +from qtoggleserver.utils import sequence as sequence_utils from qtoggleserver.utils.debounced import Debounced from qtoggleserver.utils.misc import append_traceback, stack_to_traceback @@ -150,6 +149,14 @@ class PortTimeout(PortError): pass +class WriteRequest(NamedTuple): + """An entry in a port's write queue, optionally carrying a future to be resolved with the result of + (or exception raised by) the corresponding `transform_and_write_value()` call.""" + + value: PortValue + future: asyncio.Future | None = None + + def skip_write_unavailable(func: Callable) -> Callable: @functools.wraps(func) async def wrapper(self: BasePort, value: NullablePortValue) -> None: @@ -218,7 +225,7 @@ def __init__(self, port_id: str) -> None: self._step: int | float | None = self.STEP self._internal: bool = self.INTERNAL - self._sequence: core_sequences.Sequence | None = None + self._sequence: sequence_utils.Sequence | None = None self._expression: core_expressions.Expression | None = None self._transform_read: core_expressions.Expression | None = None self._transform_write: core_expressions.Expression | None = None @@ -248,7 +255,7 @@ def __init__(self, port_id: str) -> None: self._last_written_value: tuple[NullablePortValue, int] | None = None self._write_value_lock = asyncio.Lock() - self._write_queue: deque[PortValue] = deque(maxlen=self.WRITE_QUEUE_SIZE) + self._write_queue: deque[WriteRequest] = deque(maxlen=self.WRITE_QUEUE_SIZE) self._write_task: asyncio.Task | None = None try: asyncio.get_running_loop() @@ -668,10 +675,7 @@ def set_last_read_value(self, value: NullablePortValue) -> None: async def read_transformed_value(self) -> NullablePortValue: value = None async with self._read_value_lock: - try: - value = await self.read_value() - except Exception: - raise + value = await self.read_value() if self._transform_read: eval_context = EvalContext(port_values={self.get_id(): value}) @@ -694,10 +698,8 @@ async def _write_value_safe(self, value: NullablePortValue) -> None: async with self._write_value_lock: try: self._writing_value = value - self.save_asap() await self.write_value(value) self._last_written_value = value, int(time.time() * 1000) - self.save_asap() finally: self._writing_value = None self.save_asap() @@ -706,13 +708,46 @@ def get_pending_value(self) -> NullablePortValue: """Return the most recent value that's about to be written to the port but hasn't been, yet.""" try: - return self._write_queue[-1] + return self._write_queue[-1].value except IndexError: return self._writing_value def get_last_written_value(self) -> NullablePortValue: return self._last_written_value[0] if self._last_written_value else None + def get_target_value(self) -> NullablePortValue: + """Return the value the port is expected to end up with, as a result of the writing process: the pending value, + if available, falling back to the last written value. + """ + + pending_value = self.get_pending_value() + if pending_value is not None: + return pending_value + + return self.get_last_written_value() + + def push_write(self, value: NullablePortValue, future: asyncio.Future | None = None) -> None: + """Push a value to the writing process queue. If `future` is given, it will be resolved with the result of + (`None`) or exception raised by the eventual `transform_and_write_value()` call for this value. It will be + cancelled instead if evicted from a full queue before ever being written, or if the port's write loop is + cancelled first.""" + + if len(self._write_queue) >= self._write_queue.maxlen: + evicted = self._write_queue[0] + if evicted.future and not evicted.future.done(): + evicted.future.cancel() + + self._write_queue.append(WriteRequest(value, future)) + self.save_asap() + + async def push_write_and_wait(self, value: NullablePortValue) -> None: + """Push a value to the writing process queue and wait for it to actually be written, propagating any + exception raised in the process.""" + + future = asyncio.get_running_loop().create_future() + self.push_write(value, future=future) + await future + async def eval_and_push_write(self, eval_context: EvalContext) -> None: """Evaluate the port's expression and push the resulting value to the write queue. Shield any evaluation exceptions from caller, but make sure to log them.""" @@ -738,29 +773,43 @@ async def eval_and_push_write(self, eval_context: EvalContext) -> None: json_utils.dumps(adapted_value), ) - # Only write value to port if it differs from the last value - if self.get_last_value() != adapted_value: - self._write_queue.append(adapted_value) - self.save_asap() + # Only write value to port if it differs from the target value + if self.get_target_value() != adapted_value: + self.push_write(adapted_value) async def _write_loop(self) -> None: + request: WriteRequest | None = None try: while True: try: - value = self._write_queue.pop() - self.save_asap() + request = self._write_queue.pop() + # No need to call `save_asap()` as it will be called indirectly by `transform_and_write_value()` except IndexError: + request = None await asyncio.sleep(settings.core.tick_interval / 1000.0) continue try: - await self.transform_and_write_value(value) - except Exception: + await self.transform_and_write_value(request.value) + except Exception as e: self.error("eval failed", exc_info=True) + if request.future and not request.future.done(): + request.future.set_exception(e) await asyncio.sleep(settings.core.tick_interval / 1000.0) + else: + if request.future and not request.future.done(): + request.future.set_result(None) except asyncio.CancelledError: self.debug("eval task cancelled") + # Don't leave any awaiter hanging: cancel the future of the request that was being processed, as well + # as those of any requests still sitting in the queue + if request and request.future and not request.future.done(): + request.future.cancel() + for pending_request in self._write_queue: + if pending_request.future and not pending_request.future.done(): + pending_request.future.cancel() + async def transform_and_write_value(self, value: NullablePortValue) -> None: """Apply write transform (if any) and write the value to the port.""" @@ -793,9 +842,8 @@ def get_last_value(self) -> NullablePortValue: return pending_value if self._last_written_value: - if self._last_read_value: - if self._last_read_value[1] > self._last_written_value[1]: - return self._last_read_value[0] + if self._last_read_value and self._last_read_value[1] > self._last_written_value[1]: + return self._last_read_value[0] return self._last_written_value[0] @@ -824,16 +872,11 @@ async def set_sequence(self, values: list[PortValue], delays: list[int], repeat: self._sequence = None if values: - self._sequence = core_sequences.Sequence( - values, delays, repeat, self.transform_and_write_value_fire_and_forget, self._on_sequence_finish - ) + self._sequence = sequence_utils.Sequence(values, delays, repeat, self.push_write, self._on_sequence_finish) self.debug("installing sequence") self._sequence.start() - def transform_and_write_value_fire_and_forget(self, value: NullablePortValue) -> None: - asyncio_utils.fire_and_forget(self.transform_and_write_value(value)) - async def _on_sequence_finish(self) -> None: self.debug("sequence finished") @@ -910,8 +953,7 @@ async def from_persisted(self, data: GenericJSONDict) -> None: for name, value in attr_items: if name in ( "id", - "pending_value", - "pending_queue", + "write_queue", "last_read_value", "last_read_timestamp", "last_written_value", @@ -944,15 +986,12 @@ async def from_persisted(self, data: GenericJSONDict) -> None: self._last_written_value = data["last_written_value"], data.get("last_written_timestamp", now_ms) self.debug("loaded last_written_value = %s", json_utils.dumps(data["last_written_value"])) - pending_queue = data.get("pending_queue") - if isinstance(pending_queue, list): + write_queue = data.get("write_queue") + if isinstance(write_queue, list): self._write_queue.clear() - for pending_value in pending_queue: - if pending_value is not None: - self._write_queue.append(pending_value) - - if data.get("pending_value") is not None and not isinstance(pending_queue, list): - self._write_queue.append(data["pending_value"]) + for value in write_queue: + if value is not None: + self._write_queue.append(WriteRequest(value)) # Handle legacy `value` field for backward compatibility with old persisted data if data.get("value") is not None: @@ -962,15 +1001,7 @@ async def from_persisted(self, data: GenericJSONDict) -> None: if await self.is_writable(): # Write the just-loaded value to the port - value = data["value"] - if self._transform_write: - eval_context = EvalContext(port_values={self.get_id(): value}) - try: - value = self.adapt_value_type(await self._transform_write.eval(eval_context)) - except expressions_exceptions.ValueUnavailable: - value = None - - await self.write_value(value) + await self.transform_and_write_value(data["value"]) elif not loaded_last_read and self.is_enabled(): try: value = await self.read_transformed_value() @@ -1007,8 +1038,7 @@ async def to_persisted(self) -> GenericJSONDict: d["last_read_timestamp"] = self._last_read_value[1] if self._last_read_value else None d["last_written_value"] = self._last_written_value[0] if self._last_written_value else None d["last_written_timestamp"] = self._last_written_value[1] if self._last_written_value else None - d["pending_value"] = self.get_pending_value() - d["pending_queue"] = list(self._write_queue) + d["write_queue"] = [request.value for request in self._write_queue] # attributes for name in await self.get_modifiable_attrs(): diff --git a/qtoggleserver/core/sequences.py b/qtoggleserver/utils/sequence.py similarity index 80% rename from qtoggleserver/core/sequences.py rename to qtoggleserver/utils/sequence.py index b9ca5543..8de84e05 100644 --- a/qtoggleserver/core/sequences.py +++ b/qtoggleserver/utils/sequence.py @@ -15,7 +15,14 @@ class SequenceError(Exception): class Sequence: def __init__( - self, values: list[PortValue], delays: list[int], repeat: int, callback: Callable, finish_callback: Callable + self, + values: list[PortValue], + delays: list[int], + repeat: int, + callback: Callable, + finish_callback: Callable, + callback_args: tuple = (), + callback_kwargs: dict | None = None, ) -> None: self._values: list[PortValue] = values self._delays: list[int] = delays @@ -23,6 +30,8 @@ def __init__( self._callback: Callable = callback self._finish_callback: Callable = finish_callback + self._callback_args: tuple = callback_args + self._callback_kwargs: dict = callback_kwargs or {} self._counter: int = 0 self._loop_task: asyncio.Task | None = None @@ -44,7 +53,7 @@ async def _loop(self) -> None: for i, value in enumerate(self._values): try: try: - self._callback(value) + self._callback(value, *self._callback_args, **self._callback_kwargs) except Exception as e: logger.error("sequence callback failed: %s", e, exc_info=True) diff --git a/qtoggleserver/web/base.py b/qtoggleserver/web/base.py index 113503cb..0efdeabb 100644 --- a/qtoggleserver/web/base.py +++ b/qtoggleserver/web/base.py @@ -177,11 +177,7 @@ async def _handle_api_call_exception(self, func: Callable, kwargs: dict, error: if not self._finished: # avoid finishing an already finished request await self.finish_json(error.to_json()) - if isinstance(error, core_api.APIAccepted): - self.set_status(202) - if not self._finished and error.response is not None: # avoid finishing an already finished request - await self.finish_json(error.response) - elif isinstance(error, StreamClosedError) and func.__name__ == "get_listen": + if isinstance(error, StreamClosedError) and func.__name__ == "get_listen": logger.debug("api call get_listen could not complete: stream closed") else: logger.error("api call %s failed: %s (args=%s, body=%s)", func.__name__, error, args, body, exc_info=True) diff --git a/tests/unit/qtoggleserver/core/api/test_funcs_ports.py b/tests/unit/qtoggleserver/core/api/test_funcs_ports.py index 9cf6a142..e12dd792 100644 --- a/tests/unit/qtoggleserver/core/api/test_funcs_ports.py +++ b/tests/unit/qtoggleserver/core/api/test_funcs_ports.py @@ -88,3 +88,94 @@ async def test_anonymous_user_permissions(self, mock_api_request_maker) -> None: with pytest.raises(core_api.APIError, match="authentication-required") as exc_info: await ports_api_funcs.patch_ports(request, []) assert exc_info.value.status == 401 + + +class TestPatchPortValue: + @pytest.fixture(autouse=True) + def mock_slaves(self, mocker) -> None: + mocker.patch("qtoggleserver.slaves.devices.get_all", return_value=[]) + + async def test_no_such_port(self, mock_api_request_maker) -> None: + request = mock_api_request_maker("PATCH", "/ports/nosuch/value", access_level=core_api.ACCESS_LEVEL_NORMAL) + with pytest.raises(core_api.APIError, match="no-such-port") as exc_info: + await ports_api_funcs.patch_port_value(request, "nosuch", 100) + assert exc_info.value.status == 404 + + async def test_queues_write_via_push_write_and_wait( + self, mock_api_request_maker, mock_num_port1, mock_persist_driver, mocker + ) -> None: + """Should route the write through the write queue (`push_write_and_wait`) rather than writing directly, and + wait for it to actually complete.""" + + mock_num_port1.set_writable(True) + spy = mocker.spy(mock_num_port1, "push_write_and_wait") + + request = mock_api_request_maker("PATCH", "/ports/nid1/value", access_level=core_api.ACCESS_LEVEL_NORMAL) + result = await ports_api_funcs.patch_port_value(request, "nid1", 100) + + assert result is None + spy.assert_called_once_with(100) + assert mock_num_port1.get_last_written_value() == 100 + + async def test_port_timeout_raises_504( + self, mock_api_request_maker, mock_num_port1, mock_persist_driver, mocker + ) -> None: + mock_num_port1.set_writable(True) + mocker.patch.object(mock_num_port1, "push_write_and_wait", side_effect=core_ports.PortTimeout()) + + request = mock_api_request_maker("PATCH", "/ports/nid1/value", access_level=core_api.ACCESS_LEVEL_NORMAL) + with pytest.raises(core_api.APIError, match="port-timeout") as exc_info: + await ports_api_funcs.patch_port_value(request, "nid1", 100) + assert exc_info.value.status == 504 + + async def test_port_error_raises_502( + self, mock_api_request_maker, mock_num_port1, mock_persist_driver, mocker + ) -> None: + mock_num_port1.set_writable(True) + mocker.patch.object(mock_num_port1, "push_write_and_wait", side_effect=core_ports.PortError("boom")) + + request = mock_api_request_maker("PATCH", "/ports/nid1/value", access_level=core_api.ACCESS_LEVEL_NORMAL) + with pytest.raises(core_api.APIError, match="port-error") as exc_info: + await ports_api_funcs.patch_port_value(request, "nid1", 100) + assert exc_info.value.status == 502 + + async def test_unexpected_error_raises_500( + self, mock_api_request_maker, mock_num_port1, mock_persist_driver, mocker + ) -> None: + mock_num_port1.set_writable(True) + mocker.patch.object(mock_num_port1, "push_write_and_wait", side_effect=ValueError("oops")) + + request = mock_api_request_maker("PATCH", "/ports/nid1/value", access_level=core_api.ACCESS_LEVEL_NORMAL) + with pytest.raises(core_api.APIError, match="unexpected-error") as exc_info: + await ports_api_funcs.patch_port_value(request, "nid1", 100) + assert exc_info.value.status == 500 + + async def test_succeeds_even_if_read_value_not_updated( + self, mock_api_request_maker, mock_num_port1, mock_persist_driver + ) -> None: + """Should return normally once the write completes, even if the port's read value doesn't (yet) reflect the + newly written target (e.g. not yet propagated back from hardware).""" + + mock_num_port1.set_writable(True) + mock_num_port1.set_last_read_value(100) + + request = mock_api_request_maker("PATCH", "/ports/nid1/value", access_level=core_api.ACCESS_LEVEL_NORMAL) + result = await ports_api_funcs.patch_port_value(request, "nid1", 200) + assert result is None + + async def test_disabled_port(self, mock_api_request_maker, mock_num_port1, mock_persist_driver) -> None: + mock_num_port1.set_writable(True) + mock_num_port1._enabled = False + + request = mock_api_request_maker("PATCH", "/ports/nid1/value", access_level=core_api.ACCESS_LEVEL_NORMAL) + with pytest.raises(core_api.APIError, match="port-disabled") as exc_info: + await ports_api_funcs.patch_port_value(request, "nid1", 100) + assert exc_info.value.status == 400 + + async def test_read_only_port(self, mock_api_request_maker, mock_num_port1, mock_persist_driver) -> None: + mock_num_port1.set_writable(False) + + request = mock_api_request_maker("PATCH", "/ports/nid1/value", access_level=core_api.ACCESS_LEVEL_NORMAL) + with pytest.raises(core_api.APIError, match="read-only-port") as exc_info: + await ports_api_funcs.patch_port_value(request, "nid1", 100) + assert exc_info.value.status == 400 diff --git a/tests/unit/qtoggleserver/core/test_ports.py b/tests/unit/qtoggleserver/core/test_ports.py index 01111f7e..a0e7bb30 100644 --- a/tests/unit/qtoggleserver/core/test_ports.py +++ b/tests/unit/qtoggleserver/core/test_ports.py @@ -1,6 +1,7 @@ import asyncio from qtoggleserver.core.expressions.exceptions import ValueUnavailable +from qtoggleserver.core.ports import WriteRequest class TestPortGetLastValue: @@ -39,6 +40,197 @@ def test_last_read(self, mock_num_port1, mocker): assert mock_num_port1.get_last_value() == 300 +class TestPortGetTargetValue: + def test_pending(self, mock_num_port1, mocker): + """Should return the pending value, since it's not None.""" + + mocker.patch.object(mock_num_port1, "get_pending_value", return_value=100) + mock_num_port1._last_written_value = None + assert mock_num_port1.get_target_value() == 100 + + mock_num_port1._last_written_value = (200, 2000) + assert mock_num_port1.get_target_value() == 100 + + def test_last_written(self, mock_num_port1, mocker): + """Should return the last written value, since there's no pending value.""" + + mocker.patch.object(mock_num_port1, "get_pending_value", return_value=None) + mock_num_port1._last_written_value = (200, 2000) + assert mock_num_port1.get_target_value() == 200 + + def test_unavailable(self, mock_num_port1, mocker): + """Should return `None`, since neither pending nor last written values are available.""" + + mocker.patch.object(mock_num_port1, "get_pending_value", return_value=None) + mock_num_port1._last_written_value = None + assert mock_num_port1.get_target_value() is None + + def test_last_read_ignored(self, mock_num_port1, mocker): + """Should ignore the last read value, however recent it may be.""" + + mocker.patch.object(mock_num_port1, "get_pending_value", return_value=None) + mock_num_port1._last_written_value = None + mock_num_port1._last_read_value = (300, 3000) + assert mock_num_port1.get_target_value() is None + + mock_num_port1._last_written_value = (200, 1000) # older timestamp + assert mock_num_port1.get_target_value() == 200 + + +class TestPortPushWrite: + def test(self, mock_num_port1, mocker): + """Should append the value to the write queue and mark the port for saving.""" + + mocker.patch.object(mock_num_port1, "save_asap") + mock_num_port1._write_queue.clear() + + mock_num_port1.push_write(100) + + assert list(mock_num_port1._write_queue) == [WriteRequest(100)] + mock_num_port1.save_asap.assert_called_once() + + def test_multiple(self, mock_num_port1, mocker): + """Should append every pushed value to the write queue, in order.""" + + mocker.patch.object(mock_num_port1, "save_asap") + mock_num_port1._write_queue.clear() + + mock_num_port1.push_write(100) + mock_num_port1.push_write(200) + + assert list(mock_num_port1._write_queue) == [WriteRequest(100), WriteRequest(200)] + assert mock_num_port1.save_asap.call_count == 2 + + def test_with_future(self, mock_num_port1, mocker): + """Should attach the given future to the queued write request.""" + + mocker.patch.object(mock_num_port1, "save_asap") + mock_num_port1._write_queue.clear() + future = asyncio.get_event_loop().create_future() + + mock_num_port1.push_write(100, future=future) + + assert list(mock_num_port1._write_queue) == [WriteRequest(100, future)] + + def test_evicts_oldest_when_full(self, mock_num_port1, mocker): + """Should cancel the future of the oldest request when it gets evicted because the queue is full.""" + + mocker.patch.object(mock_num_port1, "save_asap") + mock_num_port1._write_queue.clear() + evicted_future = asyncio.get_event_loop().create_future() + mock_num_port1.push_write(1, future=evicted_future) + for value in range(2, mock_num_port1._write_queue.maxlen + 1): + mock_num_port1.push_write(value) + + mock_num_port1.push_write(999) + + assert evicted_future.cancelled() + assert list(mock_num_port1._write_queue)[0].value == 2 + assert list(mock_num_port1._write_queue)[-1].value == 999 + + +class TestPortPushWriteAndWait: + async def test_resolves_on_success(self, mock_num_port1, mocker): + """Should return once the pushed value has actually been written, without raising.""" + + mock_num_port1._write_queue.clear() + mocker.patch.object(mock_num_port1, "transform_and_write_value", new=mocker.AsyncMock()) + + await mock_num_port1.push_write_and_wait(100) + + mock_num_port1.transform_and_write_value.assert_called_once_with(100) + + async def test_raises_on_failure(self, mock_num_port1, mocker): + """Should propagate any exception raised while writing the value.""" + + mock_num_port1._write_queue.clear() + mocker.patch.object( + mock_num_port1, "transform_and_write_value", new=mocker.AsyncMock(side_effect=ValueError("nope")) + ) + + try: + await mock_num_port1.push_write_and_wait(100) + assert False, "expected ValueError to be raised" + except ValueError as e: + assert str(e) == "nope" + + async def test_cleanup_cancels_pending_futures(self, mock_num_port1, mocker): + """Should cancel the future of the request currently being written, as well as those of any requests still + sitting in the queue, when the write loop is cancelled (e.g. on port removal).""" + + mock_num_port1._write_queue.clear() + started = asyncio.Event() + release = asyncio.Event() + + async def blocking_transform(value): + started.set() + await release.wait() + + mocker.patch.object(mock_num_port1, "transform_and_write_value", new=blocking_transform) + + future_a = asyncio.get_event_loop().create_future() + future_b = asyncio.get_event_loop().create_future() + mock_num_port1.push_write(1, future=future_a) + mock_num_port1.push_write(2, future=future_b) + + await started.wait() # one of the two requests is now being (indefinitely) written + await mock_num_port1.cleanup() + + assert future_a.cancelled() + assert future_b.cancelled() + + +class TestPortSetSequence: + async def test_pushes_each_value_via_push_write(self, mock_num_port1, mocker): + """Should install a sequence that pushes each value to the write queue via push_write.""" + + spy = mocker.spy(mock_num_port1, "push_write") + + await mock_num_port1.set_sequence([1, 2, 3], [1, 1, 1], 1) + loop_task = mock_num_port1._sequence._loop_task + await loop_task + + assert spy.call_args_list == [mocker.call(1), mocker.call(2), mocker.call(3)] + + async def test_clears_sequence_and_saves_when_finished(self, mock_num_port1, mocker): + """Should clear the sequence and mark the port for saving once the sequence finishes.""" + + mocker.patch.object(mock_num_port1, "save_asap") + + await mock_num_port1.set_sequence([1], [1], 1) + loop_task = mock_num_port1._sequence._loop_task + await loop_task + + assert mock_num_port1._sequence is None + mock_num_port1.save_asap.assert_called() + + async def test_cancels_previous_sequence(self, mock_num_port1, mocker): + """Should cancel any currently running sequence before installing a new one.""" + + await mock_num_port1.set_sequence([1, 2], [1000, 1000], 1) + old_sequence = mock_num_port1._sequence + spy_cancel = mocker.spy(old_sequence, "cancel") + + await mock_num_port1.set_sequence([3], [1], 1) + + spy_cancel.assert_called_once() + assert mock_num_port1._sequence is not old_sequence + + await mock_num_port1._sequence.cancel() + + async def test_empty_values_only_cancels(self, mock_num_port1, mocker): + """Should cancel any existing sequence and not install a new one when values is empty.""" + + await mock_num_port1.set_sequence([1, 2], [1000, 1000], 1) + old_sequence = mock_num_port1._sequence + spy_cancel = mocker.spy(old_sequence, "cancel") + + await mock_num_port1.set_sequence([], [], 1) + + spy_cancel.assert_called_once() + assert mock_num_port1._sequence is None + + class TestPortEvalAndPushWrite: async def test(self, mock_num_port1, mock_num_port2, mocker): """Should evaluate the expression with the provided eval context and push the result to the write queue.""" @@ -49,14 +241,47 @@ async def test(self, mock_num_port1, mock_num_port2, mocker): mocker.patch.object(mock_num_port1, "get_expression", return_value=mock_expression) mocker.patch.object(mock_num_port1, "adapt_value_type", return_value=100) - mocker.patch.object(mock_num_port1, "get_last_value", return_value=None) - mock_num_port1._write_queue = mocker.Mock() + mocker.patch.object(mock_num_port1, "get_target_value", return_value=None) + mocker.patch.object(mock_num_port1, "push_write") await mock_num_port1.eval_and_push_write(mock_eval_context) mock_expression.eval.assert_called_once_with(mock_eval_context) mock_num_port1.adapt_value_type.assert_called_once_with(mock_expression.eval.return_value) - mock_num_port1._write_queue.append.assert_called_once_with(100) + mock_num_port1.push_write.assert_called_once_with(100) + + async def test_same_as_target_value_not_written(self, mock_num_port1, mocker): + """Should not push anything to the write queue if the evaluated value equals the target value.""" + + mock_eval_context = mocker.Mock() + mock_expression = mocker.Mock() + mock_expression.eval = mocker.AsyncMock(return_value=99) + mocker.patch.object(mock_num_port1, "get_expression", return_value=mock_expression) + + mocker.patch.object(mock_num_port1, "adapt_value_type", return_value=100) + mocker.patch.object(mock_num_port1, "get_target_value", return_value=100) + mocker.patch.object(mock_num_port1, "push_write") + + await mock_num_port1.eval_and_push_write(mock_eval_context) + mock_num_port1.push_write.assert_not_called() + + async def test_differing_from_last_read_value_written(self, mock_num_port1, mocker): + """Should push the evaluated value to the write queue when it only matches the (more recent) last read value, + as the last read value plays no part in the target value.""" + + mock_eval_context = mocker.Mock() + mock_expression = mocker.Mock() + mock_expression.eval = mocker.AsyncMock(return_value=99) + mocker.patch.object(mock_num_port1, "get_expression", return_value=mock_expression) + + mocker.patch.object(mock_num_port1, "adapt_value_type", return_value=100) + mock_num_port1._writing_value = None + mock_num_port1._write_queue.clear() + mock_num_port1._last_written_value = (200, 1000) + mock_num_port1._last_read_value = (100, 2000) + + await mock_num_port1.eval_and_push_write(mock_eval_context) + assert list(mock_num_port1._write_queue) == [WriteRequest(100)] async def test_unavailable_not_written(self, mock_num_port1, mocker): """Should not push anything to the write queue if the expression evaluation raises due to value being @@ -65,10 +290,10 @@ async def test_unavailable_not_written(self, mock_num_port1, mocker): mock_eval_context = mocker.Mock() mock_num_port1._expression = mocker.Mock() mock_num_port1._expression.eval = mocker.AsyncMock(side_effect=ValueUnavailable) - mock_num_port1._write_queue = mocker.Mock() + mocker.patch.object(mock_num_port1, "push_write") await mock_num_port1.eval_and_push_write(mock_eval_context) - mock_num_port1._write_queue.append.assert_not_called() + mock_num_port1.push_write.assert_not_called() class TestPortGetPendingValue: @@ -82,9 +307,9 @@ def test_with_queue(self, mock_num_port1): """Should return the most recent value from writing queue.""" mock_num_port1._writing_value = 1 - mock_num_port1._write_queue.append(2) - mock_num_port1._write_queue.append(3) - mock_num_port1._write_queue.append(4) + mock_num_port1._write_queue.append(WriteRequest(2)) + mock_num_port1._write_queue.append(WriteRequest(3)) + mock_num_port1._write_queue.append(WriteRequest(4)) assert mock_num_port1.get_pending_value() == 4 def test_empty_queue(self, mock_num_port1): @@ -721,8 +946,8 @@ async def test_includes_state_fields(self, mock_num_port1, mocker): mock_num_port1._last_read_value = (42, 1111) mock_num_port1._last_written_value = (43, 2222) mock_num_port1._write_queue.clear() - mock_num_port1._write_queue.append(44) - mock_num_port1._write_queue.append(45) + mock_num_port1._write_queue.append(WriteRequest(44)) + mock_num_port1._write_queue.append(WriteRequest(45)) result = await mock_num_port1.to_persisted() @@ -730,8 +955,7 @@ async def test_includes_state_fields(self, mock_num_port1, mocker): assert result["last_read_timestamp"] == 1111 assert result["last_written_value"] == 43 assert result["last_written_timestamp"] == 2222 - assert result["pending_value"] == 45 - assert result["pending_queue"] == [44, 45] + assert result["write_queue"] == [44, 45] class TestPortFromPersisted: @@ -744,13 +968,54 @@ async def test_loads_state_fields(self, mock_num_port1, mocker): "last_read_timestamp": 1001, "last_written_value": 34, "last_written_timestamp": 1002, - "pending_queue": [56, 78], + "write_queue": [56, 78], } await mock_num_port1.from_persisted(data) assert mock_num_port1._last_read_value == (12, 1001) assert mock_num_port1._last_written_value == (34, 1002) - assert list(mock_num_port1._write_queue) == [56, 78] + assert list(mock_num_port1._write_queue) == [WriteRequest(56), WriteRequest(78)] assert mock_num_port1.get_pending_value() == 78 mock_num_port1.write_value.assert_not_called() + + async def test_legacy_value_field_written_via_transform_and_write_value(self, mock_num_port1, mocker): + """Should write a legacy `value` field through transform_and_write_value, rather than calling write_value + directly, so that the last written value (and other write-related state) end up correctly updated.""" + + mocker.patch.object(mock_num_port1, "is_writable", new=mocker.AsyncMock(return_value=True)) + mocker.patch.object(mock_num_port1, "write_value", new=mocker.AsyncMock()) + spy = mocker.spy(mock_num_port1, "transform_and_write_value") + + await mock_num_port1.from_persisted({"value": 42}) + + spy.assert_called_once_with(42) + mock_num_port1.write_value.assert_called_once_with(42) + assert mock_num_port1.get_last_written_value() == 42 + assert mock_num_port1._last_read_value[0] == 42 + + async def test_legacy_value_field_not_written_when_read_only(self, mock_num_port1, mocker): + """Should not attempt to write a legacy `value` field for a read-only port.""" + + mocker.patch.object(mock_num_port1, "is_writable", new=mocker.AsyncMock(return_value=False)) + mocker.patch.object(mock_num_port1, "write_value", new=mocker.AsyncMock()) + + await mock_num_port1.from_persisted({"value": 42}) + + mock_num_port1.write_value.assert_not_called() + assert mock_num_port1.get_last_written_value() is None + assert mock_num_port1._last_read_value[0] == 42 + + async def test_legacy_value_field_applies_write_transform(self, mock_num_port1, mocker): + """Should apply the write transform (if any) before writing a legacy `value` field, same as any other + write, and reflect the transformed value as the last written value.""" + + mocker.patch.object(mock_num_port1, "is_writable", new=mocker.AsyncMock(return_value=True)) + mocker.patch.object(mock_num_port1, "write_value", new=mocker.AsyncMock()) + mock_num_port1._transform_write = mocker.Mock() + mock_num_port1._transform_write.eval = mocker.AsyncMock(return_value=99) + + await mock_num_port1.from_persisted({"value": 42}) + + mock_num_port1.write_value.assert_called_once_with(99) + assert mock_num_port1.get_last_written_value() == 99 diff --git a/tests/unit/qtoggleserver/utils/test_sequence.py b/tests/unit/qtoggleserver/utils/test_sequence.py new file mode 100644 index 00000000..118ae3a6 --- /dev/null +++ b/tests/unit/qtoggleserver/utils/test_sequence.py @@ -0,0 +1,127 @@ +import asyncio + +import pytest + +from qtoggleserver.utils.sequence import Sequence, SequenceError + + +class TestSequenceRun: + async def test_calls_callback_for_each_value_in_order(self, mocker): + """Should invoke the callback once per value, in order, then the finish callback.""" + + callback = mocker.Mock() + finish_callback = mocker.AsyncMock() + + sequence = Sequence([1, 2, 3], [1, 1, 1], 1, callback, finish_callback) + sequence.start() + await sequence._loop_task + + assert callback.call_args_list == [mocker.call(1), mocker.call(2), mocker.call(3)] + finish_callback.assert_awaited_once() + + async def test_passes_custom_args_and_kwargs_to_callback(self, mocker): + """Should forward the given callback_args/callback_kwargs to every callback invocation, after the value.""" + + callback = mocker.Mock() + finish_callback = mocker.AsyncMock() + + sequence = Sequence( + [1, 2], + [1, 1], + 1, + callback, + finish_callback, + callback_args=("extra", 42), + callback_kwargs={"flag": True}, + ) + sequence.start() + await sequence._loop_task + + assert callback.call_args_list == [ + mocker.call(1, "extra", 42, flag=True), + mocker.call(2, "extra", 42, flag=True), + ] + + async def test_repeats_given_number_of_times(self, mocker): + """Should replay the entire list of values `repeat` times before calling the finish callback.""" + + callback = mocker.Mock() + finished = asyncio.Event() + + async def finish_callback(): + finished.set() + + sequence = Sequence([1, 2], [1, 1], 2, callback, finish_callback) + sequence.start() + await asyncio.wait_for(finished.wait(), timeout=1) + + assert callback.call_args_list == [mocker.call(1), mocker.call(2), mocker.call(1), mocker.call(2)] + + async def test_infinite_repeat_never_calls_finish_callback(self, mocker): + """A `repeat` of 0 should keep looping forever and never invoke the finish callback.""" + + callback = mocker.Mock() + finish_callback = mocker.AsyncMock() + + sequence = Sequence([1, 2], [1, 1], 0, callback, finish_callback) + sequence.start() + + # Let a few passes go by + while callback.call_count < 6: + await asyncio.sleep(0.01) + + finish_callback.assert_not_called() + + await sequence.cancel() + + async def test_callback_exception_is_caught_and_logged(self, mocker): + """A failing callback should not interrupt the sequence.""" + + callback = mocker.Mock(side_effect=[ValueError("boom"), None]) + finish_callback = mocker.AsyncMock() + spy_error = mocker.patch("qtoggleserver.utils.sequence.logger.error") + + sequence = Sequence([1, 2], [1, 1], 1, callback, finish_callback) + sequence.start() + await sequence._loop_task + + assert callback.call_args_list == [mocker.call(1), mocker.call(2)] + spy_error.assert_called_once() + finish_callback.assert_awaited_once() + + +class TestSequenceStart: + async def test_raises_if_already_started(self, mocker): + """Should raise SequenceError when start() is called while a loop task is already running.""" + + sequence = Sequence([1], [1], 1, mocker.Mock(), mocker.AsyncMock()) + sequence.start() + + with pytest.raises(SequenceError): + sequence.start() + + await sequence.cancel() + + +class TestSequenceCancel: + async def test_stops_further_callbacks(self, mocker): + """Should stop the sequence before it reaches subsequent values once cancelled.""" + + callback_ran = asyncio.Event() + callback = mocker.Mock(side_effect=lambda value: callback_ran.set()) + finish_callback = mocker.AsyncMock() + + sequence = Sequence([1, 2, 3], [1000, 1, 1], 1, callback, finish_callback) + sequence.start() + + await asyncio.wait_for(callback_ran.wait(), timeout=1) + await sequence.cancel() + + assert callback.call_count == 1 + finish_callback.assert_not_awaited() + + async def test_noop_when_not_started(self): + """Should do nothing if there's no loop task running.""" + + sequence = Sequence([1], [1], 1, lambda value: None, None) + await sequence.cancel() # should not raise