fix: recover Work runs after failed verification
This commit is contained in:
+188
-2
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
@@ -25,7 +26,7 @@ class ScriptedProvider:
|
||||
self.requests = []
|
||||
|
||||
async def complete(self, **kwargs):
|
||||
self.requests.append(kwargs)
|
||||
self.requests.append(copy.deepcopy(kwargs))
|
||||
return self.responses.pop(0)
|
||||
|
||||
|
||||
@@ -203,7 +204,7 @@ async def test_artifact_completion_is_rejected_until_written_and_verified(settin
|
||||
try:
|
||||
answer = await AgentLoop(provider, registry, store, max_tool_output_chars=10_000).run(
|
||||
spec=get_model_spec("work-light"),
|
||||
messages=[{"role": "user", "content": "写一个脚本并生成报告"}],
|
||||
messages=[{"role": "user", "content": "写一个报告"}],
|
||||
identity=UserIdentity("u1", "", "", "user"),
|
||||
raw_user_jwt="jwt",
|
||||
chat_id="c1",
|
||||
@@ -212,6 +213,191 @@ async def test_artifact_completion_is_rejected_until_written_and_verified(settin
|
||||
assert answer == "已完成 report.md"
|
||||
assert [name for name, _ in registry.calls] == ["write_file", "read_file"]
|
||||
assert "completion.rejected" in events
|
||||
assert provider.requests[1]["messages"][-1]["role"] == "user"
|
||||
assert "MUST call" in provider.requests[1]["messages"][-1]["content"]
|
||||
finally:
|
||||
await store.close()
|
||||
|
||||
|
||||
class FailingExecutionRegistry(RecordingRegistry):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.exec_calls = 0
|
||||
|
||||
async def execute(self, name, arguments, context):
|
||||
self.calls.append((name, arguments))
|
||||
if name == "exec":
|
||||
self.exec_calls += 1
|
||||
if self.exec_calls == 1:
|
||||
return {"ok": False, "error": "SyntaxError: invalid source"}
|
||||
return {"ok": True, "output": f"{name} complete"}
|
||||
|
||||
|
||||
async def test_failed_execution_must_be_repaired_before_completion(settings) -> None:
|
||||
write = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("1", "write_file", {"path": "sort.py", "content": "broken"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
read = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("2", "read_file", {"path": "sort.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
failed_exec = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("3", "exec", {"command": "python3 sort.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
premature = {"choices": [{"message": {"role": "assistant", "content": "已完成"}, "finish_reason": "stop"}]}
|
||||
repaired_exec = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("4", "exec", {"command": "python3 sort.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
final = {
|
||||
"choices": [{"message": {"role": "assistant", "content": "脚本运行和报告验证完成"}, "finish_reason": "stop"}]
|
||||
}
|
||||
provider = ScriptedProvider([write, read, failed_exec, premature, repaired_exec, final])
|
||||
registry = FailingExecutionRegistry()
|
||||
store = RuntimeStore(settings.database_url)
|
||||
await store.initialize()
|
||||
events = []
|
||||
|
||||
async def callback(event_type, payload):
|
||||
events.append(event_type)
|
||||
|
||||
try:
|
||||
answer = await AgentLoop(provider, registry, store, max_tool_output_chars=10_000).run(
|
||||
spec=get_model_spec("work-light"),
|
||||
messages=[{"role": "user", "content": "写脚本对比排序算法并生成报告"}],
|
||||
identity=UserIdentity("u1", "", "", "user"),
|
||||
raw_user_jwt="jwt",
|
||||
chat_id="c1",
|
||||
callback=callback,
|
||||
)
|
||||
assert answer == "脚本运行和报告验证完成"
|
||||
assert registry.exec_calls == 2
|
||||
assert events.count("completion.rejected") == 1
|
||||
recovery = provider.requests[4]["messages"][-1]
|
||||
assert recovery["role"] == "user"
|
||||
assert "SyntaxError" in recovery["content"]
|
||||
finally:
|
||||
await store.close()
|
||||
|
||||
|
||||
async def test_script_delivery_requires_successful_execution(settings) -> None:
|
||||
write = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("1", "write_file", {"path": "script.py", "content": "print('ok')"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
read = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("2", "read_file", {"path": "script.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
premature = {"choices": [{"message": {"role": "assistant", "content": "脚本完成"}, "finish_reason": "stop"}]}
|
||||
execute = {
|
||||
"choices": [
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [tool_call("3", "exec", {"command": "python3 script.py"})],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
]
|
||||
}
|
||||
final = {"choices": [{"message": {"role": "assistant", "content": "脚本执行验证完成"}, "finish_reason": "stop"}]}
|
||||
provider = ScriptedProvider([write, read, premature, execute, final])
|
||||
registry = RecordingRegistry()
|
||||
store = RuntimeStore(settings.database_url)
|
||||
await store.initialize()
|
||||
|
||||
async def callback(event_type, payload):
|
||||
return None
|
||||
|
||||
try:
|
||||
answer = await AgentLoop(provider, registry, store, max_tool_output_chars=10_000).run(
|
||||
spec=get_model_spec("work-light"),
|
||||
messages=[{"role": "user", "content": "写一个 Python 脚本"}],
|
||||
identity=UserIdentity("u1", "", "", "user"),
|
||||
raw_user_jwt="jwt",
|
||||
chat_id="c1",
|
||||
callback=callback,
|
||||
)
|
||||
assert answer == "脚本执行验证完成"
|
||||
assert [name for name, _ in registry.calls] == ["write_file", "read_file", "exec"]
|
||||
assert "no command or process completed successfully" in provider.requests[3]["messages"][-1]["content"]
|
||||
finally:
|
||||
await store.close()
|
||||
|
||||
|
||||
async def test_rejected_checkpoint_does_not_spin_without_tool_calls(settings) -> None:
|
||||
stopped = {"choices": [{"message": {"role": "assistant", "content": "已完成"}, "finish_reason": "stop"}]}
|
||||
provider = ScriptedProvider([stopped, stopped, stopped])
|
||||
store = RuntimeStore(settings.database_url)
|
||||
await store.initialize()
|
||||
events = []
|
||||
|
||||
async def callback(event_type, payload):
|
||||
events.append(event_type)
|
||||
|
||||
try:
|
||||
answer = await AgentLoop(
|
||||
provider,
|
||||
RecordingRegistry(),
|
||||
store,
|
||||
max_tool_output_chars=10_000,
|
||||
).run(
|
||||
spec=get_model_spec("work-light"),
|
||||
messages=[{"role": "user", "content": "写一个报告"}],
|
||||
identity=UserIdentity("u1", "", "", "user"),
|
||||
raw_user_jwt="jwt",
|
||||
chat_id="c1",
|
||||
callback=callback,
|
||||
)
|
||||
assert answer.startswith("未完成:")
|
||||
assert len(provider.requests) == 3
|
||||
assert events.count("completion.rejected") == 2
|
||||
assert events.count("completion.unverified") == 1
|
||||
finally:
|
||||
await store.close()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user