Skip to content

Commit 5a2d63f

Browse files
fix(core): validate tool_use_behavior callback results (#5173)
fix: validate tool_use_behavior callback results
1 parent a9b1ce7 commit 5a2d63f

2 files changed

Lines changed: 55 additions & 1 deletion

File tree

‎src/agents/run_internal/turn_resolution.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -781,7 +781,11 @@ async def check_for_final_output_from_tools(
781781
elif callable(agent.tool_use_behavior):
782782
result = agent.tool_use_behavior(context_wrapper, tool_results)
783783
if inspect.isawaitable(result):
784-
return await result
784+
result = await result
785+
if not isinstance(result, ToolsToFinalOutputResult):
786+
raise UserError(
787+
"Agent tool_use_behavior callable must return ToolsToFinalOutputResult."
788+
)
785789
return result
786790

787791
logger.error("Invalid tool_use_behavior: %s", agent.tool_use_behavior)

‎tests/test_tool_use_behavior.py‎

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,13 +11,15 @@
1111
Agent,
1212
FunctionToolResult,
1313
RunContextWrapper,
14+
Runner,
1415
ToolCallOutputItem,
1516
ToolsToFinalOutputResult,
1617
UserError,
1718
function_tool,
1819
tool_namespace,
1920
)
2021
from agents.run_internal import run_loop
22+
from agents.testing import ScriptedModel, function_call
2123

2224
from .test_responses import get_function_tool
2325

@@ -175,6 +177,54 @@ async def __call__(
175177
assert behavior.calls == 1
176178

177179

180+
@pytest.mark.asyncio
181+
@pytest.mark.parametrize("streamed", [False, True])
182+
@pytest.mark.parametrize("async_behavior", [False, True])
183+
@pytest.mark.parametrize("valid_result", [False, True])
184+
async def test_custom_tool_use_behavior_return_contract(
185+
streamed: bool, async_behavior: bool, valid_result: bool
186+
) -> None:
187+
calls = 0
188+
189+
def behavior(context: RunContextWrapper, results: list[FunctionToolResult]) -> Any:
190+
nonlocal calls
191+
calls += 1
192+
assert len(results) == 1
193+
assert results[0].output == "pong"
194+
if valid_result:
195+
return ToolsToFinalOutputResult(is_final_output=True, final_output="done")
196+
return "done"
197+
198+
async def async_callback(context: RunContextWrapper, results: list[FunctionToolResult]) -> Any:
199+
return behavior(context, results)
200+
201+
model = ScriptedModel([[function_call("ping", {}, call_id="call_1")]])
202+
agent = Agent(
203+
name="test",
204+
tools=[get_function_tool("ping", return_value="pong")],
205+
model=model,
206+
tool_use_behavior=async_callback if async_behavior else behavior,
207+
)
208+
209+
async def run() -> Any:
210+
if streamed:
211+
result = Runner.run_streamed(agent, "test")
212+
async for _ in result.stream_events():
213+
pass
214+
return result.final_output
215+
return (await Runner.run(agent, "test")).final_output
216+
217+
if valid_result:
218+
assert await run() == "done"
219+
else:
220+
with pytest.raises(
221+
UserError, match="tool_use_behavior callable must return ToolsToFinalOutputResult"
222+
):
223+
await run()
224+
assert calls == 1
225+
assert len(model.calls) == 1
226+
227+
178228
@pytest.mark.asyncio
179229
async def test_invalid_tool_use_behavior_raises() -> None:
180230
"""If tool_use_behavior is invalid, we should raise a UserError."""

0 commit comments

Comments
 (0)