test_orchestration_validation_policy.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461
  1. from __future__ import annotations
  2. import asyncio
  3. import pytest
  4. from agent.orchestration.config import OrchestrationConfig
  5. from agent.orchestration.evidence import (
  6. EvidenceBudgetExceeded,
  7. EvidenceOwnershipError,
  8. EvidenceProviderError,
  9. EvidenceProviderResult,
  10. EvidenceQuery,
  11. )
  12. from agent.orchestration.models import (
  13. AcceptanceCriterion,
  14. AgentRole,
  15. ArtifactRef,
  16. ArtifactSnapshot,
  17. AttemptSubmission,
  18. TaskAttempt,
  19. TaskRecord,
  20. TaskSpec,
  21. ValidationMode,
  22. ValidationPlan,
  23. ValidationVerdict,
  24. )
  25. from agent.orchestration.protocols import WorkerRunResult
  26. from agent.orchestration.validation_policy import (
  27. DefaultValidationPolicy,
  28. DeterministicRuleResult,
  29. RuleBasedDeterministicValidator,
  30. ValidationContext,
  31. )
  32. from test_coordinator_integration import FakeExecutor, create_task, make_coordinator
  33. def make_context(root: str = "root") -> ValidationContext:
  34. task = TaskRecord(
  35. task_id="task-1",
  36. goal_id=None,
  37. parent_task_id=None,
  38. display_path="1",
  39. specs=[TaskSpec(
  40. version=1,
  41. objective="validate",
  42. acceptance_criteria=[AcceptanceCriterion("criterion", "validate result")],
  43. )],
  44. )
  45. attempt = TaskAttempt(
  46. attempt_id="attempt-1",
  47. task_id=task.task_id,
  48. spec_version=1,
  49. worker_trace_id="worker-1",
  50. worker_preset="worker",
  51. execution_mode="new",
  52. snapshot_id="snapshot-1",
  53. )
  54. snapshot = ArtifactSnapshot(
  55. snapshot_id="snapshot-1",
  56. attempt_id=attempt.attempt_id,
  57. normalized_content={"summary": "done"},
  58. sha256="digest",
  59. artifact_refs=[],
  60. evidence_refs=[],
  61. )
  62. return ValidationContext(root, task, attempt, snapshot)
  63. class StaticRule:
  64. def __init__(self, rule_id, verdict):
  65. self.rule_id = rule_id
  66. self.verdict = verdict
  67. def evaluate(self, context):
  68. assert context.snapshot.snapshot_id
  69. return DeterministicRuleResult(self.rule_id, self.verdict, "checked")
  70. class AsyncRule(StaticRule):
  71. async def evaluate(self, context):
  72. await asyncio.sleep(0)
  73. return super().evaluate(context)
  74. class ErrorRule:
  75. rule_id = "error"
  76. def evaluate(self, context):
  77. raise RuntimeError("broken rule")
  78. class DeterministicPolicy:
  79. def __init__(self, verdict=ValidationVerdict.PASSED):
  80. self.verdict = verdict
  81. def plan(self, context):
  82. return ValidationPlan(
  83. mode=ValidationMode.DETERMINISTIC,
  84. validator_preset=None,
  85. rule_ids=[item.criterion_id for item in context.task.current_spec.acceptance_criteria],
  86. )
  87. class FixedPolicy:
  88. def __init__(self, plan):
  89. self.fixed_plan = plan
  90. def plan(self, context):
  91. return self.fixed_plan
  92. class InvalidBlankPolicy:
  93. def plan(self, context):
  94. plan = ValidationPlan()
  95. object.__setattr__(plan, "validator_preset", " ")
  96. return plan
  97. @pytest.mark.asyncio
  98. async def test_default_policy_and_validation_plan_roundtrip():
  99. plan = DefaultValidationPolicy("careful-validator").plan(make_context())
  100. assert plan.mode == ValidationMode.AGENT
  101. assert plan.validator_preset == "careful-validator"
  102. assert ValidationPlan.from_dict(plan.to_dict()) == plan
  103. deterministic = ValidationPlan(
  104. mode=ValidationMode.DETERMINISTIC,
  105. validator_preset=None,
  106. rule_ids=["schema", "checksum"],
  107. max_evidence_queries=2,
  108. )
  109. assert ValidationPlan.from_dict(deterministic.to_dict()) == deterministic
  110. assert ValidationPlan.from_dict(
  111. {"mode": "deterministic", "rule_ids": []}
  112. ).validator_preset is None
  113. with pytest.raises(ValueError, match="cannot contain"):
  114. ValidationPlan(rule_ids=["unexpected"])
  115. with pytest.raises(ValueError, match="cannot be blank"):
  116. ValidationPlan(validator_preset=" ")
  117. for kwargs in (
  118. {"max_evidence_queries": True},
  119. {"max_evidence_queries": 1.5},
  120. {"max_evidence_items_per_query": False},
  121. {"evidence_timeout_seconds": True},
  122. ):
  123. with pytest.raises(ValueError):
  124. ValidationPlan(**kwargs)
  125. with pytest.raises(ValueError, match="cannot select"):
  126. ValidationPlan(mode=ValidationMode.DETERMINISTIC)
  127. with pytest.raises(ValueError, match="unique"):
  128. ValidationPlan(
  129. mode=ValidationMode.DETERMINISTIC,
  130. validator_preset=None,
  131. rule_ids=["same", "same"],
  132. )
  133. @pytest.mark.asyncio
  134. async def test_deterministic_rules_aggregate_failures_and_errors():
  135. context = make_context()
  136. validator = RuleBasedDeterministicValidator(
  137. [
  138. StaticRule("pass", ValidationVerdict.PASSED),
  139. AsyncRule("fail", ValidationVerdict.FAILED),
  140. ErrorRule(),
  141. ]
  142. )
  143. result = await validator.validate(
  144. context,
  145. ValidationPlan(
  146. mode=ValidationMode.DETERMINISTIC,
  147. validator_preset=None,
  148. rule_ids=["pass", "missing", "error", "fail"],
  149. ),
  150. )
  151. assert result.verdict == ValidationVerdict.FAILED
  152. assert len(result.errors) == 2
  153. assert result.to_dict()["verdict"] == "failed"
  154. empty = await validator.validate(
  155. context,
  156. ValidationPlan(mode=ValidationMode.DETERMINISTIC, validator_preset=None),
  157. )
  158. assert empty.verdict == ValidationVerdict.INCONCLUSIVE
  159. with pytest.raises(ValueError, match="deterministic plan"):
  160. await validator.validate(context, ValidationPlan())
  161. with pytest.raises(ValueError, match="Duplicate"):
  162. RuleBasedDeterministicValidator(
  163. [
  164. StaticRule("same", ValidationVerdict.PASSED),
  165. StaticRule("same", ValidationVerdict.PASSED),
  166. ]
  167. )
  168. @pytest.mark.asyncio
  169. async def test_coordinator_freezes_default_agent_plan(tmp_path):
  170. executor = FakeExecutor([ValidationVerdict.PASSED])
  171. coordinator, store, _ = await make_coordinator(tmp_path, executor)
  172. task_id = await create_task(coordinator, "default agent policy")
  173. result = (await coordinator.dispatch_tasks("root", [task_id]))[0]
  174. report = (await store.load("root")).validations[result.validation_id]
  175. assert executor.validator_calls == 1
  176. assert report.validation_plan.mode == ValidationMode.AGENT
  177. assert report.validation_plan.validator_preset == "validator"
  178. assert ValidationPlan.from_dict(report.validation_plan) == report.validation_plan
  179. @pytest.mark.asyncio
  180. @pytest.mark.parametrize(
  181. "verdict", [ValidationVerdict.PASSED, ValidationVerdict.FAILED]
  182. )
  183. async def test_coordinator_executes_deterministic_plan_without_agent_validator(
  184. tmp_path, verdict
  185. ):
  186. executor = FakeExecutor([])
  187. coordinator, store, _ = await make_coordinator(tmp_path, executor)
  188. coordinator.validation_policy = DeterministicPolicy()
  189. coordinator.deterministic_validator = RuleBasedDeterministicValidator(
  190. [StaticRule("c1", verdict)]
  191. )
  192. task_id = await create_task(coordinator, "deterministic")
  193. result = (await coordinator.dispatch_tasks("root", [task_id]))[0]
  194. report = (await store.load("root")).validations[result.validation_id]
  195. assert executor.worker_calls == 1
  196. assert executor.validator_calls == 0
  197. assert report.validation_plan.mode == ValidationMode.DETERMINISTIC
  198. assert report.verdict == verdict
  199. assert report.criterion_results[0].criterion_id == "c1"
  200. @pytest.mark.asyncio
  201. @pytest.mark.parametrize("rule_ids", [[], ["unknown"]])
  202. async def test_deterministic_plan_must_cover_hard_criteria_without_unknowns(
  203. tmp_path, rule_ids
  204. ):
  205. executor = FakeExecutor([])
  206. coordinator, _, _ = await make_coordinator(tmp_path, executor)
  207. coordinator.validation_policy = FixedPolicy(
  208. ValidationPlan(
  209. mode=ValidationMode.DETERMINISTIC,
  210. validator_preset=None,
  211. rule_ids=rule_ids,
  212. )
  213. )
  214. coordinator.deterministic_validator = RuleBasedDeterministicValidator([])
  215. task_id = await create_task(coordinator, "invalid deterministic mapping")
  216. result = (await coordinator.dispatch_tasks("root", [task_id]))[0]
  217. assert "criteria" in (result.error or "")
  218. assert executor.validator_calls == 0
  219. @pytest.mark.asyncio
  220. async def test_coordinator_rejects_custom_policy_with_blank_agent_preset(tmp_path):
  221. executor = FakeExecutor([])
  222. coordinator, _, _ = await make_coordinator(tmp_path, executor)
  223. coordinator.validation_policy = InvalidBlankPolicy()
  224. task_id = await create_task(coordinator, "invalid preset")
  225. result = (await coordinator.dispatch_tasks("root", [task_id]))[0]
  226. assert "cannot be blank" in (result.error or "")
  227. assert executor.validator_calls == 0
  228. class EvidencePolicy:
  229. def __init__(self, max_queries=1, timeout=0.1):
  230. self.max_queries = max_queries
  231. self.timeout = timeout
  232. def plan(self, context):
  233. return ValidationPlan(
  234. max_evidence_queries=self.max_queries,
  235. max_evidence_items_per_query=1,
  236. evidence_timeout_seconds=self.timeout,
  237. )
  238. class EvidenceProvider:
  239. def __init__(self, *, delay=0, error=None, item_count=1, items=None):
  240. self.delay = delay
  241. self.error = error
  242. self.item_count = item_count
  243. self.items = items
  244. self.requests = []
  245. async def query(self, request):
  246. self.requests.append(request)
  247. if self.delay:
  248. await asyncio.sleep(self.delay)
  249. if self.error:
  250. raise self.error
  251. return EvidenceProviderResult(
  252. items=(
  253. self.items
  254. if self.items is not None
  255. else [{"index": index} for index in range(self.item_count)]
  256. ),
  257. evidence_refs=[ArtifactRef(uri="memory://evidence", version="1")],
  258. )
  259. class BlockingValidatorExecutor:
  260. def __init__(self):
  261. self.coordinator = None
  262. self.context_ready = asyncio.Event()
  263. self.context = None
  264. async def run_worker(self, context):
  265. await self.coordinator.submit_attempt(
  266. {
  267. **context,
  268. "role": AgentRole.WORKER.value,
  269. "trace_id": context["worker_trace_id"],
  270. "tool_call_id": f"submit-{context['attempt_id']}",
  271. },
  272. AttemptSubmission("done", [ArtifactRef(uri="memory://work", version="1")]),
  273. )
  274. return WorkerRunResult(context["worker_trace_id"], "completed")
  275. async def run_validator(self, context):
  276. self.context = context
  277. self.context_ready.set()
  278. await asyncio.Event().wait()
  279. async def stop(self, trace_id):
  280. return True
  281. async def running_evidence_validation(tmp_path, provider, *, max_queries=1, timeout=0.1):
  282. executor = BlockingValidatorExecutor()
  283. coordinator, store, _ = await make_coordinator(tmp_path, executor)
  284. coordinator.config = OrchestrationConfig(stop_grace_seconds=0)
  285. coordinator.validation_policy = EvidencePolicy(max_queries, timeout)
  286. coordinator.evidence_provider = provider
  287. task_id = await create_task(coordinator, "evidence")
  288. operation = await coordinator.start_operation("root", "dispatch", task_ids=[task_id])
  289. await asyncio.wait_for(executor.context_ready.wait(), timeout=1)
  290. context = executor.context
  291. actor = {
  292. **context,
  293. "role": AgentRole.VALIDATOR.value,
  294. "trace_id": context["validator_trace_id"],
  295. "tool_call_id": "evidence-1",
  296. }
  297. request = EvidenceQuery(
  298. root_trace_id="root",
  299. task_id=task_id,
  300. attempt_id=context["attempt_id"],
  301. snapshot_id=context["snapshot_id"],
  302. validation_id=context["validation_id"],
  303. query="facts",
  304. limit=100,
  305. )
  306. return coordinator, store, operation, actor, request
  307. @pytest.mark.asyncio
  308. async def test_coordinator_evidence_ownership_and_persistent_budget(tmp_path):
  309. provider = EvidenceProvider()
  310. coordinator, store, operation, actor, request = await running_evidence_validation(
  311. tmp_path, provider
  312. )
  313. response = await coordinator.query_evidence(actor, request)
  314. assert response.queries_used == 1
  315. assert response.queries_remaining == 0
  316. assert provider.requests[0].limit == 1
  317. report = (await store.load("root")).validations[request.validation_id]
  318. assert report.evidence_queries_used == 1
  319. assert next(iter(report.evidence_query_results.values()))["status"] == "completed"
  320. replay = await coordinator.query_evidence(actor, request)
  321. assert replay.queries_used == 1
  322. assert len(provider.requests) == 1
  323. with pytest.raises(EvidenceBudgetExceeded):
  324. await coordinator.query_evidence(
  325. {**actor, "tool_call_id": "evidence-2"},
  326. EvidenceQuery(**{**request.__dict__, "query": "more"}),
  327. )
  328. with pytest.raises(EvidenceOwnershipError):
  329. await coordinator.query_evidence(
  330. {**actor, "tool_call_id": "foreign"},
  331. EvidenceQuery(**{**request.__dict__, "snapshot_id": "foreign"}),
  332. )
  333. with pytest.raises(EvidenceOwnershipError):
  334. await coordinator.query_evidence(
  335. {**actor, "trace_id": "foreign", "tool_call_id": "foreign-trace"},
  336. request,
  337. )
  338. await coordinator.stop_operation("root", operation.operation_id)
  339. @pytest.mark.asyncio
  340. @pytest.mark.parametrize(
  341. ("provider", "message"),
  342. [
  343. (EvidenceProvider(delay=0.05), "timed out"),
  344. (EvidenceProvider(error=RuntimeError("offline")), "offline"),
  345. (EvidenceProvider(item_count=2), "more than"),
  346. (EvidenceProvider(items=[{"bad": object()}]), "non-JSON"),
  347. ],
  348. )
  349. async def test_evidence_provider_failures_are_validation_errors(
  350. tmp_path, provider, message
  351. ):
  352. timeout = 0.001 if provider.delay else 0.1
  353. coordinator, store, operation, actor, request = await running_evidence_validation(
  354. tmp_path,
  355. provider,
  356. timeout=timeout,
  357. )
  358. with pytest.raises(EvidenceProviderError, match=message):
  359. await coordinator.query_evidence(actor, request)
  360. assert (await store.load("root")).validations[request.validation_id].evidence_queries_used == 1
  361. with pytest.raises(EvidenceProviderError, match=message):
  362. await coordinator.query_evidence(actor, request)
  363. assert len(provider.requests) == 1
  364. await coordinator.stop_operation("root", operation.operation_id)
  365. @pytest.mark.asyncio
  366. async def test_concurrent_evidence_replay_calls_provider_once(tmp_path):
  367. provider = EvidenceProvider(delay=0.02)
  368. coordinator, _, operation, actor, request = await running_evidence_validation(
  369. tmp_path,
  370. provider,
  371. )
  372. responses = await asyncio.gather(
  373. coordinator.query_evidence(actor, request),
  374. coordinator.query_evidence(actor, request),
  375. )
  376. assert [item.queries_used for item in responses] == [1, 1]
  377. assert len(provider.requests) == 1
  378. await coordinator.stop_operation("root", operation.operation_id)
  379. def test_evidence_query_rejects_empty_scope_and_limit():
  380. values = {
  381. "root_trace_id": "root",
  382. "task_id": "task",
  383. "attempt_id": "attempt",
  384. "snapshot_id": "snapshot",
  385. "validation_id": "validation",
  386. "query": "facts",
  387. }
  388. with pytest.raises(ValueError, match="query"):
  389. EvidenceQuery(**{**values, "query": " "})
  390. with pytest.raises(ValueError, match="limit"):
  391. EvidenceQuery(**values, limit=0)
  392. with pytest.raises(ValueError, match="limit"):
  393. EvidenceQuery(**values, limit=True)
  394. with pytest.raises(ValueError, match="root_trace_id"):
  395. EvidenceQuery(**{**values, "root_trace_id": 123})