diff --git a/flashdreams/flashdreams/api_v2/loop.py b/flashdreams/flashdreams/api_v2/loop.py index 8a1ad759a..53f388c1c 100644 --- a/flashdreams/flashdreams/api_v2/loop.py +++ b/flashdreams/flashdreams/api_v2/loop.py @@ -182,6 +182,15 @@ def _begin_run( ) -> _LoopRunResult: """Prepare one step or return a lifecycle request to the caller.""" self._run_message_batch() + return self._incorporate_user_events(events, generation) + + @final + def _incorporate_user_events( + self, + events: UserInputEvents, + generation: int, + ) -> _LoopRunResult: + """Add newly available user events to the prepared run.""" self._pending_user_events.extend(events.get_events()) transition = _parse_lifecycle_events(self._pending_user_events) if transition is None and self._new_session_request is not None: @@ -321,6 +330,10 @@ def _run_model_loop( last_run_started = time.monotonic() if self._shutdown_event.is_set(): break + events, generation = event_buffer.read(reader_id) + run = self._incorporate_user_events(events, generation) + if run.step_index is None: + break step_started_at = time.monotonic() raw_result = self.step(run.step_index, self.user_events) step_elapsed_s = time.monotonic() - step_started_at diff --git a/flashdreams/test_v2/test_session_runner.py b/flashdreams/test_v2/test_session_runner.py index 064aa1c98..94abf08e8 100644 --- a/flashdreams/test_v2/test_session_runner.py +++ b/flashdreams/test_v2/test_session_runner.py @@ -545,6 +545,38 @@ def _lifecycle_event(event_type: type[UserInputEvent]) -> UserInputEvents: return UserInputEvents([event_type(timestamp=uint64(0))]) +def test_model_loop_collects_input_after_pacing( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = FakeSession(_session_desc(), CallLog()) + session.init() + session.model_loop.frequency = 1 + event_buffer = EventBuffer() + event_buffer.register(0) + + def cadence_wait(timeout: float) -> bool: + assert timeout > 0 + event_buffer.append(_key_event()) + return False + + monkeypatch.setattr(session.model_loop._shutdown_event, "wait", cadence_wait) + session.model_loop._run_model_loop( + event_buffer=event_buffer, + reader_id=0, + publish=lambda generation, results, elapsed: None, + max_steps=2, + ) + + assert session.observed_events[0].get_events() == [] + observed = session.observed_events[1].get_events() + assert len(observed) == 1 + assert isinstance(observed[0], KeyboardUserInputEvent) + assert (observed[0].key, observed[0].state) == ( + "a", + KeyboardInputState.PRESSED, + ) + + def test_run_session_presents_every_step_in_order() -> None: log = CallLog() session = FakeSession(_session_desc(), log)