Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 12 additions & 5 deletions burr/core/application.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
99 changes: 99 additions & 0 deletions tests/core/test_application.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading