test_orchestration_system_prompt_resolver.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388
  1. from __future__ import annotations
  2. import asyncio
  3. from hashlib import sha256
  4. from types import SimpleNamespace
  5. from typing import Any, Mapping
  6. import pytest
  7. from agent import (
  8. AgentPreset,
  9. AgentRole,
  10. AgentRunner,
  11. FileSystemArtifactStore,
  12. FileSystemTaskStore,
  13. FileSystemTraceStore,
  14. RoleSystemPromptOverride,
  15. RoleSystemPromptResolver,
  16. wire_orchestration,
  17. )
  18. from agent.core.presets import register_preset
  19. from agent.orchestration.executor import LocalAgentExecutor
  20. from agent.orchestration.models import FailureCode
  21. def _override(content: str) -> RoleSystemPromptOverride:
  22. return RoleSystemPromptOverride(
  23. content=content,
  24. prompt_identity=f"sha256:{sha256(content.encode('utf-8')).hexdigest()}",
  25. )
  26. def _worker_context(**updates: Any) -> dict[str, Any]:
  27. context = {
  28. "worker_preset": "script_worker",
  29. "worker_trace_id": "worker-trace",
  30. "root_trace_id": "root",
  31. "task_id": "task",
  32. "spec_version": 1,
  33. "attempt_id": "attempt",
  34. "task_spec": {"objective": "write"},
  35. "continue_trace_id": None,
  36. }
  37. context.update(updates)
  38. return context
  39. def _validator_context(**updates: Any) -> dict[str, Any]:
  40. context = {
  41. "validator_preset": "script_validator",
  42. "validator_trace_id": "validator-trace",
  43. "root_trace_id": "root",
  44. "task_id": "task",
  45. "spec_version": 1,
  46. "attempt_id": "attempt",
  47. "snapshot_id": "snapshot",
  48. "validation_id": "validation",
  49. "task_spec": {"objective": "validate"},
  50. "artifact_snapshot": {"artifact_refs": []},
  51. }
  52. context.update(updates)
  53. return context
  54. class PromptResolver:
  55. def __init__(self, override: RoleSystemPromptOverride | None) -> None:
  56. self.override = override
  57. self.calls: list[tuple[AgentRole, str, Mapping[str, Any]]] = []
  58. async def resolve(self, *, role, preset, context):
  59. with pytest.raises(TypeError):
  60. context["unsafe"] = True
  61. self.calls.append((role, preset, context))
  62. return self.override
  63. class CapturingRunner:
  64. trace_store = None
  65. task_coordinator = None
  66. def __init__(self) -> None:
  67. self.configs = []
  68. async def run_result(self, *, messages, config):
  69. self.configs.append(config)
  70. return {"status": "completed", "summary": "done"}
  71. def test_system_prompt_types_are_public_and_validate_content_identity():
  72. assert RoleSystemPromptResolver.__name__ == "RoleSystemPromptResolver"
  73. value = _override("frozen prompt")
  74. assert value.content == "frozen prompt"
  75. with pytest.raises(ValueError, match="does not match"):
  76. RoleSystemPromptOverride(
  77. content="frozen prompt",
  78. prompt_identity=f"sha256:{'0' * 64}",
  79. )
  80. with pytest.raises(ValueError, match="lowercase"):
  81. RoleSystemPromptOverride(content="x", prompt_identity=f"sha256:{'A' * 64}")
  82. @pytest.mark.asyncio
  83. async def test_new_role_trace_receives_only_prompt_and_protected_identity():
  84. runner = CapturingRunner()
  85. resolver = PromptResolver(_override("business policy"))
  86. result = await LocalAgentExecutor(
  87. runner, role_system_prompt_resolver=resolver
  88. ).run_worker(_worker_context(untrusted="not protected"))
  89. assert result.status == "completed"
  90. assert resolver.calls[0][0:2] == (AgentRole.WORKER, "script_worker")
  91. config = runner.configs[0]
  92. assert config.system_prompt == "business policy"
  93. assert (
  94. config.context["role_prompt_identity"]
  95. == _override("business policy").prompt_identity
  96. )
  97. assert "untrusted" not in config.context
  98. assert config.tools is None and config.tool_groups is None
  99. @pytest.mark.asyncio
  100. async def test_new_validator_trace_receives_frozen_prompt_identity():
  101. runner = CapturingRunner()
  102. expected = _override("validator business policy")
  103. resolver = PromptResolver(expected)
  104. result = await LocalAgentExecutor(
  105. runner, role_system_prompt_resolver=resolver
  106. ).run_validator(_validator_context(untrusted="not protected"))
  107. assert result.status == "completed"
  108. assert resolver.calls[0][0:2] == (AgentRole.VALIDATOR, "script_validator")
  109. config = runner.configs[0]
  110. assert config.system_prompt == expected.content
  111. assert config.context["role_prompt_identity"] == expected.prompt_identity
  112. assert config.context["validation_id"] == "validation"
  113. assert "untrusted" not in config.context
  114. @pytest.mark.asyncio
  115. async def test_missing_or_empty_prompt_resolver_keeps_existing_defaults():
  116. no_resolver_runner = CapturingRunner()
  117. empty_resolver_runner = CapturingRunner()
  118. no_resolver = await LocalAgentExecutor(no_resolver_runner).run_worker(
  119. _worker_context()
  120. )
  121. empty_resolver = await LocalAgentExecutor(
  122. empty_resolver_runner,
  123. role_system_prompt_resolver=PromptResolver(None),
  124. ).run_worker(_worker_context())
  125. assert no_resolver.status == empty_resolver.status == "completed"
  126. assert no_resolver_runner.configs[0].system_prompt is None
  127. assert empty_resolver_runner.configs[0].system_prompt is None
  128. assert "role_prompt_identity" not in no_resolver_runner.configs[0].context
  129. assert "role_prompt_identity" not in empty_resolver_runner.configs[0].context
  130. @pytest.mark.asyncio
  131. async def test_invalid_prompt_resolver_result_is_an_executor_failure():
  132. class InvalidResolver:
  133. async def resolve(self, **kwargs):
  134. return {"content": "unsafe", "prompt_identity": "forged"}
  135. runner = CapturingRunner()
  136. result = await LocalAgentExecutor(
  137. runner, role_system_prompt_resolver=InvalidResolver()
  138. ).run_worker(_worker_context())
  139. assert result.status == "failed"
  140. assert "must return RoleSystemPromptOverride or None" in result.error
  141. assert runner.configs == []
  142. @pytest.mark.asyncio
  143. async def test_prompt_resolver_exception_is_normalized_as_executor_failure():
  144. class BrokenResolver:
  145. async def resolve(self, **kwargs):
  146. raise RuntimeError("frozen prompt manifest unavailable")
  147. runner = CapturingRunner()
  148. result = await LocalAgentExecutor(
  149. runner, role_system_prompt_resolver=BrokenResolver()
  150. ).run_validator(_validator_context())
  151. assert result.status == "failed"
  152. assert result.error == "frozen prompt manifest unavailable"
  153. assert result.execution_stats.failure_code == FailureCode.EXECUTOR_ERROR
  154. assert runner.configs == []
  155. @pytest.mark.asyncio
  156. async def test_concurrent_roles_keep_distinct_frozen_prompts_and_identities():
  157. class PerTaskResolver:
  158. async def resolve(self, *, role, preset, context):
  159. del preset
  160. await asyncio.sleep(0)
  161. return _override(f"{role.value} policy for {context['task_id']}")
  162. runner = CapturingRunner()
  163. executor = LocalAgentExecutor(runner, role_system_prompt_resolver=PerTaskResolver())
  164. worker, validator = await asyncio.gather(
  165. executor.run_worker(
  166. _worker_context(task_id="worker-task", worker_trace_id="worker-trace")
  167. ),
  168. executor.run_validator(
  169. _validator_context(
  170. task_id="validator-task", validator_trace_id="validator-trace"
  171. )
  172. ),
  173. )
  174. assert worker.status == validator.status == "completed"
  175. configs = {config.name: config for config in runner.configs}
  176. worker_config = configs["Worker worker-task"]
  177. validator_config = configs["Validator validator-task"]
  178. assert worker_config.system_prompt == "worker policy for worker-task"
  179. assert validator_config.system_prompt == "validator policy for validator-task"
  180. assert (
  181. worker_config.context["role_prompt_identity"]
  182. == _override("worker policy for worker-task").prompt_identity
  183. )
  184. assert (
  185. validator_config.context["role_prompt_identity"]
  186. == _override("validator policy for validator-task").prompt_identity
  187. )
  188. @pytest.mark.asyncio
  189. async def test_real_runner_persists_worker_and_validator_prompt_identity(tmp_path):
  190. worker_preset = "test_prompt_identity_worker"
  191. validator_preset = "test_prompt_identity_validator"
  192. register_preset(
  193. worker_preset,
  194. AgentPreset(role=AgentRole.WORKER, allowed_tools=[], max_iterations=1),
  195. )
  196. register_preset(
  197. validator_preset,
  198. AgentPreset(role=AgentRole.VALIDATOR, allowed_tools=[], max_iterations=1),
  199. )
  200. async def llm_call(**_kwargs):
  201. return {"content": "no terminal tool", "tool_calls": None}
  202. store = FileSystemTraceStore(str(tmp_path))
  203. runner = AgentRunner(
  204. trace_store=store,
  205. llm_call=llm_call,
  206. task_coordinator=object(),
  207. )
  208. class PerRoleResolver:
  209. async def resolve(self, *, role, preset, context):
  210. del preset, context
  211. return _override(f"frozen {role.value} prompt")
  212. executor = LocalAgentExecutor(runner, role_system_prompt_resolver=PerRoleResolver())
  213. worker = await executor.run_worker(_worker_context(worker_preset=worker_preset))
  214. validator = await executor.run_validator(
  215. _validator_context(validator_preset=validator_preset)
  216. )
  217. assert worker.status == validator.status == "failed"
  218. worker_trace = await store.get_trace("worker-trace")
  219. validator_trace = await store.get_trace("validator-trace")
  220. assert (
  221. worker_trace.context["role_prompt_identity"]
  222. == _override("frozen worker prompt").prompt_identity
  223. )
  224. assert (
  225. validator_trace.context["role_prompt_identity"]
  226. == _override("frozen validator prompt").prompt_identity
  227. )
  228. worker_path = await store.get_main_path_messages(
  229. worker_trace.trace_id, worker_trace.head_sequence
  230. )
  231. validator_path = await store.get_main_path_messages(
  232. validator_trace.trace_id, validator_trace.head_sequence
  233. )
  234. assert "frozen worker prompt" in str(worker_path[0].content)
  235. assert "frozen validator prompt" in str(validator_path[0].content)
  236. @pytest.mark.asyncio
  237. async def test_repair_validates_identity_without_resetting_system_prompt():
  238. expected = _override("original frozen prompt")
  239. trace = SimpleNamespace(
  240. context={"role_prompt_identity": expected.prompt_identity},
  241. total_tokens=0,
  242. total_cost=0.0,
  243. model="worker-model",
  244. )
  245. class Coordinator:
  246. async def validate_continue_from(self, *args):
  247. return "worker-trace"
  248. class Store:
  249. async def get_trace(self, trace_id):
  250. return trace
  251. async def update_trace(self, trace_id, **updates):
  252. trace.context = updates["context"]
  253. class Runner(CapturingRunner):
  254. def __init__(self):
  255. super().__init__()
  256. self.trace_store = Store()
  257. self.task_coordinator = Coordinator()
  258. runner = Runner()
  259. result = await LocalAgentExecutor(
  260. runner, role_system_prompt_resolver=PromptResolver(expected)
  261. ).run_worker(
  262. _worker_context(
  263. attempt_id="attempt-2",
  264. prior_attempt_id="attempt-1",
  265. continue_trace_id="worker-trace",
  266. )
  267. )
  268. assert result.status == "completed"
  269. assert runner.configs[0].system_prompt is None
  270. assert trace.context["role_prompt_identity"] == expected.prompt_identity
  271. assert trace.context["attempt_id"] == "attempt-2"
  272. @pytest.mark.asyncio
  273. async def test_repair_prompt_mismatch_fails_before_trace_context_mutation():
  274. original = _override("original")
  275. trace = SimpleNamespace(
  276. context={"role_prompt_identity": original.prompt_identity, "attempt_id": "old"},
  277. total_tokens=0,
  278. total_cost=0.0,
  279. model="worker-model",
  280. )
  281. class Coordinator:
  282. async def validate_continue_from(self, *args):
  283. return "worker-trace"
  284. class Store:
  285. updated = False
  286. async def get_trace(self, trace_id):
  287. return trace
  288. async def update_trace(self, trace_id, **updates):
  289. self.updated = True
  290. runner = CapturingRunner()
  291. runner.trace_store = Store()
  292. runner.task_coordinator = Coordinator()
  293. result = await LocalAgentExecutor(
  294. runner, role_system_prompt_resolver=PromptResolver(_override("changed"))
  295. ).run_worker(
  296. _worker_context(
  297. attempt_id="new",
  298. prior_attempt_id="old",
  299. continue_trace_id="worker-trace",
  300. )
  301. )
  302. assert result.status == "failed"
  303. assert "ROLE_SYSTEM_PROMPT_MISMATCH" in result.error
  304. assert runner.trace_store.updated is False
  305. assert runner.configs == []
  306. @pytest.mark.asyncio
  307. async def test_wire_injects_optional_system_prompt_resolver(tmp_path):
  308. runner = CapturingRunner()
  309. resolver = PromptResolver(_override("wired prompt"))
  310. coordinator = wire_orchestration(
  311. runner,
  312. FileSystemTaskStore(str(tmp_path)),
  313. FileSystemArtifactStore(str(tmp_path)),
  314. role_system_prompt_resolver=resolver,
  315. )
  316. result = await coordinator.executor.run_worker(_worker_context())
  317. assert result.status == "completed"
  318. assert runner.configs[0].system_prompt == "wired prompt"