|
11 | 11 | Agent, |
12 | 12 | FunctionToolResult, |
13 | 13 | RunContextWrapper, |
| 14 | + Runner, |
14 | 15 | ToolCallOutputItem, |
15 | 16 | ToolsToFinalOutputResult, |
16 | 17 | UserError, |
17 | 18 | function_tool, |
18 | 19 | tool_namespace, |
19 | 20 | ) |
20 | 21 | from agents.run_internal import run_loop |
| 22 | +from agents.testing import ScriptedModel, function_call |
21 | 23 |
|
22 | 24 | from .test_responses import get_function_tool |
23 | 25 |
|
@@ -175,6 +177,54 @@ async def __call__( |
175 | 177 | assert behavior.calls == 1 |
176 | 178 |
|
177 | 179 |
|
| 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 | + |
178 | 228 | @pytest.mark.asyncio |
179 | 229 | async def test_invalid_tool_use_behavior_raises() -> None: |
180 | 230 | """If tool_use_behavior is invalid, we should raise a UserError.""" |
|
0 commit comments