diff --git a/src/agents/util/_approvals.py b/src/agents/util/_approvals.py index 8992f5ada2..a2b3b3b203 100644 --- a/src/agents/util/_approvals.py +++ b/src/agents/util/_approvals.py @@ -42,7 +42,9 @@ async def evaluate_needs_approval_setting( maybe_result = needs_approval_setting(*args) if inspect.isawaitable(maybe_result): maybe_result = await maybe_result - return bool(maybe_result) + if isinstance(maybe_result, bool): + return maybe_result + return True if strict: raise UserError( f"Invalid needs_approval value: expected a bool or callable, " diff --git a/tests/test_approval_utils.py b/tests/test_approval_utils.py new file mode 100644 index 0000000000..ff40d5b334 --- /dev/null +++ b/tests/test_approval_utils.py @@ -0,0 +1,29 @@ +import pytest + +from agents.util._approvals import evaluate_needs_approval_setting + + +@pytest.mark.asyncio +@pytest.mark.parametrize("result", [None, 0, "", [], {}]) +async def test_callable_non_bool_falsy_results_fail_closed(result: object) -> None: + def policy() -> object: + return result + + assert await evaluate_needs_approval_setting(policy) is True + + +@pytest.mark.asyncio +async def test_async_callable_non_bool_result_fails_closed() -> None: + async def policy() -> None: + return None + + assert await evaluate_needs_approval_setting(policy) is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize("result", [True, False]) +async def test_callable_bool_results_are_preserved(result: bool) -> None: + def policy() -> bool: + return result + + assert await evaluate_needs_approval_setting(policy) is result