584 lines
20 KiB
Python
584 lines
20 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import copy
|
|
import json
|
|
from typing import Any
|
|
|
|
from agent_platform.auth import UserIdentity
|
|
from agent_platform.models import get_model_spec
|
|
from agent_platform.runtime.loop import (
|
|
AgentLoop,
|
|
is_substantive_execution,
|
|
normalize_write_file_content,
|
|
tool_event_details,
|
|
)
|
|
from agent_platform.runtime.tools import TOOL_METADATA
|
|
from agent_platform.store import RuntimeStore
|
|
|
|
|
|
def tool_call(call_id: str, name: str, arguments: dict[str, Any]) -> dict[str, Any]:
|
|
return {
|
|
"id": call_id,
|
|
"type": "function",
|
|
"function": {"name": name, "arguments": json.dumps(arguments)},
|
|
}
|
|
|
|
|
|
def test_write_file_normalizes_double_escaped_code_layout() -> None:
|
|
escaped = r'import time\n\n# benchmark\ndef main():\n\tprint("keep\ninside")\n\nmain()\n'
|
|
assert normalize_write_file_content("benchmark.py", escaped) == (
|
|
'import time\n\n# benchmark\ndef main():\n\tprint("keep\\ninside")\n\nmain()\n'
|
|
)
|
|
|
|
|
|
def test_write_file_normalizes_double_escaped_plain_text() -> None:
|
|
assert normalize_write_file_content("report.md", r"# Report\n\nMeasured: 1.2 s\n") == (
|
|
"# Report\n\nMeasured: 1.2 s\n"
|
|
)
|
|
|
|
|
|
def test_read_only_shell_commands_do_not_satisfy_execution_evidence() -> None:
|
|
assert not is_substantive_execution("exec", {"command": "cat report.txt"})
|
|
assert not is_substantive_execution("exec", {"command": "sed -n '1,20p' script.py"})
|
|
assert is_substantive_execution("exec", {"command": "python3 script.py"})
|
|
assert is_substantive_execution("start_process", {"command": "python3 server.py"})
|
|
|
|
|
|
class ScriptedProvider:
|
|
def __init__(self, responses: list[dict[str, Any]]) -> None:
|
|
self.responses = responses
|
|
self.requests = []
|
|
|
|
async def complete(self, **kwargs):
|
|
self.requests.append(copy.deepcopy(kwargs))
|
|
return self.responses.pop(0)
|
|
|
|
|
|
class ParallelRegistry:
|
|
def __init__(self) -> None:
|
|
self.started = 0
|
|
self.both_started = asyncio.Event()
|
|
|
|
def specs(self, *, read_only=False, allow_delegate=True):
|
|
return [TOOL_METADATA["read_file"].openai_spec()]
|
|
|
|
async def execute(self, name, arguments, context):
|
|
self.started += 1
|
|
if self.started == 2:
|
|
self.both_started.set()
|
|
await asyncio.wait_for(self.both_started.wait(), timeout=1)
|
|
return {"ok": True, "path": arguments["path"]}
|
|
|
|
|
|
async def test_read_only_tool_calls_run_in_parallel(settings) -> None:
|
|
first = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
tool_call("1", "read_file", {"path": "a"}),
|
|
tool_call("2", "read_file", {"path": "b"}),
|
|
],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
final = {"choices": [{"message": {"role": "assistant", "content": "done"}, "finish_reason": "stop"}]}
|
|
store = RuntimeStore(settings.database_url)
|
|
await store.initialize()
|
|
registry = ParallelRegistry()
|
|
loop = AgentLoop(ScriptedProvider([first, final]), registry, store, max_tool_output_chars=10_000)
|
|
events = []
|
|
|
|
async def callback(event_type, payload):
|
|
events.append(event_type)
|
|
|
|
try:
|
|
answer = await loop.run(
|
|
spec=get_model_spec("work-light"),
|
|
messages=[{"role": "user", "content": "inspect both"}],
|
|
identity=UserIdentity("u1", "", "", "user"),
|
|
raw_user_jwt="jwt",
|
|
chat_id="c1",
|
|
callback=callback,
|
|
)
|
|
assert answer == "done"
|
|
assert registry.started == 2
|
|
assert events.count("tool.completed") == 2
|
|
finally:
|
|
await store.close()
|
|
|
|
|
|
class SerialRegistry:
|
|
def __init__(self) -> None:
|
|
self.active = 0
|
|
self.max_active = 0
|
|
|
|
def specs(self, *, read_only=False, allow_delegate=True):
|
|
return [TOOL_METADATA["write_file"].openai_spec()]
|
|
|
|
async def execute(self, name, arguments, context):
|
|
self.active += 1
|
|
self.max_active = max(self.max_active, self.active)
|
|
await asyncio.sleep(0.02)
|
|
self.active -= 1
|
|
return {"ok": True}
|
|
|
|
|
|
async def test_mutating_tool_calls_are_serialized(settings) -> None:
|
|
first = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
tool_call("1", "write_file", {"path": "a", "content": "a"}),
|
|
tool_call("2", "write_file", {"path": "b", "content": "b"}),
|
|
],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
verify = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [tool_call("3", "read_file", {"path": "a"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
final = {"choices": [{"message": {"role": "assistant", "content": "done"}, "finish_reason": "stop"}]}
|
|
store = RuntimeStore(settings.database_url)
|
|
await store.initialize()
|
|
registry = SerialRegistry()
|
|
loop = AgentLoop(ScriptedProvider([first, verify, final]), registry, store, max_tool_output_chars=10_000)
|
|
|
|
async def callback(event_type, payload):
|
|
return None
|
|
|
|
try:
|
|
await loop.run(
|
|
spec=get_model_spec("work-light"),
|
|
messages=[{"role": "user", "content": "write both"}],
|
|
identity=UserIdentity("u1", "", "", "user"),
|
|
raw_user_jwt="jwt",
|
|
chat_id="c1",
|
|
callback=callback,
|
|
)
|
|
assert registry.max_active == 1
|
|
finally:
|
|
await store.close()
|
|
|
|
|
|
class RecordingRegistry:
|
|
def __init__(self) -> None:
|
|
self.calls = []
|
|
|
|
def specs(self, *, read_only=False, allow_delegate=True):
|
|
return [metadata.openai_spec() for metadata in TOOL_METADATA.values()]
|
|
|
|
async def execute(self, name, arguments, context):
|
|
self.calls.append((name, arguments))
|
|
return {"ok": True, "output": f"{name} complete"}
|
|
|
|
|
|
async def test_artifact_completion_is_rejected_until_written_and_verified(settings) -> None:
|
|
premature = {"choices": [{"message": {"role": "assistant", "content": "报告已经完成"}, "finish_reason": "stop"}]}
|
|
write = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("1", "write_file", {"path": "report.md", "content": "ok"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
verify = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("2", "read_file", {"path": "report.md"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
final = {"choices": [{"message": {"role": "assistant", "content": "已完成 report.md"}, "finish_reason": "stop"}]}
|
|
provider = ScriptedProvider([premature, write, verify, final])
|
|
registry = RecordingRegistry()
|
|
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 == "已完成 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"}]}
|
|
repair = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("4", "write_file", {"path": "sort.py", "content": "fixed"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
repaired_exec = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("5", "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, repair, 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
|
|
direct_recovery = provider.requests[3]["messages"][-1]
|
|
assert direct_recovery["role"] == "user"
|
|
assert "Execution recovery is required" in direct_recovery["content"]
|
|
recovery = provider.requests[4]["messages"][-1]
|
|
assert recovery["role"] == "user"
|
|
assert "SyntaxError" in recovery["content"]
|
|
finally:
|
|
await store.close()
|
|
|
|
|
|
async def test_unchanged_failed_command_is_blocked_until_a_repair(settings) -> None:
|
|
write = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("1", "write_file", {"path": "sort.py", "content": "broken"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
failed_exec = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("2", "exec", {"command": "python3 sort.py"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
repeated_exec = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("3", "exec", {"command": "python3 sort.py"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
repair = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("4", "write_file", {"path": "sort.py", "content": "fixed"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
successful_exec = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("5", "exec", {"command": "python3 sort.py"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
final = {"choices": [{"message": {"role": "assistant", "content": "验证完成"}, "finish_reason": "stop"}]}
|
|
provider = ScriptedProvider([write, failed_exec, repeated_exec, repair, successful_exec, final])
|
|
registry = FailingExecutionRegistry()
|
|
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": "写脚本并运行"}],
|
|
identity=UserIdentity("u1", "", "", "user"),
|
|
raw_user_jwt="jwt",
|
|
chat_id="c1",
|
|
callback=callback,
|
|
)
|
|
assert answer == "验证完成"
|
|
assert registry.exec_calls == 2
|
|
assert [name for name, _ in registry.calls] == ["write_file", "exec", "write_file", "exec"]
|
|
blocked_result = provider.requests[3]["messages"][-2]
|
|
assert blocked_result["role"] == "tool"
|
|
assert "unchanged retry" in blocked_result["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()
|
|
|
|
|
|
async def test_bare_interactive_exec_is_blocked_before_registry(settings) -> None:
|
|
first = {
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"tool_calls": [tool_call("1", "exec", {"command": "python3"})],
|
|
},
|
|
"finish_reason": "tool_calls",
|
|
}
|
|
]
|
|
}
|
|
final = {"choices": [{"message": {"role": "assistant", "content": "无法执行"}, "finish_reason": "stop"}]}
|
|
registry = RecordingRegistry()
|
|
store = RuntimeStore(settings.database_url)
|
|
await store.initialize()
|
|
|
|
async def callback(event_type, payload):
|
|
return None
|
|
|
|
try:
|
|
answer = await AgentLoop(
|
|
ScriptedProvider([first, final]),
|
|
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.calls == []
|
|
finally:
|
|
await store.close()
|
|
|
|
|
|
def test_tool_events_render_one_friendly_completed_card() -> None:
|
|
assert tool_event_details("tool.started", {"name": "exec"}) is None
|
|
rendered = tool_event_details(
|
|
"tool.completed",
|
|
{
|
|
"call_id": "call-1",
|
|
"name": "exec",
|
|
"arguments": {"command": "pytest -q"},
|
|
"ok": True,
|
|
"summary": "2 passed",
|
|
},
|
|
)
|
|
assert rendered is not None
|
|
assert 'name="执行命令"' in rendered
|
|
assert "2 passed" in rendered
|
|
assert 'done="false"' not in rendered
|