Skip to content

Commit 47b5c19

Browse files
authored
Ensure that the agent sees warnings about bad function calls (#52)
1 parent 02eb29a commit 47b5c19

5 files changed

Lines changed: 193 additions & 3 deletions

File tree

‎src/triframe_inspect/phases/actor.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@
2929
PhaseResult,
3030
TriframeSettings,
3131
TriframeStateSnapshot,
32+
WarningMessage,
3233
format_limit_info,
3334
)
3435
from triframe_inspect.util import get_content_str, generate_choices
@@ -152,6 +153,11 @@ def prepare_messages_for_actor(
152153
cast(ExecutedOption, executed_entry) if executed_entry else None,
153154
)
154155
history_messages.extend(processed_messages)
156+
elif history_entry.type == "warning":
157+
warning = cast(WarningMessage, history_entry)
158+
history_messages.append(
159+
ChatMessageUser(content=f"<warning>{warning.warning}</warning>")
160+
)
155161

156162
# Return messages in chronological order
157163
return messages + list(reversed(history_messages))

‎src/triframe_inspect/phases/process.py‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
PhaseResult,
1717
ToolOutput,
1818
TriframeStateSnapshot,
19+
WarningMessage,
1920
)
2021
from triframe_inspect.limits import calculate_limits
2122

@@ -152,6 +153,15 @@ async def execute_regular_tools(
152153
option_id: str,
153154
) -> PhaseResult:
154155
"""Execute a sequence of regular tool calls"""
156+
if not chosen_option.tool_calls:
157+
state.history.append(
158+
WarningMessage(
159+
type="warning",
160+
warning="No tool calls found in the last response",
161+
)
162+
)
163+
return {"next_phase": "advisor", "state": state}
164+
155165
tool_outputs: Dict[str, ToolOutput] = {}
156166

157167
for tool_call in chosen_option.tool_calls:

‎src/triframe_inspect/type_defs/state.py‎

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -157,14 +157,22 @@ class FinalRatings(BaseModel):
157157
best_rating: Rating # Store the best rating directly
158158

159159

160+
class WarningMessage(BaseModel):
161+
"""Represents a warning to be displayed to the agent"""
162+
163+
type: Literal["warning"]
164+
warning: str
165+
166+
160167
HistoryEntry = Union[
161168
AdvisorChoice,
162169
ActorOptions,
163170
ActorChoice,
164-
FinalRatings,
165-
ToolOutput,
166171
ExecutedOption,
172+
FinalRatings,
167173
Rating,
174+
ToolOutput,
175+
WarningMessage,
168176
]
169177

170178

‎tests/test_phases/test_actor.py‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@
3737
from triframe_inspect.type_defs.state import (
3838
ActorOptions,
3939
TriframeStateSnapshot,
40+
WarningMessage,
4041
)
4142

4243

@@ -373,10 +374,11 @@ async def test_actor_no_options(
373374

374375
@pytest.mark.asyncio
375376
async def test_actor_message_preparation(file_operation_history):
376-
"""Test that actor message preparation includes executed options and tool outputs"""
377+
"""Test that actor message preparation includes executed options, tool outputs, and warnings"""
377378
base_state = create_base_state()
378379
base_state.task_string = BASIC_TASK
379380
base_state.history.extend(file_operation_history)
381+
base_state.history.append(WarningMessage(type="warning", warning="hello"))
380382
messages = actor.prepare_messages_for_actor(base_state)
381383

382384
assert messages[0].role == "system"
@@ -427,6 +429,10 @@ async def test_actor_message_preparation(file_operation_history):
427429
)
428430
assert cat_output.tool_call_id == "cat_call"
429431

432+
warning_output = messages[-1]
433+
assert warning_output.role == "user"
434+
assert warning_output.content == "<warning>hello</warning>"
435+
430436
tool_outputs = [msg for msg in messages[2:] if isinstance(msg, ChatMessageTool)]
431437

