test_mission_service_boundaries.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570
  1. from __future__ import annotations
  2. import asyncio
  3. from types import SimpleNamespace
  4. import pytest
  5. from agent.orchestration import (
  6. ArtifactRef,
  7. DecisionAction,
  8. OperationStatus,
  9. TaskStatus,
  10. ValidationVerdict,
  11. )
  12. from script_build_host.application.mission_service import (
  13. BuildTransitionGate,
  14. DirectionReconciler,
  15. ScriptMissionService,
  16. StartScriptBuildCommand,
  17. )
  18. from script_build_host.domain.artifacts import (
  19. ArtifactKind,
  20. ArtifactState,
  21. DirectionArtifact,
  22. DirectionGoal,
  23. )
  24. from script_build_host.domain.errors import (
  25. DirectionProjectionConflict,
  26. InputRelationMismatch,
  27. ProtocolViolation,
  28. )
  29. from script_build_host.domain.records import BuildStatus, Principal, PublicationState
  30. class _RejectingInputs:
  31. async def validate_source(self, **_: object) -> None:
  32. raise InputRelationMismatch()
  33. class _SourceAuthorizer:
  34. async def require_source_access(self, *_args: object, **_kwargs: object) -> None:
  35. return None
  36. async def require_access(self, *_args: object, **_kwargs: object) -> None:
  37. return None
  38. class _LegacyState:
  39. def __init__(self) -> None:
  40. self.created = 0
  41. async def create(self, **_: object) -> int:
  42. self.created += 1
  43. return 1
  44. @pytest.mark.asyncio
  45. async def test_relation_mismatch_produces_no_build_or_binding_write() -> None:
  46. legacy = _LegacyState()
  47. bindings = SimpleNamespace(created=0)
  48. service = ScriptMissionService(
  49. runner=SimpleNamespace(),
  50. coordinator=SimpleNamespace(),
  51. factory=SimpleNamespace(),
  52. input_snapshots=_RejectingInputs(), # type: ignore[arg-type]
  53. bindings=bindings,
  54. legacy_state=legacy, # type: ignore[arg-type]
  55. authorizer=_SourceAuthorizer(),
  56. direction_reconciler=SimpleNamespace(),
  57. )
  58. with pytest.raises(InputRelationMismatch):
  59. await service.start(
  60. StartScriptBuildCommand(
  61. execution_id=1,
  62. topic_build_id=2,
  63. topic_id=3,
  64. principal=Principal("tester"),
  65. )
  66. )
  67. assert legacy.created == 0
  68. assert bindings.created == 0
  69. @pytest.mark.asyncio
  70. async def test_source_authorization_runs_before_input_or_build_access() -> None:
  71. legacy = _LegacyState()
  72. inputs = SimpleNamespace(validated=False)
  73. class Denied:
  74. async def require_source_access(self, *_args: object, **_kwargs: object) -> None:
  75. raise PermissionError("source forbidden")
  76. service = ScriptMissionService(
  77. runner=SimpleNamespace(),
  78. coordinator=SimpleNamespace(),
  79. factory=SimpleNamespace(),
  80. input_snapshots=inputs,
  81. bindings=SimpleNamespace(),
  82. legacy_state=legacy, # type: ignore[arg-type]
  83. authorizer=Denied(), # type: ignore[arg-type]
  84. direction_reconciler=SimpleNamespace(),
  85. )
  86. with pytest.raises(PermissionError, match="source forbidden"):
  87. await service.start(StartScriptBuildCommand(1, 2, 3, Principal("intruder")))
  88. assert legacy.created == 0
  89. assert inputs.validated is False
  90. class _StopState:
  91. def __init__(self, status: BuildStatus) -> None:
  92. self.status = status
  93. async def get_status(self, _build: int) -> BuildStatus:
  94. return self.status
  95. async def set_status(self, _build: int, status: BuildStatus, **_: object) -> None:
  96. self.status = status
  97. @pytest.mark.asyncio
  98. async def test_stop_waits_for_planner_even_when_there_are_no_operations() -> None:
  99. state = _StopState(BuildStatus.RUNNING)
  100. pending: asyncio.Future[None] = asyncio.get_running_loop().create_future()
  101. trace = SimpleNamespace(status="running")
  102. runner = SimpleNamespace(
  103. stop=lambda _root: _async_value(True),
  104. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(trace)),
  105. )
  106. coordinator = SimpleNamespace(
  107. task_store=SimpleNamespace(load=lambda _root: _async_value(SimpleNamespace(operations={})))
  108. )
  109. service = ScriptMissionService(
  110. runner=runner,
  111. coordinator=coordinator,
  112. factory=SimpleNamespace(),
  113. input_snapshots=SimpleNamespace(),
  114. bindings=SimpleNamespace(
  115. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  116. ),
  117. legacy_state=state,
  118. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  119. direction_reconciler=SimpleNamespace(),
  120. stop_timeout_seconds=0.01,
  121. )
  122. service._runs[7] = pending # type: ignore[assignment]
  123. result = await service.stop(7, Principal("owner"))
  124. assert result.status == BuildStatus.STOPPING
  125. assert state.status == BuildStatus.STOPPING
  126. pending.cancel()
  127. @pytest.mark.asyncio
  128. async def test_stop_is_idempotent_and_publication_is_gated_by_durable_status() -> None:
  129. state = _StopState(BuildStatus.STOPPED)
  130. service = ScriptMissionService(
  131. runner=SimpleNamespace(),
  132. coordinator=SimpleNamespace(),
  133. factory=SimpleNamespace(),
  134. input_snapshots=SimpleNamespace(),
  135. bindings=SimpleNamespace(),
  136. legacy_state=state,
  137. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  138. direction_reconciler=SimpleNamespace(),
  139. )
  140. result = await service.stop(7, Principal("owner"))
  141. assert result.status == BuildStatus.STOPPED
  142. state.status = BuildStatus.STOPPING
  143. publications = SimpleNamespace(prepared=False)
  144. reconciler = DirectionReconciler(
  145. coordinator=SimpleNamespace(),
  146. bindings=SimpleNamespace(),
  147. artifacts=SimpleNamespace(),
  148. publications=publications,
  149. legacy_state=state,
  150. )
  151. with pytest.raises(ProtocolViolation, match="forbidden while stopping"):
  152. await reconciler.reconcile(7, "root")
  153. assert publications.prepared is False
  154. @pytest.mark.asyncio
  155. @pytest.mark.parametrize(
  156. ("publication_state", "has_projection", "active"),
  157. [
  158. (PublicationState.PENDING, True, None),
  159. (PublicationState.FAILED, True, 3),
  160. (PublicationState.PUBLISHED, True, 3),
  161. ],
  162. )
  163. async def test_direction_reconciler_reenters_each_publication_crash_point(
  164. publication_state: PublicationState,
  165. has_projection: bool,
  166. active: int | None,
  167. ) -> None:
  168. ref = ArtifactRef(
  169. "script-build://artifact-versions/3",
  170. ArtifactKind.DIRECTION.value,
  171. "3",
  172. "sha256:" + "a" * 64,
  173. )
  174. direction = DirectionArtifact(
  175. goals=(
  176. DirectionGoal(
  177. "goal", "statement", "the statement is testable", success_criteria=("done",)
  178. ),
  179. ),
  180. evidence_refs=("script-build://artifact-versions/2",),
  181. )
  182. projected = direction.legacy_markdown if has_projection else None
  183. version = SimpleNamespace(
  184. artifact_version_id=3,
  185. canonical_sha256=ref.digest,
  186. state=(
  187. ArtifactState.PUBLISHED
  188. if publication_state is PublicationState.PUBLISHED
  189. else ArtifactState.FROZEN
  190. ),
  191. artifact=direction,
  192. )
  193. task = SimpleNamespace(
  194. task_id="direction-task",
  195. current_spec=SimpleNamespace(context_refs=("script-build://task-kinds/direction",)),
  196. status=TaskStatus.COMPLETED,
  197. decision_ids=["accept"],
  198. )
  199. attempt = SimpleNamespace(
  200. attempt_id="attempt",
  201. snapshot_id="snapshot",
  202. submission=SimpleNamespace(artifact_refs=(ref,)),
  203. )
  204. ledger = SimpleNamespace(
  205. tasks={task.task_id: task},
  206. decisions={
  207. "accept": SimpleNamespace(
  208. action=DecisionAction.ACCEPT,
  209. decision_id="accept",
  210. attempt_id="attempt",
  211. validation_id="validation",
  212. )
  213. },
  214. attempts={"attempt": attempt},
  215. validations={
  216. "validation": SimpleNamespace(
  217. attempt_id="attempt",
  218. verdict=ValidationVerdict.PASSED,
  219. snapshot_id="snapshot",
  220. )
  221. },
  222. )
  223. class Bindings:
  224. current = active
  225. async def get_by_build(self, _build: int):
  226. return SimpleNamespace(active_direction_artifact_version_id=self.current)
  227. async def set_active_direction(self, *, artifact_version_id: int, **_kwargs):
  228. self.current = artifact_version_id
  229. class Publications:
  230. state = publication_state
  231. published_calls = 0
  232. async def prepare(self, **_kwargs):
  233. return SimpleNamespace(publication_id=1, state=self.state)
  234. async def mark_published(self, _publication_id: int):
  235. self.state = PublicationState.PUBLISHED
  236. self.published_calls += 1
  237. async def mark_failed(self, *_args, **_kwargs):
  238. self.state = PublicationState.FAILED
  239. class State(_StopState):
  240. direction = projected
  241. async def get_direction(self, _build: int):
  242. return self.direction
  243. async def project_direction(self, _build: int, value: str):
  244. self.direction = value
  245. bindings = Bindings()
  246. publications = Publications()
  247. state = State(BuildStatus.PARTIAL)
  248. reconciler = DirectionReconciler(
  249. coordinator=SimpleNamespace(
  250. task_store=SimpleNamespace(load=lambda _root: _async_value(ledger))
  251. ),
  252. bindings=bindings,
  253. artifacts=SimpleNamespace(read_by_ref=lambda *_args, **_kwargs: _async_value(version)),
  254. publications=publications,
  255. legacy_state=state,
  256. )
  257. assert await reconciler.reconcile(7, "root") == 3
  258. assert state.direction == direction.legacy_markdown
  259. assert bindings.current == 3
  260. assert publications.state is PublicationState.PUBLISHED
  261. assert publications.published_calls == (
  262. 0 if publication_state is PublicationState.PUBLISHED else 1
  263. )
  264. @pytest.mark.asyncio
  265. async def test_direction_reconciler_fails_closed_on_projection_conflict() -> None:
  266. ref = ArtifactRef(
  267. "script-build://artifact-versions/3",
  268. ArtifactKind.DIRECTION.value,
  269. "3",
  270. "sha256:" + "a" * 64,
  271. )
  272. direction = DirectionArtifact(
  273. goals=(
  274. DirectionGoal(
  275. "goal", "statement", "the statement is testable", success_criteria=("done",)
  276. ),
  277. ),
  278. evidence_refs=("script-build://artifact-versions/2",),
  279. )
  280. task = SimpleNamespace(
  281. task_id="direction-task",
  282. current_spec=SimpleNamespace(context_refs=("script-build://task-kinds/direction",)),
  283. status=TaskStatus.COMPLETED,
  284. decision_ids=["accept"],
  285. )
  286. ledger = SimpleNamespace(
  287. tasks={task.task_id: task},
  288. decisions={
  289. "accept": SimpleNamespace(
  290. action=DecisionAction.ACCEPT,
  291. decision_id="accept",
  292. attempt_id="attempt",
  293. validation_id="validation",
  294. )
  295. },
  296. attempts={
  297. "attempt": SimpleNamespace(
  298. attempt_id="attempt",
  299. snapshot_id="snapshot",
  300. submission=SimpleNamespace(artifact_refs=(ref,)),
  301. )
  302. },
  303. validations={
  304. "validation": SimpleNamespace(
  305. attempt_id="attempt",
  306. verdict=ValidationVerdict.PASSED,
  307. snapshot_id="snapshot",
  308. )
  309. },
  310. )
  311. class Publications:
  312. failed = False
  313. async def prepare(self, **_kwargs):
  314. return SimpleNamespace(publication_id=1, state=PublicationState.PENDING)
  315. async def mark_failed(self, *_args, **_kwargs):
  316. self.failed = True
  317. publications = Publications()
  318. reconciler = DirectionReconciler(
  319. coordinator=SimpleNamespace(
  320. task_store=SimpleNamespace(load=lambda _root: _async_value(ledger))
  321. ),
  322. bindings=SimpleNamespace(),
  323. artifacts=SimpleNamespace(
  324. read_by_ref=lambda *_args, **_kwargs: _async_value(
  325. SimpleNamespace(
  326. artifact_version_id=3,
  327. canonical_sha256=ref.digest,
  328. state=ArtifactState.FROZEN,
  329. artifact=direction,
  330. )
  331. )
  332. ),
  333. publications=publications,
  334. legacy_state=SimpleNamespace(
  335. get_status=lambda _build: _async_value(BuildStatus.PARTIAL),
  336. get_direction=lambda _build: _async_value("# conflicting direction"),
  337. ),
  338. )
  339. with pytest.raises(DirectionProjectionConflict):
  340. await reconciler.reconcile(7, "root")
  341. assert publications.failed is True
  342. @pytest.mark.asyncio
  343. @pytest.mark.parametrize("status", [BuildStatus.FAILED, BuildStatus.SUCCESS])
  344. async def test_stop_does_not_rewrite_terminal_failed_or_success(status: BuildStatus) -> None:
  345. state = _StopState(status)
  346. service = ScriptMissionService(
  347. runner=SimpleNamespace(),
  348. coordinator=SimpleNamespace(),
  349. factory=SimpleNamespace(),
  350. input_snapshots=SimpleNamespace(),
  351. bindings=SimpleNamespace(),
  352. legacy_state=state,
  353. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  354. direction_reconciler=SimpleNamespace(),
  355. )
  356. result = await service.stop(7, Principal("owner"))
  357. assert result.status is status
  358. assert state.status is status
  359. @pytest.mark.asyncio
  360. async def test_stop_intent_and_direction_reconcile_are_serialized() -> None:
  361. gate = BuildTransitionGate()
  362. entered = asyncio.Event()
  363. release = asyncio.Event()
  364. state = _StopState(BuildStatus.RUNNING)
  365. class BlockingReconciler:
  366. transition_gate = gate
  367. async def reconcile(self, script_build_id: int, _root: str) -> int:
  368. async with gate.hold(script_build_id):
  369. entered.set()
  370. await release.wait()
  371. assert state.status is BuildStatus.RUNNING
  372. return 1
  373. runner = SimpleNamespace(
  374. stop=lambda _root: _async_value(True),
  375. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(None)),
  376. )
  377. coordinator = SimpleNamespace(
  378. task_store=SimpleNamespace(load=lambda _root: _async_value(SimpleNamespace(operations={})))
  379. )
  380. service = ScriptMissionService(
  381. runner=runner,
  382. coordinator=coordinator,
  383. factory=SimpleNamespace(),
  384. input_snapshots=SimpleNamespace(),
  385. bindings=SimpleNamespace(
  386. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  387. ),
  388. legacy_state=state,
  389. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  390. direction_reconciler=BlockingReconciler(), # type: ignore[arg-type]
  391. transition_gate=gate,
  392. )
  393. reconcile = asyncio.create_task(service.direction_reconciler.reconcile(7, "root"))
  394. await entered.wait()
  395. stop = asyncio.create_task(service.stop(7, Principal("owner")))
  396. await asyncio.sleep(0)
  397. assert state.status is BuildStatus.RUNNING
  398. release.set()
  399. assert await reconcile == 1
  400. assert (await stop).status is BuildStatus.STOPPED
  401. @pytest.mark.asyncio
  402. @pytest.mark.parametrize(
  403. "initial_status",
  404. [BuildStatus.RUNNING, BuildStatus.STOPPING, BuildStatus.PARTIAL],
  405. )
  406. async def test_stop_closes_all_nonterminal_states_when_no_execution_remains(
  407. initial_status: BuildStatus,
  408. ) -> None:
  409. state = _StopState(initial_status)
  410. service = ScriptMissionService(
  411. runner=SimpleNamespace(
  412. stop=lambda _root: _async_value(True),
  413. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(None)),
  414. ),
  415. coordinator=SimpleNamespace(
  416. task_store=SimpleNamespace(
  417. load=lambda _root: _async_value(SimpleNamespace(operations={}))
  418. )
  419. ),
  420. factory=SimpleNamespace(),
  421. input_snapshots=SimpleNamespace(),
  422. bindings=SimpleNamespace(
  423. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  424. ),
  425. legacy_state=state,
  426. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  427. direction_reconciler=SimpleNamespace(),
  428. )
  429. result = await service.stop(7, Principal("owner"))
  430. assert result.status is BuildStatus.STOPPED
  431. assert state.status is BuildStatus.STOPPED
  432. @pytest.mark.asyncio
  433. async def test_stop_timeout_keeps_intent_and_reports_active_operation_ids() -> None:
  434. state = _StopState(BuildStatus.RUNNING)
  435. operation = SimpleNamespace(
  436. operation_id="operation-1",
  437. status=OperationStatus.RUNNING,
  438. )
  439. service = ScriptMissionService(
  440. runner=SimpleNamespace(
  441. stop=lambda _root: _async_value(True),
  442. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(None)),
  443. ),
  444. coordinator=SimpleNamespace(
  445. task_store=SimpleNamespace(
  446. load=lambda _root: _async_value(
  447. SimpleNamespace(operations={"operation-1": operation})
  448. )
  449. ),
  450. stop_operation=lambda *_args, **_kwargs: _async_value(None),
  451. get_operation=lambda *_args, **_kwargs: _async_value(operation),
  452. ),
  453. factory=SimpleNamespace(),
  454. input_snapshots=SimpleNamespace(),
  455. bindings=SimpleNamespace(
  456. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  457. ),
  458. legacy_state=state,
  459. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  460. direction_reconciler=SimpleNamespace(),
  461. stop_timeout_seconds=0.01,
  462. )
  463. result = await service.stop(7, Principal("owner"))
  464. assert result.status is BuildStatus.STOPPING
  465. assert result.stopped_operation_ids == ("operation-1",)
  466. assert state.status is BuildStatus.STOPPING
  467. @pytest.mark.asyncio
  468. async def test_stop_with_missing_binding_records_terminal_failure() -> None:
  469. class RecordingState(_StopState):
  470. error_summary: str | None = None
  471. async def set_status(
  472. self,
  473. _build: int,
  474. status: BuildStatus,
  475. *,
  476. error_summary: str | None = None,
  477. ) -> None:
  478. self.status = status
  479. self.error_summary = error_summary
  480. state = RecordingState(BuildStatus.PARTIAL)
  481. async def missing_binding(_build: int):
  482. raise LookupError("missing")
  483. service = ScriptMissionService(
  484. runner=SimpleNamespace(),
  485. coordinator=SimpleNamespace(),
  486. factory=SimpleNamespace(),
  487. input_snapshots=SimpleNamespace(),
  488. bindings=SimpleNamespace(get_by_build=missing_binding),
  489. legacy_state=state,
  490. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  491. direction_reconciler=SimpleNamespace(),
  492. )
  493. with pytest.raises(LookupError, match="missing"):
  494. await service.stop(7, Principal("owner"))
  495. assert state.status is BuildStatus.FAILED
  496. assert state.error_summary == "MISSION_BINDING_MISSING"
  497. async def _async_value(value):
  498. return value