mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
refactor(event_loop): route message appends through Agent._append_messages (#3131)
This commit is contained in:
@@ -18,7 +18,7 @@ from opentelemetry import trace as trace_api
|
||||
|
||||
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
|
||||
@@ -36,7 +36,7 @@ from ..types._events import (
|
||||
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,
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
yield ToolResultMessageEvent(message=_cancelled_msg)
|
||||
|
||||
agent.event_loop_metrics.end_cycle(cycle_start_time, cycle_trace)
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user