diff --git a/burr/core/application.py b/burr/core/application.py index 415e0984e..8756cb7f2 100644 --- a/burr/core/application.py +++ b/burr/core/application.py @@ -158,6 +158,14 @@ def _remap_dunder_parameters( return inputs +def _remap_injected_inputs(method: Callable, inputs: Dict[str, Any]) -> Dict[str, Any]: + """Remaps the framework-injected ``__context``/``__tracer`` inputs to the name-mangled + parameter names a class-based action's method may use. See ``_remap_dunder_parameters``.""" + if "__context" in inputs or "__tracer" in inputs: + return _remap_dunder_parameters(method, inputs, ["__context", "__tracer"]) + return inputs + + def _run_function(function: Function, state: State, inputs: Dict[str, Any], name: str) -> dict: """Runs a function, returning the result of running the function. Note this restricts the keys in the state to only those that the @@ -176,9 +184,7 @@ def _run_function(function: Function, state: State, inputs: Dict[str, Any], name ) state_to_use = state.subset(*function.reads) function.validate_inputs(inputs) - if "__context" in inputs or "__tracer" in inputs: - # potentially need to remap the __context & __tracer variables - inputs = _remap_dunder_parameters(function.run, inputs, ["__context", "__tracer"]) + inputs = _remap_injected_inputs(function.run, inputs) result = function.run(state_to_use, **inputs) _validate_result(result, name) return result @@ -191,6 +197,7 @@ async def _arun_function( Async version of the above.""" state_to_use = state.subset(*function.reads) function.validate_inputs(inputs) + inputs = _remap_injected_inputs(function.run, inputs) result = await function.run(state_to_use, **inputs) _validate_result(result, name) return result @@ -477,7 +484,7 @@ def _run_multi_step_streaming_action( """ action.validate_inputs(inputs) stream_initialize_time = system.now() - generator = action.stream_run(state, **inputs) + generator = action.stream_run(state, **_remap_injected_inputs(action.stream_run, inputs)) result = None first_stream_start_time = None count = 0 @@ -535,7 +542,7 @@ async def _arun_multi_step_streaming_action( """Runs a multi-step streaming action in async. See the synchronous version for more details.""" action.validate_inputs(inputs) stream_initialize_time = system.now() - generator = action.stream_run(state, **inputs) + generator = action.stream_run(state, **_remap_injected_inputs(action.stream_run, inputs)) result = None first_stream_start_time = None count = 0 diff --git a/tests/core/test_application.py b/tests/core/test_application.py index 5a252f778..efb7d205b 100644 --- a/tests/core/test_application.py +++ b/tests/core/test_application.py @@ -4256,6 +4256,105 @@ def test_remap_context_variable_without_mangled_context(): assert _remap_dunder_parameters(_action.run, inputs, ["__context", "__tracer"]) == expected +class AsyncActionWithContext(Action): + """Class-based async action whose run() takes the injected ``__context``. + Python name-mangles the parameter to ``_AsyncActionWithContext__context``.""" + + @property + def reads(self) -> list[str]: + return [] + + @property + def writes(self) -> list[str]: + return ["app_id"] + + @property + def inputs(self) -> list[str]: + return ["__context"] + + async def run(self, state: State, __context: ApplicationContext) -> dict: + return {"app_id": __context.app_id} + + def update(self, result: dict, state: State) -> State: + return state.update(**result) + + +class StreamingActionWithContext(StreamingAction): + @property + def reads(self) -> list[str]: + return [] + + @property + def writes(self) -> list[str]: + return ["app_id"] + + @property + def inputs(self) -> list[str]: + return ["__context"] + + def stream_run( + self, state: State, __context: ApplicationContext + ) -> Generator[dict, None, None]: + yield {"app_id": __context.app_id} + + def update(self, result: dict, state: State) -> State: + return state.update(**result) + + +class AsyncStreamingActionWithContext(AsyncStreamingAction): + @property + def reads(self) -> list[str]: + return [] + + @property + def writes(self) -> list[str]: + return ["app_id"] + + @property + def inputs(self) -> list[str]: + return ["__context"] + + async def stream_run(self, state: State, __context: ApplicationContext) -> AsyncGenerator: + yield {"app_id": __context.app_id} + + def update(self, result: dict, state: State) -> State: + return state.update(**result) + + +def _build_context_app(action_: Action) -> Application: + return ( + ApplicationBuilder() + .with_actions(ctx_action=action_) + .with_transitions() + .with_entrypoint("ctx_action") + .with_identifiers(app_id="context-app-id") + .build() + ) + + +async def test_astep_class_based_action_receives_context(): + app = _build_context_app(AsyncActionWithContext()) + _, result, state = await app.astep() + assert result == {"app_id": "context-app-id"} + assert state["app_id"] == "context-app-id" + + +def test_stream_result_class_based_streaming_action_receives_context(): + app = _build_context_app(StreamingActionWithContext()) + _, container = app.stream_result(halt_after=["ctx_action"]) + result, state = container.get() + assert result == {"app_id": "context-app-id"} + assert state["app_id"] == "context-app-id" + + +async def test_astream_result_class_based_async_streaming_action_receives_context(): + app = _build_context_app(AsyncStreamingActionWithContext()) + _, container = await app.astream_result(halt_after=["ctx_action"]) + result, state = await container.get() + assert result == {"app_id": "context-app-id"} + assert state["app_id"] == "context-app-id" + + async def test_async_application_builder_initialize_raises_on_broken_persistor(): """Persisters should return None when there is no state to be loaded and the default used.""" await asyncio.sleep(0.00001)