Support runtime guidance tool interrupts
This commit is contained in:
+384
-1
@@ -7,6 +7,7 @@ import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from src.agent_tool_core import ToolStreamUpdate
|
||||
from src.agent_session import AgentMessage
|
||||
from src.agent_runtime import LocalCodingAgent, MAX_AUTO_CONTINUATIONS
|
||||
from src.agent_tools import build_tool_context, default_tool_registry, execute_tool
|
||||
@@ -16,6 +17,7 @@ from src.agent_types import (
|
||||
BudgetConfig,
|
||||
ModelConfig,
|
||||
OutputSchemaConfig,
|
||||
ToolExecutionResult,
|
||||
UsageStats,
|
||||
)
|
||||
from src.compact import CompactionResult
|
||||
@@ -282,9 +284,324 @@ class AgentRuntimeTests(unittest.TestCase):
|
||||
|
||||
self.assertEqual(result.final_output, 'The file contains hello world.')
|
||||
self.assertEqual(result.tool_calls, 1)
|
||||
self.assertGreaterEqual(len(result.transcript), 5)
|
||||
transcript_roles = [entry.get('role') for entry in result.transcript]
|
||||
self.assertEqual(transcript_roles, ['user', 'assistant', 'tool', 'assistant'])
|
||||
self.assertIn('I will inspect the file first.', result.transcript[1]['content'])
|
||||
self.assertIn('hello world', result.transcript[2]['content'])
|
||||
self.assertIn('The file contains hello world.', result.transcript[3]['content'])
|
||||
self.assertGreaterEqual(len(result.file_history), 0)
|
||||
|
||||
def test_runtime_guidance_before_tools_skips_stale_tool_plan(self) -> None:
|
||||
responses = [
|
||||
{
|
||||
'choices': [
|
||||
{
|
||||
'message': {
|
||||
'role': 'assistant',
|
||||
'content': 'I will inspect the old file.',
|
||||
'tool_calls': [
|
||||
{
|
||||
'id': 'call_1',
|
||||
'type': 'function',
|
||||
'function': {
|
||||
'name': 'read_file',
|
||||
'arguments': '{"path": "old.txt"}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
'finish_reason': 'tool_calls',
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
'choices': [
|
||||
{
|
||||
'message': {
|
||||
'role': 'assistant',
|
||||
'content': 'Replanned with the guidance.',
|
||||
},
|
||||
'finish_reason': 'stop',
|
||||
}
|
||||
]
|
||||
},
|
||||
]
|
||||
recorded_payloads: list[dict[str, object]] = []
|
||||
provider_calls = 0
|
||||
|
||||
def guidance_provider() -> tuple[dict[str, object], ...]:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
if provider_calls == 2:
|
||||
return ({'id': 7, 'content': 'Use new.txt instead.'},)
|
||||
return ()
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
workspace = Path(tmp_dir)
|
||||
(workspace / 'old.txt').write_text('old content\n', encoding='utf-8')
|
||||
with patch(
|
||||
'src.openai_compat.request.urlopen',
|
||||
side_effect=make_recording_urlopen_side_effect(responses, recorded_payloads),
|
||||
):
|
||||
agent = LocalCodingAgent(
|
||||
model_config=ModelConfig(
|
||||
model='Qwen/Qwen3-Coder-30B-A3B-Instruct',
|
||||
base_url='http://127.0.0.1:8000/v1',
|
||||
),
|
||||
runtime_config=AgentRuntimeConfig(cwd=workspace),
|
||||
)
|
||||
agent.runtime_guidance_provider = guidance_provider
|
||||
result = agent.run('Inspect old.txt')
|
||||
|
||||
self.assertEqual(result.final_output, 'Replanned with the guidance.')
|
||||
self.assertEqual(result.tool_calls, 0)
|
||||
self.assertFalse(result.file_history)
|
||||
replan_events = [
|
||||
event
|
||||
for event in result.events
|
||||
if event.get('type') == 'runtime_guidance_replan_before_tools'
|
||||
]
|
||||
self.assertEqual(len(replan_events), 1)
|
||||
injected_events = [
|
||||
event
|
||||
for event in result.events
|
||||
if event.get('type') == 'runtime_guidance_injected'
|
||||
]
|
||||
self.assertEqual(injected_events[-1].get('stage'), 'before_tools')
|
||||
self.assertEqual(len(recorded_payloads), 2)
|
||||
second_messages = recorded_payloads[1]['messages']
|
||||
assert isinstance(second_messages, list)
|
||||
roles = [
|
||||
message.get('role')
|
||||
for message in second_messages
|
||||
if isinstance(message, dict)
|
||||
]
|
||||
self.assertEqual(roles[-3:], ['assistant', 'tool', 'user'])
|
||||
self.assertIn('runtime_guidance_replan', str(second_messages[-2]))
|
||||
self.assertIn('Use new.txt instead.', str(second_messages[-1]))
|
||||
|
||||
def test_runtime_guidance_interrupts_running_streaming_tool(self) -> None:
|
||||
responses = [
|
||||
{
|
||||
'choices': [
|
||||
{
|
||||
'message': {
|
||||
'role': 'assistant',
|
||||
'content': 'I will run the long tool.',
|
||||
'tool_calls': [
|
||||
{
|
||||
'id': 'call_1',
|
||||
'type': 'function',
|
||||
'function': {
|
||||
'name': 'long_tool',
|
||||
'arguments': '{}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
'finish_reason': 'tool_calls',
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
'choices': [
|
||||
{
|
||||
'message': {
|
||||
'role': 'assistant',
|
||||
'content': 'Replanned after the interruption.',
|
||||
},
|
||||
'finish_reason': 'stop',
|
||||
}
|
||||
]
|
||||
},
|
||||
]
|
||||
recorded_payloads: list[dict[str, object]] = []
|
||||
provider_calls = 0
|
||||
cancel_observed_after_guidance: list[bool] = []
|
||||
|
||||
def guidance_provider() -> tuple[dict[str, object], ...]:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
if provider_calls == 3:
|
||||
return ({'id': 12, 'content': '停一下,改成只生成 3 条。'},)
|
||||
return ()
|
||||
|
||||
def fake_execute_tool_streaming(registry, name, arguments, context): # noqa: ANN001
|
||||
self.assertEqual(name, 'long_tool')
|
||||
yield ToolStreamUpdate(kind='delta', content='working\n', stream='stdout')
|
||||
cancel_observed_after_guidance.append(bool(context.cancel_event.is_set()))
|
||||
yield ToolStreamUpdate(
|
||||
kind='result',
|
||||
result=ToolExecutionResult(
|
||||
name=name,
|
||||
ok=False,
|
||||
content='Tool saw cancellation.',
|
||||
metadata={'cancel_seen': bool(context.cancel_event.is_set())},
|
||||
),
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
workspace = Path(tmp_dir)
|
||||
with (
|
||||
patch(
|
||||
'src.openai_compat.request.urlopen',
|
||||
side_effect=make_recording_urlopen_side_effect(
|
||||
responses,
|
||||
recorded_payloads,
|
||||
),
|
||||
),
|
||||
patch(
|
||||
'src.agent_runtime.execute_tool_streaming',
|
||||
side_effect=fake_execute_tool_streaming,
|
||||
),
|
||||
):
|
||||
agent = LocalCodingAgent(
|
||||
model_config=ModelConfig(
|
||||
model='Qwen/Qwen3-Coder-30B-A3B-Instruct',
|
||||
base_url='http://127.0.0.1:8000/v1',
|
||||
),
|
||||
runtime_config=AgentRuntimeConfig(cwd=workspace),
|
||||
)
|
||||
agent.runtime_guidance_provider = guidance_provider
|
||||
result = agent.run('Run a long tool')
|
||||
|
||||
self.assertEqual(result.final_output, 'Replanned after the interruption.')
|
||||
self.assertEqual(result.tool_calls, 1)
|
||||
self.assertEqual(cancel_observed_after_guidance, [True])
|
||||
interrupt_events = [
|
||||
event
|
||||
for event in result.events
|
||||
if event.get('type') == 'runtime_guidance_interrupt_tool'
|
||||
]
|
||||
self.assertEqual(len(interrupt_events), 1)
|
||||
injected_events = [
|
||||
event
|
||||
for event in result.events
|
||||
if event.get('type') == 'runtime_guidance_injected'
|
||||
]
|
||||
self.assertEqual(injected_events[-1].get('stage'), 'during_tool_interrupted')
|
||||
self.assertEqual(len(recorded_payloads), 2)
|
||||
second_messages = recorded_payloads[1]['messages']
|
||||
assert isinstance(second_messages, list)
|
||||
roles = [
|
||||
message.get('role')
|
||||
for message in second_messages
|
||||
if isinstance(message, dict)
|
||||
]
|
||||
self.assertEqual(roles[-3:], ['assistant', 'tool', 'user'])
|
||||
self.assertIn('runtime_guidance_interrupted_tool', str(second_messages[-2]))
|
||||
self.assertIn('改成只生成 3 条', str(second_messages[-1]))
|
||||
|
||||
def test_runtime_guidance_during_tool_can_defer_until_tool_result(self) -> None:
|
||||
responses = [
|
||||
{
|
||||
'choices': [
|
||||
{
|
||||
'message': {
|
||||
'role': 'assistant',
|
||||
'content': 'I will run the long tool.',
|
||||
'tool_calls': [
|
||||
{
|
||||
'id': 'call_1',
|
||||
'type': 'function',
|
||||
'function': {
|
||||
'name': 'long_tool',
|
||||
'arguments': '{}',
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
'finish_reason': 'tool_calls',
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
'choices': [
|
||||
{
|
||||
'message': {
|
||||
'role': 'assistant',
|
||||
'content': 'Replanned after the tool result.',
|
||||
},
|
||||
'finish_reason': 'stop',
|
||||
}
|
||||
]
|
||||
},
|
||||
]
|
||||
recorded_payloads: list[dict[str, object]] = []
|
||||
provider_calls = 0
|
||||
cancel_observed_after_guidance: list[bool] = []
|
||||
|
||||
def guidance_provider() -> tuple[dict[str, object], ...]:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
if provider_calls == 3:
|
||||
return ({'id': 13, 'content': '完成后顺便整理成 Markdown。'},)
|
||||
return ()
|
||||
|
||||
def fake_execute_tool_streaming(registry, name, arguments, context): # noqa: ANN001
|
||||
self.assertEqual(name, 'long_tool')
|
||||
yield ToolStreamUpdate(kind='delta', content='working\n', stream='stdout')
|
||||
cancel_observed_after_guidance.append(bool(context.cancel_event.is_set()))
|
||||
yield ToolStreamUpdate(
|
||||
kind='result',
|
||||
result=ToolExecutionResult(
|
||||
name=name,
|
||||
ok=True,
|
||||
content='Tool completed.',
|
||||
metadata={'cancel_seen': bool(context.cancel_event.is_set())},
|
||||
),
|
||||
)
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
workspace = Path(tmp_dir)
|
||||
with (
|
||||
patch(
|
||||
'src.openai_compat.request.urlopen',
|
||||
side_effect=make_recording_urlopen_side_effect(
|
||||
responses,
|
||||
recorded_payloads,
|
||||
),
|
||||
),
|
||||
patch(
|
||||
'src.agent_runtime.execute_tool_streaming',
|
||||
side_effect=fake_execute_tool_streaming,
|
||||
),
|
||||
):
|
||||
agent = LocalCodingAgent(
|
||||
model_config=ModelConfig(
|
||||
model='Qwen/Qwen3-Coder-30B-A3B-Instruct',
|
||||
base_url='http://127.0.0.1:8000/v1',
|
||||
),
|
||||
runtime_config=AgentRuntimeConfig(cwd=workspace),
|
||||
)
|
||||
agent.runtime_guidance_provider = guidance_provider
|
||||
result = agent.run('Run a long tool')
|
||||
|
||||
self.assertEqual(result.final_output, 'Replanned after the tool result.')
|
||||
self.assertEqual(cancel_observed_after_guidance, [False])
|
||||
deferred_events = [
|
||||
event
|
||||
for event in result.events
|
||||
if event.get('type') == 'runtime_guidance_deferred_after_tool'
|
||||
]
|
||||
self.assertEqual(len(deferred_events), 1)
|
||||
interrupt_events = [
|
||||
event
|
||||
for event in result.events
|
||||
if event.get('type') == 'runtime_guidance_interrupt_tool'
|
||||
]
|
||||
self.assertFalse(interrupt_events)
|
||||
injected_events = [
|
||||
event
|
||||
for event in result.events
|
||||
if event.get('type') == 'runtime_guidance_injected'
|
||||
]
|
||||
self.assertEqual(injected_events[-1].get('stage'), 'after_tool')
|
||||
second_messages = recorded_payloads[1]['messages']
|
||||
assert isinstance(second_messages, list)
|
||||
self.assertIn('runtime_guidance_deferred_after_tool', str(second_messages[-2]))
|
||||
self.assertIn('完成后顺便整理成 Markdown', str(second_messages[-1]))
|
||||
|
||||
def test_agent_persists_max_turns_explanation_after_tool_result(self) -> None:
|
||||
responses = [
|
||||
{
|
||||
@@ -331,6 +648,72 @@ class AgentRuntimeTests(unittest.TestCase):
|
||||
self.assertEqual(result.transcript[-1]['stop_reason'], 'max_turns')
|
||||
self.assertIn('最大步骤上限', result.transcript[-1]['content'])
|
||||
|
||||
def test_runtime_guidance_before_finish_continues_instead_of_finishing(self) -> None:
|
||||
responses = [
|
||||
{
|
||||
'choices': [
|
||||
{
|
||||
'message': {
|
||||
'role': 'assistant',
|
||||
'content': 'Draft answer.',
|
||||
},
|
||||
'finish_reason': 'stop',
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
'choices': [
|
||||
{
|
||||
'message': {
|
||||
'role': 'assistant',
|
||||
'content': 'Revised final answer.',
|
||||
},
|
||||
'finish_reason': 'stop',
|
||||
}
|
||||
]
|
||||
},
|
||||
]
|
||||
recorded_payloads: list[dict[str, object]] = []
|
||||
provider_calls = 0
|
||||
|
||||
def guidance_provider() -> tuple[dict[str, object], ...]:
|
||||
nonlocal provider_calls
|
||||
provider_calls += 1
|
||||
if provider_calls == 2:
|
||||
return ({'id': 9, 'content': 'Add the timing summary before finalizing.'},)
|
||||
return ()
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
workspace = Path(tmp_dir)
|
||||
with patch(
|
||||
'src.openai_compat.request.urlopen',
|
||||
side_effect=make_recording_urlopen_side_effect(responses, recorded_payloads),
|
||||
):
|
||||
agent = LocalCodingAgent(
|
||||
model_config=ModelConfig(
|
||||
model='Qwen/Qwen3-Coder-30B-A3B-Instruct',
|
||||
base_url='http://127.0.0.1:8000/v1',
|
||||
),
|
||||
runtime_config=AgentRuntimeConfig(cwd=workspace),
|
||||
)
|
||||
agent.runtime_guidance_provider = guidance_provider
|
||||
result = agent.run('Prepare an answer')
|
||||
|
||||
self.assertEqual(result.final_output, 'Revised final answer.')
|
||||
injected_events = [
|
||||
event
|
||||
for event in result.events
|
||||
if event.get('type') == 'runtime_guidance_injected'
|
||||
]
|
||||
self.assertEqual(injected_events[-1].get('stage'), 'before_finish')
|
||||
self.assertEqual(len(recorded_payloads), 2)
|
||||
second_messages = recorded_payloads[1]['messages']
|
||||
assert isinstance(second_messages, list)
|
||||
self.assertEqual(second_messages[-2].get('role'), 'assistant')
|
||||
self.assertEqual(second_messages[-2].get('content'), 'Draft answer.')
|
||||
self.assertEqual(second_messages[-1].get('role'), 'user')
|
||||
self.assertIn('Add the timing summary', str(second_messages[-1].get('content')))
|
||||
|
||||
def test_agent_stops_after_data_agent_goal_requires_review(self) -> None:
|
||||
responses = [
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user