test_orchestration_validation_policy.py 21 KB

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