diff --git a/burr/core/application.py b/burr/core/application.py index 415e0984e..7f06cec57 100644 --- a/burr/core/application.py +++ b/burr/core/application.py @@ -2037,6 +2037,21 @@ def _set_state(self, new_state: State[ApplicationStateType]): def get_next_action(self) -> Optional[Action]: return self._graph.get_next_node(self._state.get(PRIOR_STEP), self._state, self.entrypoint) + def get_prior_action(self) -> Optional[Action]: + """Returns the last action that ran, or None if nothing has run yet. + + In a pre_run_step hook this is the action before the current one, in a post_run_step + hook it is the action that just finished. For streaming actions it only changes once + the stream is fully consumed. Raises a ValueError if the prior action is not in the + graph anymore, for example after restoring state saved with an older graph. + + :return: The last action that ran, or None. + """ + prior_step = self._state.get(PRIOR_STEP) + if prior_step is None: + return None + return self._graph.get_action(prior_step) + def update_state(self, new_state: State[ApplicationStateType]): """Updates state -- this is meant to be called if you need to do anything with the state. For example: diff --git a/tests/core/test_application.py b/tests/core/test_application.py index 5a252f778..ad9e1e0f7 100644 --- a/tests/core/test_application.py +++ b/tests/core/test_application.py @@ -3184,6 +3184,49 @@ def test_app_get_next_step(): assert app.get_next_action().name == "counter_1" +def test_app_get_prior_action(): + counter_action_1 = base_counter_action.with_name("counter_1") + counter_action_2 = base_counter_action.with_name("counter_2") + app = Application( + state=State(), + entrypoint="counter_1", + partition_key="test", + uid="test-123", + sequence_id=0, + graph=Graph( + actions=[counter_action_1, counter_action_2], + transitions=[ + Transition(counter_action_1, counter_action_2, default), + Transition(counter_action_2, counter_action_1, default), + ], + ), + ) + assert app.get_prior_action() is None + app.step() + assert app.get_prior_action().name == "counter_1" + app.step() + assert app.get_prior_action().name == "counter_2" + app.reset_to_entrypoint() + assert app.get_prior_action() is None + + +def test_app_get_prior_action_not_in_graph(): + counter_action = base_counter_action.with_name("counter") + app = Application( + state=State({PRIOR_STEP: "removed"}), + entrypoint="counter", + partition_key="test", + uid="test-123", + sequence_id=0, + graph=Graph( + actions=[counter_action], + transitions=[Transition(counter_action, counter_action, default)], + ), + ) + with pytest.raises(ValueError): + app.get_prior_action() + + def test_application_builder_complete(): app = ( ApplicationBuilder()