test_agent_surface.py 16 KB

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