Files
zk-data-agent/tests/test_agent_loop.py
T

275 lines
8.7 KiB
Python

from __future__ import annotations
import asyncio
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, 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)},
}
class ScriptedProvider:
def __init__(self, responses: list[dict[str, Any]]) -> None:
self.responses = responses
self.requests = []
async def complete(self, **kwargs):
self.requests.append(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
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