677 lines
28 KiB
Python
677 lines
28 KiB
Python
#!/usr/bin/env python3
|
||
"""dispatch-agent-task 的离线测试(不连库、不连网、不启动真 pi)。
|
||
|
||
用假 pi 事件流(与 pi --mode json 真实线格式一致)与假数据库连接固定:
|
||
任务包合同、argv 构造、事件归一、run_dispatch 全链失败关闭与回执形状。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import pathlib
|
||
import stat
|
||
import sys
|
||
import tempfile
|
||
import unittest
|
||
|
||
|
||
PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[3]
|
||
SKILL_DIR = PROJECT_ROOT / ".agent" / "skills" / "dispatch-agent-task" / "scripts"
|
||
EVIDENCE_DIR = PROJECT_ROOT / ".agent" / "skills" / "record-run-evidence" / "scripts"
|
||
for path in (SKILL_DIR, EVIDENCE_DIR):
|
||
if str(path) not in sys.path:
|
||
sys.path.insert(0, str(path))
|
||
|
||
import agent_task # noqa: E402
|
||
from agent_task import ( # noqa: E402
|
||
OutputInvalidError,
|
||
TaskSpecError,
|
||
build_task_package,
|
||
load_spec,
|
||
validate_structured_output,
|
||
)
|
||
from pi_runner import ( # noqa: E402
|
||
ExecutionPolicy,
|
||
FrameworkError,
|
||
PiAgentRunner,
|
||
build_pi_argv,
|
||
)
|
||
from dispatch_agent_task import EXIT_OUTPUT_INVALID, EXIT_SPEC_INVALID, run_dispatch # noqa: E402
|
||
from agent_trace import AgentTraceWriter # noqa: E402
|
||
|
||
REPO_ROOT = PROJECT_ROOT
|
||
|
||
OUTPUT_SCHEMA = {
|
||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||
"type": "object",
|
||
"properties": {
|
||
"title": {"type": "string", "minLength": 1},
|
||
"beats": {"type": "array", "minItems": 1, "items": {"type": "string"}},
|
||
},
|
||
"required": ["title", "beats"],
|
||
"additionalProperties": False,
|
||
}
|
||
|
||
|
||
def make_spec(tmp: pathlib.Path, **overrides) -> pathlib.Path:
|
||
spec = {
|
||
"specVersion": "agent-task-v1",
|
||
"role": "planner",
|
||
"taskPrompt": "为一幕场景给出三拍结构",
|
||
"input": {"premise": "深空站失联前的最后八小时"},
|
||
"outputSchema": OUTPUT_SCHEMA,
|
||
"outputSchemaId": "planner-mini-outline-v1",
|
||
"toolAllowlist": [],
|
||
"maxDurationSeconds": 120,
|
||
}
|
||
spec.update(overrides)
|
||
path = tmp / "task.json"
|
||
path.write_text(json.dumps(spec, ensure_ascii=False), encoding="utf-8")
|
||
return path
|
||
|
||
|
||
def pi_stream_lines(final_text: str, *, with_tool: bool = False, model: str = "m-a"):
|
||
"""构造与 pi --mode json 一致的假事件流(含 usage/model/stopReason)。"""
|
||
assistant_final = {
|
||
"role": "assistant",
|
||
"content": [{"type": "text", "text": final_text}],
|
||
"model": model,
|
||
"provider": "prov",
|
||
"usage": {"input": 100, "output": 40, "cacheRead": 10, "cost": {"total": 0.012}},
|
||
"stopReason": "stop",
|
||
}
|
||
lines = [
|
||
{"type": "session", "version": 3, "id": "sess-1", "cwd": "/tmp"},
|
||
{"type": "agent_start"},
|
||
{"type": "turn_start"},
|
||
{"type": "message_end", "message": {"role": "user", "content": []}},
|
||
]
|
||
if with_tool:
|
||
assistant_tool = {
|
||
"role": "assistant",
|
||
"content": [{"type": "tool_call", "id": "t1", "name": "read", "arguments": {}}],
|
||
"model": model,
|
||
"provider": "prov",
|
||
"usage": {"input": 90, "output": 5, "cost": {"total": 0.001}},
|
||
"stopReason": "tool_use",
|
||
}
|
||
lines += [
|
||
{"type": "message_end", "message": assistant_tool},
|
||
{"type": "tool_execution_start", "toolCallId": "t1", "toolName": "read", "args": {}},
|
||
{"type": "tool_execution_end", "toolCallId": "t1", "toolName": "read", "result": "ok", "isError": False},
|
||
{"type": "turn_start"},
|
||
]
|
||
lines += [
|
||
{"type": "message_end", "message": assistant_final},
|
||
{"type": "turn_end", "message": assistant_final, "toolResults": []},
|
||
{"type": "agent_end", "messages": [{"role": "user", "content": []}, assistant_final]},
|
||
{"type": "agent_settled"},
|
||
]
|
||
return [json.dumps(line, ensure_ascii=False).encode("utf-8") + b"\n" for line in lines]
|
||
|
||
|
||
class FakeStream:
|
||
def __init__(self, lines, exit_code=0, timed_out=False):
|
||
self._lines = lines
|
||
self.exit_code = exit_code
|
||
self.timed_out = timed_out
|
||
|
||
def __iter__(self):
|
||
yield from self._lines
|
||
|
||
def close(self):
|
||
return None
|
||
|
||
|
||
def fake_launcher(lines, exit_code=0, timed_out=False):
|
||
def _launch(argv, timeout, cwd):
|
||
assert argv[0] == "pi", argv
|
||
return FakeStream(list(lines), exit_code=exit_code, timed_out=timed_out)
|
||
|
||
return _launch
|
||
|
||
|
||
class RecordingConnect:
|
||
"""捕获全部写入语句的假连接工厂(事件 + 证据 + run 登记)。"""
|
||
|
||
def __init__(self):
|
||
self.log = []
|
||
|
||
def __call__(self, *args, **kwargs):
|
||
outer = self
|
||
|
||
class _Cursor:
|
||
def execute(self, sql, params=None):
|
||
outer.log.append((sql, params))
|
||
return self
|
||
|
||
def fetchone(self):
|
||
sql = outer.log[-1][0]
|
||
if sql.startswith("SELECT run_id, work_id"):
|
||
# start_run 回读:返回与入参一致的绑定行。
|
||
params = outer.log[-1][1]
|
||
return ("row", None, None, "running")
|
||
if sql.startswith("INSERT INTO example_run"):
|
||
return ("row",)
|
||
if sql.startswith("UPDATE example_run"):
|
||
return ("row", "failed" if "failed" in (params := outer.log[-1][1]) else "completed", None)
|
||
RecordingConnect.next_id += 1
|
||
return (RecordingConnect.next_id,)
|
||
|
||
def fetchall(self):
|
||
return []
|
||
|
||
def commit(self):
|
||
outer.log.append(("COMMIT", None))
|
||
|
||
def rollback(self):
|
||
outer.log.append(("ROLLBACK", None))
|
||
|
||
class _Ctx:
|
||
def __enter__(self):
|
||
return _Cursor()
|
||
|
||
def __exit__(self, *exc):
|
||
return False
|
||
|
||
return _Ctx()
|
||
|
||
next_id = 5000
|
||
|
||
def event_rows(self):
|
||
return [p for sql, p in self.log if sql.startswith("INSERT INTO example_agent_event")]
|
||
|
||
def run_rows(self):
|
||
return [p for sql, p in self.log if sql.startswith("INSERT INTO example_run")]
|
||
|
||
|
||
class TaskSpecTest(unittest.TestCase):
|
||
def setUp(self):
|
||
self.tmp = pathlib.Path(tempfile.mkdtemp())
|
||
|
||
def test_load_and_hash_binding(self):
|
||
path = make_spec(self.tmp)
|
||
spec = load_spec(path)
|
||
self.assertEqual(spec.role, "planner")
|
||
# 写错 input 哈希必须失败关闭。
|
||
bad = json.loads(path.read_text())
|
||
bad["inputSha256"] = "sha256:" + "0" * 64
|
||
path2 = self.tmp / "bad.json"
|
||
path2.write_text(json.dumps(bad), encoding="utf-8")
|
||
with self.assertRaises(TaskSpecError):
|
||
load_spec(path2)
|
||
|
||
def test_accepts_role_file_names_and_rejects_unknown_role(self):
|
||
for role in ("writer", "planner", "detector", "judge", "extractor"):
|
||
with self.subTest(role=role):
|
||
spec = load_spec(make_spec(self.tmp, role=role))
|
||
self.assertEqual(build_task_package(spec, REPO_ROOT).spec.role, role)
|
||
with self.assertRaises(TaskSpecError):
|
||
load_spec(make_spec(self.tmp, role="semantic_detector"))
|
||
with self.assertRaises(TaskSpecError):
|
||
load_spec(make_spec(self.tmp, role="hacker"))
|
||
|
||
def test_rejects_bad_schema_and_empty_prompt(self):
|
||
with self.assertRaises(TaskSpecError):
|
||
load_spec(make_spec(self.tmp, outputSchema={"type": "no-such-type"}))
|
||
with self.assertRaises(TaskSpecError):
|
||
load_spec(make_spec(self.tmp, taskPrompt=" "))
|
||
|
||
def test_rejects_nonportable_or_malformed_fields(self):
|
||
invalid_overrides = (
|
||
{"provider": "p"},
|
||
{"toolAllowlist": "read"},
|
||
{"toolAllowlist": ["read tool"]},
|
||
{"toolAllowlist": ["read", "read"]},
|
||
{"maxDurationSeconds": 0},
|
||
{"maxDurationSeconds": float("nan")},
|
||
{"workId": "1"},
|
||
{"targetChapter": True},
|
||
)
|
||
for overrides in invalid_overrides:
|
||
with self.subTest(overrides=overrides), self.assertRaises(TaskSpecError):
|
||
load_spec(make_spec(self.tmp, **overrides))
|
||
malformed = self.tmp / "malformed.json"
|
||
malformed.write_text("{", encoding="utf-8")
|
||
with self.assertRaises(TaskSpecError):
|
||
load_spec(malformed)
|
||
nonstandard = self.tmp / "nonstandard.json"
|
||
nonstandard.write_text(json.dumps({"specVersion": float("nan")}), encoding="utf-8")
|
||
with self.assertRaises(TaskSpecError):
|
||
load_spec(nonstandard)
|
||
|
||
def test_spec_hash_binds_scope_and_deadline(self):
|
||
base = build_task_package(load_spec(make_spec(self.tmp)), REPO_ROOT)
|
||
scoped = build_task_package(
|
||
load_spec(make_spec(self.tmp, workId=7, targetChapter=3)), REPO_ROOT
|
||
)
|
||
slower = build_task_package(
|
||
load_spec(make_spec(self.tmp, maxDurationSeconds=121)), REPO_ROOT
|
||
)
|
||
self.assertNotEqual(base.spec_sha256, scoped.spec_sha256)
|
||
self.assertNotEqual(base.spec_sha256, slower.spec_sha256)
|
||
|
||
|
||
class PackageTest(unittest.TestCase):
|
||
def setUp(self):
|
||
self.tmp = pathlib.Path(tempfile.mkdtemp())
|
||
|
||
def test_system_prompt_is_identity_plus_role_contract_plus_schema(self):
|
||
package = build_task_package(load_spec(make_spec(self.tmp)), REPO_ROOT)
|
||
role_text = (REPO_ROOT / ".agent" / "agents" / "planner.md").read_text(encoding="utf-8")
|
||
self.assertTrue(package.system_prompt.startswith(role_text.rstrip()))
|
||
self.assertIn("--- 角色合同(唯一事实源) ---", package.system_prompt)
|
||
self.assertIn("## planner:规划师", package.system_prompt)
|
||
self.assertIn("--- 结构化输出合同 ---", package.system_prompt)
|
||
self.assertIn('"$schema"', package.system_prompt)
|
||
self.assertIn("--- 冻结输入 ---", package.user_message)
|
||
self.assertIn("深空站失联前的最后八小时", package.user_message)
|
||
identity = package.as_identity()
|
||
self.assertEqual(identity["toolAllowlist"], [])
|
||
self.assertEqual(identity["roleContractVersion"], "role-contracts-v1")
|
||
self.assertEqual(identity["roleContractSource"], ".agent/docs/architecture/角色合同.md")
|
||
self.assertTrue(all(len(v) == 71 and v.startswith("sha256:") for k, v in identity.items() if k.endswith("Sha256")))
|
||
|
||
|
||
class ExecutionPolicyTest(unittest.TestCase):
|
||
def test_provider_and_model_are_explicit(self):
|
||
with self.assertRaisesRegex(ValueError, "provider"):
|
||
ExecutionPolicy(model="m")
|
||
with self.assertRaisesRegex(ValueError, "model"):
|
||
ExecutionPolicy(provider="p")
|
||
self.assertEqual(
|
||
ExecutionPolicy(provider="p", model="m").requested_model_id,
|
||
"p/m",
|
||
)
|
||
|
||
|
||
class ArgvTest(unittest.TestCase):
|
||
def setUp(self):
|
||
self.tmp = pathlib.Path(tempfile.mkdtemp())
|
||
|
||
def test_argv_injects_prompt_tools_and_isolation(self):
|
||
package = build_task_package(load_spec(make_spec(self.tmp)), REPO_ROOT)
|
||
argv = build_pi_argv(package, ExecutionPolicy(provider="p", model="m", thinking="low"))
|
||
joined = " ".join(argv)
|
||
self.assertIn("--system-prompt", argv)
|
||
self.assertEqual(argv[argv.index("--system-prompt") + 1], package.system_prompt)
|
||
self.assertEqual(argv[-1], package.user_message)
|
||
self.assertIn("--no-context-files", joined)
|
||
self.assertIn("--no-skills", joined)
|
||
self.assertIn("--no-extensions", joined)
|
||
self.assertIn("--mode", joined)
|
||
self.assertEqual(argv[argv.index("--provider") + 1], "p")
|
||
self.assertEqual(argv[argv.index("--model") + 1], "m")
|
||
self.assertIn("--no-tools", joined)
|
||
|
||
def test_tool_allowlist_maps_to_tools_flag(self):
|
||
spec_path = make_spec(pathlib.Path(tempfile.mkdtemp()), toolAllowlist=["read", "bash"])
|
||
package = build_task_package(load_spec(spec_path), REPO_ROOT)
|
||
argv = build_pi_argv(package, ExecutionPolicy(provider="p", model="m"))
|
||
self.assertIn("--tools", argv)
|
||
self.assertEqual(argv[argv.index("--tools") + 1], "read,bash")
|
||
|
||
|
||
class RunnerParseTest(unittest.TestCase):
|
||
def setUp(self):
|
||
self.tmp = pathlib.Path(tempfile.mkdtemp())
|
||
|
||
def _sink(self):
|
||
events = []
|
||
|
||
class Sink:
|
||
emit = AgentTraceWriter(run_id="t", framework="pi", agent_role="planner", connect=object())
|
||
|
||
# 直接构造一个记录型 sink,绕开数据库。
|
||
class RecordingSink(AgentTraceWriter):
|
||
def __init__(self):
|
||
super().__init__(
|
||
run_id="t", framework="pi", agent_role="planner",
|
||
connect=lambda: (_ for _ in ()).throw(AssertionError("离线测试不应触库")),
|
||
)
|
||
self.rows = []
|
||
|
||
def emit(self, event_type, **kwargs):
|
||
self.rows.append((event_type, kwargs))
|
||
return len(self.rows)
|
||
|
||
return RecordingSink()
|
||
|
||
def test_stream_parses_model_tools_and_final_text(self):
|
||
package = build_task_package(
|
||
load_spec(make_spec(self.tmp, toolAllowlist=["read"])), REPO_ROOT
|
||
)
|
||
sink = self._sink()
|
||
outcome = PiAgentRunner(launcher=fake_launcher(pi_stream_lines('{"title":"t","beats":["a"]}', with_tool=True))).run(
|
||
package, ExecutionPolicy(provider="p", model="m"), sink, timeout_seconds=10
|
||
)
|
||
self.assertEqual(outcome.final_text, '{"title":"t","beats":["a"]}')
|
||
self.assertEqual(len(outcome.model_calls), 2)
|
||
self.assertEqual(outcome.model_calls[-1].actual_model_id, "prov/m-a")
|
||
self.assertEqual(outcome.model_calls[-1].cost_usd, 0.012)
|
||
self.assertEqual(outcome.tool_calls[0].name, "read")
|
||
self.assertFalse(outcome.tool_calls[0].is_error)
|
||
types = [row[0] for row in sink.rows]
|
||
self.assertEqual(types, ["agent.started", "model.completed", "tool.started", "tool.completed", "model.completed", "agent.completed"])
|
||
|
||
def test_framework_failures_raise_with_stable_codes(self):
|
||
package = build_task_package(load_spec(make_spec(self.tmp)), REPO_ROOT)
|
||
for launcher, code in (
|
||
(fake_launcher(pi_stream_lines("x"), exit_code=1), "FRAMEWORK_EXIT_NONZERO"),
|
||
(fake_launcher(pi_stream_lines("x"), timed_out=True), "FRAMEWORK_TIMEOUT"),
|
||
(fake_launcher([b"not json\n"]), "STREAM_PARSE_ERROR"),
|
||
(fake_launcher([json.dumps({"type": "agent_end", "messages": []}).encode() + b"\n"]), "NO_MODEL_RESPONSE"),
|
||
(fake_launcher(pi_stream_lines("x", with_tool=True)), "TOOL_NOT_ALLOWED"),
|
||
(
|
||
fake_launcher(
|
||
pi_stream_lines("x")[:-4]
|
||
+ [
|
||
json.dumps(
|
||
{
|
||
"type": "message_end",
|
||
"message": {
|
||
"role": "assistant",
|
||
"content": [],
|
||
"model": "m-a",
|
||
"provider": "prov",
|
||
"usage": {},
|
||
"stopReason": "error",
|
||
"errorMessage": "provider failed",
|
||
},
|
||
}
|
||
).encode()
|
||
+ b"\n"
|
||
]
|
||
),
|
||
"MODEL_TURN_FAILED",
|
||
),
|
||
):
|
||
with self.assertRaises(FrameworkError) as ctx:
|
||
PiAgentRunner(launcher=launcher).run(
|
||
package,
|
||
ExecutionPolicy(provider="p", model="m"),
|
||
self._sink(),
|
||
timeout_seconds=5,
|
||
)
|
||
self.assertEqual(ctx.exception.error_code, code)
|
||
|
||
|
||
class RunDispatchTest(unittest.TestCase):
|
||
def setUp(self):
|
||
self.tmp = pathlib.Path(tempfile.mkdtemp())
|
||
|
||
def _dispatch(self, launcher, connect):
|
||
return run_dispatch(
|
||
make_spec(self.tmp),
|
||
repo_root=REPO_ROOT,
|
||
policy=ExecutionPolicy(provider="p", model="claude-opus-test"),
|
||
run_id="unittest-agent-dispatch-1",
|
||
run_dir=self.tmp / "run",
|
||
connect_factory=connect,
|
||
launcher=launcher,
|
||
trigger_source="diagnostic",
|
||
)
|
||
|
||
def test_success_path_events_and_receipt(self):
|
||
connect = RecordingConnect()
|
||
receipt, code = self._dispatch(
|
||
fake_launcher(pi_stream_lines('{"title":"重启","beats":["警报","分歧","决断"]}')), connect
|
||
)
|
||
self.assertEqual(code, 0)
|
||
self.assertEqual(receipt["status"], "completed")
|
||
self.assertEqual(receipt["requestedModelId"], "p/claude-opus-test")
|
||
self.assertEqual(receipt["usage"], {"inputTokens": 110, "outputTokens": 40, "cachedTokens": 10, "reasoningTokens": 0})
|
||
self.assertEqual(receipt["totalCostUsd"], 0.012)
|
||
self.assertTrue(receipt["costComplete"])
|
||
self.assertEqual(
|
||
[p[2] for p in connect.event_rows()],
|
||
["run.started", "agent.started", "model.completed", "agent.completed", "run.completed"],
|
||
)
|
||
self.assertEqual(receipt["evidence"]["status"], "written")
|
||
llm_calls = [p for sql, p in connect.log if sql.startswith("INSERT INTO example_llm_call")]
|
||
self.assertIsNone(llm_calls[0][0]) # window_key 只属于 16 字符额度窗,run_id 走专列。
|
||
self.assertEqual(llm_calls[0][1], "unittest-agent-dispatch-1")
|
||
# 运行目录审计件齐全且只对当前用户开放。
|
||
run_dir = self.tmp / "run"
|
||
self.assertEqual(stat.S_IMODE(run_dir.stat().st_mode), 0o700)
|
||
for name in ("task-spec.json", "system-prompt.txt", "user-message.txt", "transcript.jsonl", "output.json", "receipt.json"):
|
||
path = run_dir / name
|
||
self.assertTrue(path.exists(), name)
|
||
self.assertEqual(stat.S_IMODE(path.stat().st_mode), 0o600, name)
|
||
|
||
def test_schema_violation_fails_closed(self):
|
||
connect = RecordingConnect()
|
||
receipt, code = self._dispatch(
|
||
fake_launcher(pi_stream_lines('{"title":"重启"}')), connect # 缺 beats
|
||
)
|
||
self.assertEqual(code, EXIT_OUTPUT_INVALID)
|
||
self.assertEqual(receipt["errorCode"], "OUTPUT_SCHEMA_INVALID")
|
||
self.assertEqual(receipt["evidence"]["status"], "written")
|
||
self.assertEqual(
|
||
[p[2] for p in connect.event_rows()],
|
||
["run.started", "agent.started", "model.completed", "agent.completed", "run.failed"],
|
||
)
|
||
self.assertTrue(any(sql.startswith("UPDATE example_run") for sql, _ in connect.log))
|
||
|
||
def test_secret_in_transcript_fails_without_leaking_receipt(self):
|
||
connect = RecordingConnect()
|
||
secret = "sk-abcdef0123456789abcdef012345"
|
||
receipt, code = self._dispatch(
|
||
fake_launcher(pi_stream_lines(json.dumps({"title": secret, "beats": ["x"]}))),
|
||
connect,
|
||
)
|
||
self.assertEqual(code, 5)
|
||
self.assertEqual(receipt["errorCode"], "RAW_SECRET_DETECTED")
|
||
self.assertNotIn(secret, json.dumps(receipt, ensure_ascii=False))
|
||
self.assertFalse((self.tmp / "run" / "transcript.jsonl").exists())
|
||
|
||
def test_framework_runs_from_repo_root(self):
|
||
seen = {}
|
||
|
||
def launcher(argv, timeout, cwd):
|
||
seen["cwd"] = cwd
|
||
return FakeStream(pi_stream_lines('{"title":"t","beats":["b"]}'))
|
||
|
||
receipt, code = self._dispatch(launcher, RecordingConnect())
|
||
self.assertEqual(code, 0)
|
||
self.assertEqual(receipt["status"], "completed")
|
||
self.assertEqual(seen["cwd"], str(REPO_ROOT.resolve()))
|
||
|
||
def test_framework_failure_marks_run_failed_and_persists_trace(self):
|
||
connect = RecordingConnect()
|
||
receipt, code = self._dispatch(fake_launcher([b"garbage\n"]), connect)
|
||
self.assertEqual(code, 3)
|
||
self.assertEqual(receipt["errorCode"], "STREAM_PARSE_ERROR")
|
||
self.assertEqual(receipt["evidence"]["status"], "written")
|
||
self.assertEqual(receipt["evidence"]["llmCallIds"], [])
|
||
|
||
def test_role_model_policy_is_checked_before_run_start(self):
|
||
connect = RecordingConnect()
|
||
receipt, code = run_dispatch(
|
||
make_spec(self.tmp),
|
||
repo_root=REPO_ROOT,
|
||
policy=ExecutionPolicy(provider="p", model="gpt-5.6-sol"),
|
||
run_id="role-policy-test",
|
||
run_dir=self.tmp / "role-policy-run",
|
||
connect_factory=connect,
|
||
launcher=fake_launcher(pi_stream_lines('{"title":"t","beats":["b"]}')),
|
||
)
|
||
self.assertEqual(code, 2)
|
||
self.assertEqual(receipt["errorCode"], "ROLE_MODEL_POLICY_MISMATCH")
|
||
self.assertFalse(any(sql.startswith("INSERT INTO example_run") for sql, _ in connect.log))
|
||
|
||
def test_trigger_detail_secret_is_rejected_before_run_start(self):
|
||
connect = RecordingConnect()
|
||
receipt, code = run_dispatch(
|
||
make_spec(self.tmp),
|
||
repo_root=REPO_ROOT,
|
||
policy=ExecutionPolicy(provider="p", model="claude-opus-test"),
|
||
run_id="safe-trigger-test",
|
||
run_dir=self.tmp / "trigger-run",
|
||
connect_factory=connect,
|
||
launcher=fake_launcher(pi_stream_lines('{"title":"t","beats":["b"]}')),
|
||
trigger_detail={"api_key": "sk-abcdef0123456789abcdef012345"},
|
||
)
|
||
self.assertEqual(code, 2)
|
||
self.assertEqual(receipt["errorCode"], "TRIGGER_DETAIL_INVALID")
|
||
self.assertFalse(any(sql.startswith("INSERT INTO example_run") for sql, _ in connect.log))
|
||
|
||
def test_run_id_cannot_escape_audit_root(self):
|
||
receipt, code = run_dispatch(
|
||
make_spec(self.tmp),
|
||
repo_root=REPO_ROOT,
|
||
policy=ExecutionPolicy(provider="p", model="claude-opus-test"),
|
||
run_id="../escape",
|
||
run_dir=self.tmp / "should-not-exist",
|
||
connect_factory=RecordingConnect(),
|
||
launcher=fake_launcher(pi_stream_lines('{"title":"t","beats":["b"]}')),
|
||
)
|
||
self.assertEqual(code, 2)
|
||
self.assertEqual(receipt["errorCode"], "RUN_ID_INVALID")
|
||
self.assertFalse((self.tmp / "should-not-exist").exists())
|
||
|
||
|
||
class ValidateOutputTest(unittest.TestCase):
|
||
def test_extract_and_validate(self):
|
||
from agent_task import AgentTaskSpec
|
||
|
||
spec = AgentTaskSpec(
|
||
role="planner",
|
||
task_prompt="p",
|
||
input={},
|
||
output_schema=OUTPUT_SCHEMA,
|
||
output_schema_id="s1",
|
||
tool_allowlist=(),
|
||
)
|
||
out = validate_structured_output('前言```json\n{"title":"t","beats":["b"]}\n```', spec)
|
||
self.assertEqual(out["title"], "t")
|
||
with self.assertRaises(OutputInvalidError):
|
||
validate_structured_output('{"title":"t"}', spec)
|
||
|
||
|
||
|
||
|
||
class SessionAndReadToolsTest(unittest.TestCase):
|
||
"""阶段 D:会话复用、工具 server 开关与依赖清单。"""
|
||
|
||
def setUp(self):
|
||
self.tmp = pathlib.Path(tempfile.mkdtemp())
|
||
|
||
def test_session_requires_dir_and_maps_to_argv(self):
|
||
with self.assertRaises(ValueError):
|
||
ExecutionPolicy(provider="p", model="m", session_id="s1")
|
||
package = build_task_package(load_spec(make_spec(self.tmp)), REPO_ROOT)
|
||
argv = build_pi_argv(
|
||
package,
|
||
ExecutionPolicy(provider="p", model="m", session_id="s1", session_dir="/tmp/sd"),
|
||
)
|
||
self.assertIn("--session-id", argv)
|
||
self.assertEqual(argv[argv.index("--session-id") + 1], "s1")
|
||
self.assertEqual(argv[argv.index("--session-dir") + 1], "/tmp/sd")
|
||
self.assertNotIn("--no-session", argv)
|
||
|
||
def test_extension_path_maps_to_e_flag_with_isolation(self):
|
||
package = build_task_package(load_spec(make_spec(self.tmp)), REPO_ROOT)
|
||
argv = build_pi_argv(
|
||
package,
|
||
ExecutionPolicy(provider="p", model="m", extension_path="/x/ext.ts"),
|
||
)
|
||
self.assertIn("-e", argv)
|
||
self.assertEqual(argv[argv.index("-e") + 1], "/x/ext.ts")
|
||
# 自动发现仍被禁:只加载显式扩展。
|
||
self.assertIn("--no-extensions", argv)
|
||
|
||
def test_tool_args_captured_for_dependency_manifest(self):
|
||
sink_holder = {}
|
||
|
||
class RecordingSink(AgentTraceWriter):
|
||
def __init__(self):
|
||
super().__init__(
|
||
run_id="t", framework="pi", agent_role="planner",
|
||
connect=lambda: (_ for _ in ()).throw(AssertionError("离线测试不应触库")),
|
||
)
|
||
self.rows = []
|
||
|
||
def emit(self, event_type, **kwargs):
|
||
self.rows.append((event_type, kwargs))
|
||
return len(self.rows)
|
||
|
||
sink = RecordingSink()
|
||
sink_holder["sink"] = sink
|
||
lines = pi_stream_lines('{"title":"t","beats":["b"]}', with_tool=True)
|
||
# 给工具事件注入真实参数形状(依赖清单材料)。
|
||
patched = []
|
||
for raw in lines:
|
||
event = json.loads(raw)
|
||
if event.get("type") == "tool_execution_start":
|
||
event["args"] = {"work_id": 12, "target_chapter": 4}
|
||
patched.append(json.dumps(event, ensure_ascii=False).encode("utf-8") + b"\n")
|
||
package = build_task_package(
|
||
load_spec(make_spec(self.tmp, toolAllowlist=["read"])), REPO_ROOT
|
||
)
|
||
runner = PiAgentRunner(launcher=fake_launcher(patched))
|
||
outcome = runner.run(
|
||
package,
|
||
ExecutionPolicy(provider="p", model="m"),
|
||
sink,
|
||
timeout_seconds=30,
|
||
)
|
||
self.assertEqual(outcome.tool_calls[0].args, {"work_id": 12, "target_chapter": 4})
|
||
started = [kwargs for kind, kwargs in sink.rows if kind == "tool.started"]
|
||
self.assertIn("work_id", started[0]["details"]["args"])
|
||
|
||
def test_read_tools_disabled_fails_closed(self):
|
||
spec_path = make_spec(self.tmp, toolAllowlist=["read_fine_outline"])
|
||
receipt, code = run_dispatch(
|
||
spec_path,
|
||
repo_root=REPO_ROOT,
|
||
policy=ExecutionPolicy(provider="p", model="claude-opus-test"),
|
||
run_id="unittest-agent-dispatch-rt",
|
||
run_dir=self.tmp / "run",
|
||
connect_factory=RecordingConnect(),
|
||
launcher=fake_launcher(pi_stream_lines('{"title":"t","beats":["b"]}')),
|
||
trigger_source="diagnostic",
|
||
)
|
||
self.assertEqual(code, EXIT_SPEC_INVALID)
|
||
self.assertEqual(receipt["errorCode"], "READ_TOOLS_NOT_ENABLED")
|
||
|
||
def test_dependency_manifest_written_on_tool_use(self):
|
||
spec_path = make_spec(self.tmp, toolAllowlist=["read"])
|
||
run_dir = self.tmp / "run"
|
||
receipt, code = run_dispatch(
|
||
spec_path,
|
||
repo_root=REPO_ROOT,
|
||
policy=ExecutionPolicy(provider="p", model="claude-opus-test"),
|
||
run_id="unittest-agent-dispatch-dep",
|
||
run_dir=run_dir,
|
||
connect_factory=RecordingConnect(),
|
||
launcher=fake_launcher(pi_stream_lines('{"title":"t","beats":["b"]}', with_tool=True)),
|
||
trigger_source="diagnostic",
|
||
enable_read_tools=True,
|
||
)
|
||
self.assertEqual(code, 0, receipt)
|
||
dependencies = json.loads((run_dir / "dependencies.json").read_text(encoding="utf-8"))
|
||
self.assertEqual(len(dependencies), 1)
|
||
self.assertEqual(dependencies[0]["tool"], "read")
|
||
self.assertEqual(dependencies[0]["seq"], 1)
|
||
self.assertEqual(receipt["dependencies"], {"count": 1, "file": "dependencies.json"})
|
||
# 无工具调用的运行不产依赖清单文件。
|
||
tmp2 = pathlib.Path(tempfile.mkdtemp())
|
||
run_dir2 = tmp2 / "run"
|
||
receipt2, code2 = run_dispatch(
|
||
make_spec(tmp2),
|
||
repo_root=REPO_ROOT,
|
||
policy=ExecutionPolicy(provider="p", model="claude-opus-test"),
|
||
run_id="unittest-agent-dispatch-nodep",
|
||
run_dir=run_dir2,
|
||
connect_factory=RecordingConnect(),
|
||
launcher=fake_launcher(pi_stream_lines('{"title":"t","beats":["b"]}')),
|
||
trigger_source="diagnostic",
|
||
)
|
||
self.assertEqual(code2, 0, receipt2)
|
||
self.assertFalse((run_dir2 / "dependencies.json").exists())
|
||
self.assertIsNone(receipt2["dependencies"])
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main(verbosity=2)
|