432438
all_have_limit_info = all(

‎tests/test_phases/test_process.py‎

Lines changed: 160 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,160 @@
1+
import pytest
2+
from inspect_ai.tool import ToolCall
3+
4+
from tests.utils import create_base_state, create_task_state
5+
from triframe_inspect.phases.process import create_phase_request
6+
from triframe_inspect.type_defs.state import (
7+
ActorChoice,
8+
ActorOption,
9+
ActorOptions,
10+
TriframeStateSnapshot,
11+
WarningMessage,
12+
)
13+
14+
15+
def create_state_with_no_tool_calls() -> TriframeStateSnapshot:
16+
"""Create a state that simulates going through advisor and actor phases with no tool calls"""
17+
state = create_base_state(
18+
task_string="Test task with no tool calls",
19+
include_advisor=False,
20+
)
21+
state.settings["enable_advising"] = False
22+
23+
option = ActorOption(
24+
id="no_tools_option",
25+
content="This option has no tool calls",
26+
tool_calls=[], # Empty tool calls list
27+
)
28+
29+
actor_options = ActorOptions(
30+
type="actor_options",
31+
options_by_id={"no_tools_option": option}
32+
)
33+
34+
actor_choice = ActorChoice(
35+
type="actor_choice",
36+
option_id="no_tools_option",
37+
rationale="Selected option with no tool calls for testing",
38+
)
39+
40+
state.history = [actor_options, actor_choice]
41+
return state
42+
43+
44+
def create_state_with_tool_calls(tool_calls: list[ToolCall]) -> TriframeStateSnapshot:
45+
"""Create a state that simulates going through advisor and actor phases with tool calls"""
46+
state = create_base_state(
47+
task_string="Test task with tool calls",
48+
include_advisor=False,
49+
)
50+
51+
state.settings["enable_advising"] = False
52+
53+
option = ActorOption(
54+
id="with_tools_option",
55+
content="This option has tool calls",
56+
tool_calls=tool_calls,
57+
)
58+
59+
actor_options = ActorOptions(
60+
type="actor_options",
61+
options_by_id={"with_tools_option": option}
62+
)
63+
64+
actor_choice = ActorChoice(
65+
type="actor_choice",
66+
option_id="with_tools_option",
67+
rationale="Selected option with tool calls for testing",
68+
)
69+
70+
state.history = [actor_options, actor_choice]
71+
return state
72+
73+
74+
@pytest.mark.asyncio
75+
async def test_process_phase_no_tool_calls():
76+
"""Test that process phase adds warning when actor choice contains no tool calls"""
77+
state = create_state_with_no_tool_calls()
78+
task_state = create_task_state("Test task with no tool calls")
79+
80+
result = await create_phase_request(task_state, state)
81+
82+
assert result["next_phase"] == "advisor"
83+
assert result["state"] == state
84+
85+
warning_entries = [entry for entry in state.history if entry.type == "warning"]
86+
assert len(warning_entries) == 1
87+
88+
warning = warning_entries[0]
89+
assert isinstance(warning, WarningMessage)
90+
assert warning.warning == "No tool calls found in the last response"
91+
92+
assert len(state.history) == 3 # actor_options, actor_choice, warning
93+
assert state.history[0].type == "actor_options"
94+
assert state.history[1].type == "actor_choice"
95+
assert state.history[2].type == "warning"
96+
97+
98+
@pytest.mark.asyncio
99+
async def test_process_phase_with_invalid_tool_call():
100+
"""Test that process phase proceeds normally when actor choice contains tool calls"""
101+
state = create_state_with_tool_calls(
102+
tool_calls=[
103+
ToolCall(
104+
id="test_invalid_call",
105+
type="function",
106+
function="not_found",
107+
arguments={},
108+
parse_error=None,
109+
),
110+
],
111+
)
112+
task_state = create_task_state("Test task with invalid tool call")
113+
114+
result = await create_phase_request(task_state, state)
115+
116+
assert result["next_phase"] == "advisor"
117+
assert result["state"] == state
118+
119+
assert len(state.history) == 3 # actor_options, actor_choice, executed_option
120+
assert state.history[0].type == "actor_options"
121+
assert state.history[1].type == "actor_choice"
122+
assert state.history[2].type == "executed_option"
123+
124+
assert len(state.history[2].tool_outputs) == 1
125+
assert "Tool not_found not found" in (
126+
state.history[2].tool_outputs["test_invalid_call"].error
127+
)
128+
129+
130+
@pytest.mark.asyncio
131+
async def test_process_phase_with_submit_call():
132+
"""Test that process phase proceeds normally when actor choice contains tool calls"""
133+
state = create_state_with_tool_calls(
134+
tool_calls=[
135+
ToolCall(
136+
id="test_submit_call",
137+
type="function",
138+
function="submit",
139+
arguments={"answer": "Test answer"},
140+
parse_error=None,
141+
),
142+
],
143+
)
144+
task_state = create_task_state("Test task with tool calls")
145+
146+
result = await create_phase_request(task_state, state)
147+
148+
assert result["next_phase"] == "complete"
149+
assert result["state"] == state
150+
151+
warning_entries = [entry for entry in state.history if entry.type == "warning"]
152+
assert len(warning_entries) == 0
153+
154+
assert len(state.history) == 3 # actor_options, actor_choice, executed_option
155+
assert state.history[0].type == "actor_options"
156+
assert state.history[1].type == "actor_choice"
157+
assert state.history[2].type == "executed_option"
158+
159+
assert len(state.history[2].tool_outputs) == 1
160+
assert state.history[2].tool_outputs["test_submit_call"].output == "Test answer"

0 commit comments

Comments
 (0)