Support runtime guidance tool interrupts

This commit is contained in:
wuyang6
2026-06-23 15:20:59 +08:00
parent 712bccf814
commit 45792c8fd5
10 changed files with 955 additions and 24 deletions
+384 -1
View File
@@ -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 = [
{