Skip to content
Merged
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
16 changes: 5 additions & 11 deletions strands-py/src/strands/event_loop/event_loop.py
Comment thread
opieter-aws marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@

from .._middleware.stages import InvokeModelContext, InvokeModelStage
from ..experimental.checkpoint import Checkpoint, CheckpointPosition
from ..hooks import AfterModelCallEvent, BeforeModelCallEvent, MessageAddedEvent
from ..hooks import AfterModelCallEvent, BeforeModelCallEvent
from ..telemetry.metrics import Trace
from ..telemetry.tracer import Tracer, get_tracer
from ..tools._validator import validate_and_prepare_tools
Expand All @@ -36,7 +36,7 @@
TypedEvent,
)
from ..types.agent import Limits
from ..types.content import Message, Messages, _ensure_tracking_id, split_system_prompt
from ..types.content import Message, Messages, split_system_prompt
from ..types.event_loop import Metrics, Usage
from ..types.exceptions import (
ContextWindowOverflowException,
Expand Down Expand Up @@ -637,9 +637,7 @@ async def _handle_model_execution(
stream_trace.end()

# Add the response message to the conversation
_ensure_tracking_id(message)
agent.messages.append(message)
await agent.hooks.invoke_callbacks_async(MessageAddedEvent(agent=agent, message=message))
await agent._append_messages(message)

# Update metrics
agent.event_loop_metrics.update_usage(usage)
Expand Down Expand Up @@ -771,9 +769,7 @@ async def _handle_tool_execution(
"content": [{"toolResult": result} for result in tool_results],
}
cancelled_tool_result_message = _cancelled_msg
_ensure_tracking_id(_cancelled_msg)
agent.messages.append(_cancelled_msg)
await agent.hooks.invoke_callbacks_async(MessageAddedEvent(agent=agent, message=_cancelled_msg))
await agent._append_messages(_cancelled_msg)
Comment thread
opieter-aws marked this conversation as resolved.
yield ToolResultMessageEvent(message=_cancelled_msg)

agent.event_loop_metrics.end_cycle(cycle_start_time, cycle_trace)
Expand Down Expand Up @@ -833,9 +829,7 @@ async def _handle_tool_execution(
"content": [{"toolResult": result} for result in tool_results],
}

_ensure_tracking_id(tool_result_message)
agent.messages.append(tool_result_message)
await agent.hooks.invoke_callbacks_async(MessageAddedEvent(agent=agent, message=tool_result_message))
await agent._append_messages(tool_result_message)

yield ToolResultMessageEvent(message=tool_result_message)

Expand Down
4 changes: 4 additions & 0 deletions strands-py/tests/strands/event_loop/test_event_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,6 +167,9 @@ def agent(model, system_prompt, messages, tool_registry, thread_pool, hook_regis
mock._checkpoint_resume_position = None
mock.trace_attributes = {}
mock.retry_strategy = ModelRetryStrategy()
# Bind the real _append_messages chokepoint so appends assign tracking ids
# and fire MessageAddedEvent exactly as production does.
mock._append_messages = Agent._append_messages.__get__(mock, Agent)

return mock

Expand Down Expand Up @@ -928,6 +931,7 @@ async def test_request_state_initialization(alist):
mock_agent.tool_registry.get_all_tool_specs.return_value = []
mock_agent.event_loop_metrics.start_cycle.return_value = (0, MagicMock())
mock_agent.hooks.invoke_callbacks_async = AsyncMock()
mock_agent._append_messages = Agent._append_messages.__get__(mock_agent, Agent)

# Call without providing request_state
stream = strands.event_loop.event_loop.event_loop_cycle(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,9 @@ def agent(model, messages, tool_registry, hook_registry):
mock._middleware_registry = strands._middleware.MiddlewareRegistry()
mock.trace_attributes = {}
mock.retry_strategy = ModelRetryStrategy()
# Bind the real _append_messages chokepoint so appends assign tracking ids
# and fire MessageAddedEvent exactly as production does.
mock._append_messages = Agent._append_messages.__get__(mock, Agent)
return mock


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -205,12 +205,10 @@ async def test_event_loop_forces_structured_output_on_end_turn(
)
await alist(stream)

# Should have appended a message to force structured output
mock_agent._append_messages.assert_called_once()
args = mock_agent._append_messages.call_args[0][0]
assert args["role"] == "user"
# Should use the default prompt
assert args["content"][0]["text"] == DEFAULT_STRUCTURED_OUTPUT_PROMPT
# The force-structured-output prompt should have been appended (among other messages)
appended_messages = [call.args[0] for call in mock_agent._append_messages.call_args_list]
expected_force_prompt = {"role": "user", "content": [{"text": DEFAULT_STRUCTURED_OUTPUT_PROMPT}]}
assert appended_messages.count(expected_force_prompt) == 1

# Should have called recurse_event_loop with the context
mock_recurse.assert_called_once()
Expand Down Expand Up @@ -260,11 +258,10 @@ async def test_event_loop_forces_structured_output_with_custom_prompt(mock_agent
)
await alist(stream)

# Should have appended a message with the custom prompt
mock_agent._append_messages.assert_called_once()
args = mock_agent._append_messages.call_args[0][0]
assert args["role"] == "user"
assert args["content"][0]["text"] == custom_prompt
# The custom force prompt should have been appended (among other messages)
appended_messages = [call.args[0] for call in mock_agent._append_messages.call_args_list]
expected_force_prompt = {"role": "user", "content": [{"text": custom_prompt}]}
assert appended_messages.count(expected_force_prompt) == 1


@patch("strands.event_loop.event_loop.get_tracer")
Expand Down
Loading