test_agent_surface.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430
  1. from __future__ import annotations
  2. from dataclasses import replace
  3. from datetime import UTC, datetime
  4. from hashlib import sha256
  5. from types import SimpleNamespace
  6. import pytest
  7. from agent import AgentRole, ToolRegistry, get_preset
  8. from script_build_host.agents.model_resolver import (
  9. ScriptRoleRunConfigResolver,
  10. SnapshotModelManifestSource,
  11. )
  12. from script_build_host.agents.presets import register_script_presets
  13. from script_build_host.agents.prompt_resolver import SnapshotRoleSystemPromptResolver
  14. from script_build_host.agents.prompts import (
  15. CONTEXT_ACCESS_PROTOCOL,
  16. phase_three_prompt_manifest,
  17. script_build_prompt_manifest,
  18. )
  19. from script_build_host.domain.errors import ProtocolViolation
  20. from script_build_host.domain.input_snapshot import ScriptBuildInputSnapshotV1
  21. from script_build_host.domain.workbench import protect_model_value
  22. from script_build_host.tools.registry import register_script_tools
  23. def test_script_presets_are_exact_and_never_expose_legacy_control_tools() -> None:
  24. register_script_presets()
  25. expected = {
  26. "script_planner",
  27. "script_direction_worker",
  28. "script_evidence_retrieval_worker",
  29. "script_retrieval_validator",
  30. "script_candidate_validator",
  31. "script_structure_worker",
  32. "script_paragraph_worker",
  33. "script_element_set_worker",
  34. "script_candidate_compare_worker",
  35. "script_compose_worker",
  36. "script_candidate_portfolio_worker",
  37. }
  38. forbidden = {
  39. "agent",
  40. "evaluate",
  41. "bash_command",
  42. "write_file",
  43. "dispatch_tasks",
  44. "implement_paths",
  45. "begin_round",
  46. "multipath",
  47. }
  48. for name in expected:
  49. preset = get_preset(name)
  50. assert preset.role in {AgentRole.PLANNER, AgentRole.WORKER, AgentRole.VALIDATOR}
  51. assert forbidden.isdisjoint(set(preset.allowed_tools or []))
  52. assert set(preset.allowed_tools or []).isdisjoint(set(preset.denied_tools or []))
  53. planner = get_preset("script_planner")
  54. assert set(planner.allowed_tools or []) == {
  55. "get_current_context",
  56. "plan_script_tasks",
  57. "decide_script_task",
  58. "search_mission_context",
  59. "read_mission_context",
  60. "dispatch_script_tasks",
  61. "validate_attempt",
  62. }
  63. compose = get_preset("script_compose_worker")
  64. assert set(compose.allowed_tools or []) == {"submit_attempt"}
  65. compare = get_preset("script_candidate_compare_worker")
  66. assert "save_comparison_candidate" in (compare.allowed_tools or [])
  67. assert "create_script_paragraph" not in (compare.allowed_tools or [])
  68. portfolio = get_preset("script_candidate_portfolio_worker")
  69. assert set(portfolio.allowed_tools or []) == {"submit_attempt"}
  70. assert "view_frozen_images" not in (
  71. get_preset("script_retrieval_validator").allowed_tools or []
  72. )
  73. assert "view_frozen_images" in (get_preset("script_candidate_validator").allowed_tools or [])
  74. element_tools = set(get_preset("script_element_set_worker").allowed_tools or [])
  75. assert "save_script_elements" in element_tools
  76. assert "create_script_element" not in element_tools
  77. assert "batch_link_paragraph_elements" not in element_tools
  78. assert set(get_preset("script_root_worker").allowed_tools or []) == {
  79. "search_mission_context",
  80. "read_mission_context",
  81. "legacy_projection_dry_run",
  82. "save_root_delivery_manifest",
  83. "submit_attempt",
  84. }
  85. assert set(get_preset("script_root_validator").allowed_tools or []) == {
  86. "search_mission_context",
  87. "read_mission_context",
  88. "query_validation_evidence",
  89. "deterministic_precheck",
  90. "submit_validation",
  91. }
  92. def test_new_build_prompt_manifest_freezes_every_phase_one_and_phase_two_role() -> None:
  93. manifest = script_build_prompt_manifest()
  94. assert len(manifest) == 12
  95. assert len({str(item["preset"]) for item in manifest}) == len(manifest)
  96. assert {str(item["preset"]) for item in manifest} >= {
  97. "script_planner",
  98. "script_context_broker",
  99. "script_structure_worker",
  100. "script_paragraph_worker",
  101. "script_element_set_worker",
  102. "script_candidate_compare_worker",
  103. "script_compose_worker",
  104. "script_candidate_portfolio_worker",
  105. }
  106. for item in manifest:
  107. assert str(item["content"])
  108. assert str(item["content_sha256"]).startswith("sha256:")
  109. assert len(str(item["content_sha256"])) == 71
  110. def test_phase_three_prompt_manifest_is_versioned_separately() -> None:
  111. manifest = phase_three_prompt_manifest()
  112. assert {str(item["preset"]) for item in manifest} == {
  113. "script_root_worker",
  114. "script_root_validator",
  115. }
  116. assert all(str(item["content_sha256"]).startswith("sha256:") for item in manifest)
  117. def test_every_model_role_freezes_one_shared_context_access_protocol() -> None:
  118. manifest = {
  119. str(item["preset"]): str(item["content"])
  120. for item in (*script_build_prompt_manifest(), *phase_three_prompt_manifest())
  121. }
  122. deterministic_or_utility = {
  123. "script_context_broker",
  124. "script_compose_worker",
  125. "script_candidate_portfolio_worker",
  126. "script_evidence_retrieval_worker",
  127. }
  128. for preset, content in manifest.items():
  129. if preset in deterministic_or_utility:
  130. assert not content.startswith(CONTEXT_ACCESS_PROTOCOL)
  131. else:
  132. assert content.startswith(CONTEXT_ACCESS_PROTOCOL), preset
  133. assert content.count("## Context Access Protocol") == 1
  134. assert "active_has_more" in manifest["script_planner"]
  135. assert "不由模型选择" in manifest["script_evidence_retrieval_worker"]
  136. assert "Direction、CandidatePortfolio" in manifest["script_root_worker"]
  137. assert "required_handles" in manifest["script_candidate_validator"]
  138. assert "required_handles" in manifest["script_root_validator"]
  139. class _Gateway:
  140. coordinator = SimpleNamespace(task_store=None)
  141. def test_physical_retrieval_workers_are_not_agent_tools() -> None:
  142. registry = ToolRegistry()
  143. register_script_tools(registry, _Gateway()) # type: ignore[arg-type]
  144. schemas = {
  145. item["function"]["name"]: item["function"]["parameters"] for item in registry.get_schemas()
  146. }
  147. descriptions = {
  148. item["function"]["name"]: item["function"]["description"] for item in registry.get_schemas()
  149. }
  150. assert {
  151. "search_script_decode_case",
  152. "external_search_case",
  153. "search_knowledge",
  154. "query_pattern_qa",
  155. }.isdisjoint(schemas)
  156. assert schemas["search_mission_context"]["required"] == ["collection"]
  157. assert schemas["read_mission_context"]["required"] == ["handle"]
  158. assert "sufficiency" in descriptions["search_mission_context"]
  159. assert "exhausted=true" in descriptions["read_mission_context"]
  160. assert "load_images" not in schemas
  161. assert "load_frozen_strategy" not in schemas
  162. assert "read_workbench_detail" not in schemas
  163. assert schemas["save_script_elements"]["required"] == [
  164. "elements",
  165. "links",
  166. "expected_state_revision",
  167. ]
  168. assert schemas["save_script_elements"]["properties"]["links"]["items"]["required"] == [
  169. "paragraph_target_key",
  170. "element_client_keys",
  171. ]
  172. assert schemas["submit_attempt"].get("properties") == {}
  173. assert schemas["submit_validation"]["required"] == [
  174. "verdict",
  175. "criterion_results",
  176. "defects",
  177. ]
  178. assert "summary" not in schemas["submit_attempt"].get("properties", {})
  179. assert "artifact_refs" not in schemas["submit_attempt"].get("properties", {})
  180. def test_agent_visible_write_schemas_never_expose_mechanical_identity_fields() -> None:
  181. registry = ToolRegistry()
  182. register_script_tools(registry, _Gateway()) # type: ignore[arg-type]
  183. schemas = {
  184. item["function"]["name"]: item["function"]["parameters"] for item in registry.get_schemas()
  185. }
  186. write_tools = {
  187. "plan_script_tasks",
  188. "decide_script_task",
  189. "save_direction_candidate",
  190. "save_script_paragraphs",
  191. "save_script_elements",
  192. "save_comparison_candidate",
  193. "save_root_delivery_manifest",
  194. "submit_validation",
  195. }
  196. forbidden = {
  197. "paragraph_id",
  198. "element_id",
  199. "branch_id",
  200. "artifact_ref",
  201. "evidence_refs",
  202. "write_scope",
  203. "output_schema",
  204. "compose_order",
  205. "context_refs",
  206. }
  207. def property_names(value: object) -> set[str]:
  208. if isinstance(value, dict):
  209. names = (
  210. set(value.get("properties", {}))
  211. if isinstance(value.get("properties"), dict)
  212. else set()
  213. )
  214. return names | set().union(*(property_names(item) for item in value.values()))
  215. if isinstance(value, list):
  216. return set().union(*(property_names(item) for item in value))
  217. return set()
  218. for name in write_tools:
  219. assert forbidden.isdisjoint(property_names(schemas[name])), name
  220. register_script_presets()
  221. forbidden_tools = {
  222. "read_input_snapshot",
  223. "read_attempt_workspace",
  224. "read_active_frontier",
  225. "read_pinned_candidates",
  226. "read_accepted_artifact",
  227. "read_workbench_detail",
  228. "load_frozen_strategy",
  229. "load_images",
  230. "create_script_paragraph",
  231. "create_script_element",
  232. "update_script_element",
  233. "batch_link_paragraph_elements",
  234. "append_paragraph_atoms",
  235. "delete_paragraph_atom",
  236. }
  237. for preset_name in (
  238. "script_direction_worker",
  239. "script_structure_worker",
  240. "script_paragraph_worker",
  241. "script_element_set_worker",
  242. "script_candidate_compare_worker",
  243. "script_compose_worker",
  244. "script_candidate_portfolio_worker",
  245. "script_root_worker",
  246. ):
  247. assert forbidden_tools.isdisjoint(get_preset(preset_name).allowed_tools or []), preset_name
  248. def test_model_visible_values_replace_storage_identity_and_uri_fields() -> None:
  249. protected = protect_model_value(
  250. {
  251. "uri": "script-build://artifact-versions/9",
  252. "url": "https://private.example/source",
  253. "raw_artifact_ref": "script-build://raw-artifacts/sha256/secret",
  254. "artifact_version_id": 9,
  255. "branch_id": 7,
  256. "strategy_id": 11,
  257. "topic_id": 12,
  258. "topic_build_id": 13,
  259. "execution_id": 14,
  260. "content_sha256": "sha256:secret",
  261. "goal_id": "goal-allowed",
  262. "decision_id": "decision-allowed",
  263. },
  264. namespace="surface-test",
  265. )
  266. encoded = str(protected)
  267. assert "script-build://" not in encoded
  268. assert "https://" not in encoded
  269. assert "sha256:" not in encoded
  270. assert "artifact_version_id" not in protected
  271. assert "branch_id" not in protected
  272. assert "strategy_id" not in protected
  273. assert "topic_id" not in protected
  274. assert "topic_build_id" not in protected
  275. assert "execution_id" not in protected
  276. assert protected["goal_id"] == "goal-allowed"
  277. assert protected["decision_id"] == "decision-allowed"
  278. class _Bindings:
  279. async def get_by_root(self, root_trace_id: str) -> SimpleNamespace:
  280. assert root_trace_id == "root-a"
  281. return SimpleNamespace(script_build_id=9, input_snapshot_id=3)
  282. class _Snapshots:
  283. async def get(self, snapshot_id: str, *, script_build_id: int) -> ScriptBuildInputSnapshotV1:
  284. assert (snapshot_id, script_build_id) == ("3", 9)
  285. return ScriptBuildInputSnapshotV1(
  286. snapshot_id="3",
  287. script_build_id=9,
  288. execution_id=1,
  289. topic_build_id=2,
  290. topic_id=3,
  291. topic={},
  292. account={},
  293. persona_points=(),
  294. section_patterns=(),
  295. strategies=(),
  296. prompt_manifest=(),
  297. datasource_manifest={},
  298. model_manifest={
  299. "presets": {
  300. "script_direction_worker": {
  301. "model": "frozen-direction-model",
  302. "temperature": 0.2,
  303. "max_iterations": 18,
  304. }
  305. }
  306. },
  307. canonical_sha256="sha256:" + "a" * 64,
  308. created_at=datetime.now(UTC),
  309. )
  310. @pytest.mark.asyncio
  311. async def test_role_model_resolver_reads_each_roots_frozen_manifest() -> None:
  312. resolver = ScriptRoleRunConfigResolver(
  313. {}, manifest_source=SnapshotModelManifestSource(_Bindings(), _Snapshots())
  314. )
  315. value = await resolver.resolve(
  316. role=AgentRole.WORKER,
  317. preset="script_direction_worker",
  318. context={"root_trace_id": "root-a"},
  319. )
  320. assert value.model == "frozen-direction-model"
  321. assert value.temperature == 0.2
  322. assert value.max_iterations == 18
  323. @pytest.mark.asyncio
  324. async def test_role_prompt_resolver_uses_task_pinned_snapshot_and_rechecks_digest() -> None:
  325. content = "Frozen paragraph worker policy"
  326. snapshot = ScriptBuildInputSnapshotV1(
  327. snapshot_id="3",
  328. script_build_id=9,
  329. execution_id=1,
  330. topic_build_id=2,
  331. topic_id=3,
  332. topic={},
  333. account={},
  334. persona_points=(),
  335. section_patterns=(),
  336. strategies=(),
  337. prompt_manifest=(
  338. {
  339. "preset": "script_paragraph_worker",
  340. "role": "worker",
  341. "content": content,
  342. "content_sha256": f"sha256:{sha256(content.encode('utf-8')).hexdigest()}",
  343. },
  344. ),
  345. datasource_manifest={},
  346. model_manifest={},
  347. canonical_sha256="sha256:" + "a" * 64,
  348. created_at=datetime.now(UTC),
  349. )
  350. class Snapshots:
  351. async def get(self, snapshot_id: str, *, script_build_id: int):
  352. assert (snapshot_id, script_build_id) == ("3", 9)
  353. return snapshot
  354. resolver = SnapshotRoleSystemPromptResolver(_Bindings(), Snapshots())
  355. value = await resolver.resolve(
  356. role=AgentRole.WORKER,
  357. preset="script_paragraph_worker",
  358. context={
  359. "root_trace_id": "root-a",
  360. "task_spec": {"context_refs": ["script-build://inputs/3"]},
  361. },
  362. )
  363. assert value is not None
  364. assert value.content == content
  365. tampered = replace(
  366. snapshot,
  367. prompt_manifest=(
  368. {
  369. **snapshot.prompt_manifest[0],
  370. "content_sha256": "sha256:" + "f" * 64,
  371. },
  372. ),
  373. )
  374. class TamperedSnapshots:
  375. async def get(self, _snapshot_id: str, *, script_build_id: int):
  376. assert script_build_id == 9
  377. return tampered
  378. with pytest.raises(ProtocolViolation, match="digest"):
  379. await SnapshotRoleSystemPromptResolver(_Bindings(), TamperedSnapshots()).resolve(
  380. role=AgentRole.WORKER,
  381. preset="script_paragraph_worker",
  382. context={
  383. "root_trace_id": "root-a",
  384. "task_spec": {"context_refs": ["script-build://inputs/3"]},
  385. },
  386. )
  387. async def _async(value):
  388. return value