test_mission_service_boundaries.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563
  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. DirectionGoal,
  22. ScriptDirectionArtifactV1,
  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", "projected", "active"),
  157. [
  158. (PublicationState.PENDING, "# direction", None),
  159. (PublicationState.FAILED, "# direction", 3),
  160. (PublicationState.PUBLISHED, "# direction", 3),
  161. ],
  162. )
  163. async def test_direction_reconciler_reenters_each_publication_crash_point(
  164. publication_state: PublicationState,
  165. projected: str | None,
  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 = ScriptDirectionArtifactV1(
  175. goals=(DirectionGoal("goal", "statement"),),
  176. evidence_refs=("script-build://artifact-versions/2",),
  177. legacy_markdown="# direction",
  178. )
  179. version = SimpleNamespace(
  180. artifact_version_id=3,
  181. canonical_sha256=ref.digest,
  182. state=(
  183. ArtifactState.PUBLISHED
  184. if publication_state is PublicationState.PUBLISHED
  185. else ArtifactState.FROZEN
  186. ),
  187. artifact=direction,
  188. )
  189. task = SimpleNamespace(
  190. task_id="direction-task",
  191. current_spec=SimpleNamespace(context_refs=("script-build://task-kinds/direction",)),
  192. status=TaskStatus.COMPLETED,
  193. decision_ids=["accept"],
  194. )
  195. attempt = SimpleNamespace(
  196. attempt_id="attempt",
  197. snapshot_id="snapshot",
  198. submission=SimpleNamespace(artifact_refs=(ref,)),
  199. )
  200. ledger = SimpleNamespace(
  201. tasks={task.task_id: task},
  202. decisions={
  203. "accept": SimpleNamespace(
  204. action=DecisionAction.ACCEPT,
  205. decision_id="accept",
  206. attempt_id="attempt",
  207. validation_id="validation",
  208. )
  209. },
  210. attempts={"attempt": attempt},
  211. validations={
  212. "validation": SimpleNamespace(
  213. attempt_id="attempt",
  214. verdict=ValidationVerdict.PASSED,
  215. snapshot_id="snapshot",
  216. )
  217. },
  218. )
  219. class Bindings:
  220. current = active
  221. async def get_by_build(self, _build: int):
  222. return SimpleNamespace(active_direction_artifact_version_id=self.current)
  223. async def set_active_direction(self, *, artifact_version_id: int, **_kwargs):
  224. self.current = artifact_version_id
  225. class Publications:
  226. state = publication_state
  227. published_calls = 0
  228. async def prepare(self, **_kwargs):
  229. return SimpleNamespace(publication_id=1, state=self.state)
  230. async def mark_published(self, _publication_id: int):
  231. self.state = PublicationState.PUBLISHED
  232. self.published_calls += 1
  233. async def mark_failed(self, *_args, **_kwargs):
  234. self.state = PublicationState.FAILED
  235. class State(_StopState):
  236. direction = projected
  237. async def get_direction(self, _build: int):
  238. return self.direction
  239. async def project_direction(self, _build: int, value: str):
  240. self.direction = value
  241. bindings = Bindings()
  242. publications = Publications()
  243. state = State(BuildStatus.PARTIAL)
  244. reconciler = DirectionReconciler(
  245. coordinator=SimpleNamespace(
  246. task_store=SimpleNamespace(load=lambda _root: _async_value(ledger))
  247. ),
  248. bindings=bindings,
  249. artifacts=SimpleNamespace(read_by_ref=lambda *_args, **_kwargs: _async_value(version)),
  250. publications=publications,
  251. legacy_state=state,
  252. )
  253. assert await reconciler.reconcile(7, "root") == 3
  254. assert state.direction == "# direction"
  255. assert bindings.current == 3
  256. assert publications.state is PublicationState.PUBLISHED
  257. assert publications.published_calls == (
  258. 0 if publication_state is PublicationState.PUBLISHED else 1
  259. )
  260. @pytest.mark.asyncio
  261. async def test_direction_reconciler_fails_closed_on_projection_conflict() -> None:
  262. ref = ArtifactRef(
  263. "script-build://artifact-versions/3",
  264. ArtifactKind.DIRECTION.value,
  265. "3",
  266. "sha256:" + "a" * 64,
  267. )
  268. direction = ScriptDirectionArtifactV1(
  269. goals=(DirectionGoal("goal", "statement"),),
  270. evidence_refs=("script-build://artifact-versions/2",),
  271. legacy_markdown="# accepted direction",
  272. )
  273. task = SimpleNamespace(
  274. task_id="direction-task",
  275. current_spec=SimpleNamespace(context_refs=("script-build://task-kinds/direction",)),
  276. status=TaskStatus.COMPLETED,
  277. decision_ids=["accept"],
  278. )
  279. ledger = SimpleNamespace(
  280. tasks={task.task_id: task},
  281. decisions={
  282. "accept": SimpleNamespace(
  283. action=DecisionAction.ACCEPT,
  284. decision_id="accept",
  285. attempt_id="attempt",
  286. validation_id="validation",
  287. )
  288. },
  289. attempts={
  290. "attempt": SimpleNamespace(
  291. attempt_id="attempt",
  292. snapshot_id="snapshot",
  293. submission=SimpleNamespace(artifact_refs=(ref,)),
  294. )
  295. },
  296. validations={
  297. "validation": SimpleNamespace(
  298. attempt_id="attempt",
  299. verdict=ValidationVerdict.PASSED,
  300. snapshot_id="snapshot",
  301. )
  302. },
  303. )
  304. class Publications:
  305. failed = False
  306. async def prepare(self, **_kwargs):
  307. return SimpleNamespace(publication_id=1, state=PublicationState.PENDING)
  308. async def mark_failed(self, *_args, **_kwargs):
  309. self.failed = True
  310. publications = Publications()
  311. reconciler = DirectionReconciler(
  312. coordinator=SimpleNamespace(
  313. task_store=SimpleNamespace(load=lambda _root: _async_value(ledger))
  314. ),
  315. bindings=SimpleNamespace(),
  316. artifacts=SimpleNamespace(
  317. read_by_ref=lambda *_args, **_kwargs: _async_value(
  318. SimpleNamespace(
  319. artifact_version_id=3,
  320. canonical_sha256=ref.digest,
  321. state=ArtifactState.FROZEN,
  322. artifact=direction,
  323. )
  324. )
  325. ),
  326. publications=publications,
  327. legacy_state=SimpleNamespace(
  328. get_status=lambda _build: _async_value(BuildStatus.PARTIAL),
  329. get_direction=lambda _build: _async_value("# conflicting direction"),
  330. ),
  331. )
  332. with pytest.raises(DirectionProjectionConflict):
  333. await reconciler.reconcile(7, "root")
  334. assert publications.failed is True
  335. @pytest.mark.asyncio
  336. @pytest.mark.parametrize("status", [BuildStatus.FAILED, BuildStatus.SUCCESS])
  337. async def test_stop_does_not_rewrite_terminal_failed_or_success(status: BuildStatus) -> None:
  338. state = _StopState(status)
  339. service = ScriptMissionService(
  340. runner=SimpleNamespace(),
  341. coordinator=SimpleNamespace(),
  342. factory=SimpleNamespace(),
  343. input_snapshots=SimpleNamespace(),
  344. bindings=SimpleNamespace(),
  345. legacy_state=state,
  346. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  347. direction_reconciler=SimpleNamespace(),
  348. )
  349. result = await service.stop(7, Principal("owner"))
  350. assert result.status is status
  351. assert state.status is status
  352. @pytest.mark.asyncio
  353. async def test_stop_intent_and_direction_reconcile_are_serialized() -> None:
  354. gate = BuildTransitionGate()
  355. entered = asyncio.Event()
  356. release = asyncio.Event()
  357. state = _StopState(BuildStatus.RUNNING)
  358. class BlockingReconciler:
  359. transition_gate = gate
  360. async def reconcile(self, script_build_id: int, _root: str) -> int:
  361. async with gate.hold(script_build_id):
  362. entered.set()
  363. await release.wait()
  364. assert state.status is BuildStatus.RUNNING
  365. return 1
  366. runner = SimpleNamespace(
  367. stop=lambda _root: _async_value(True),
  368. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(None)),
  369. )
  370. coordinator = SimpleNamespace(
  371. task_store=SimpleNamespace(load=lambda _root: _async_value(SimpleNamespace(operations={})))
  372. )
  373. service = ScriptMissionService(
  374. runner=runner,
  375. coordinator=coordinator,
  376. factory=SimpleNamespace(),
  377. input_snapshots=SimpleNamespace(),
  378. bindings=SimpleNamespace(
  379. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  380. ),
  381. legacy_state=state,
  382. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  383. direction_reconciler=BlockingReconciler(), # type: ignore[arg-type]
  384. transition_gate=gate,
  385. )
  386. reconcile = asyncio.create_task(service.direction_reconciler.reconcile(7, "root"))
  387. await entered.wait()
  388. stop = asyncio.create_task(service.stop(7, Principal("owner")))
  389. await asyncio.sleep(0)
  390. assert state.status is BuildStatus.RUNNING
  391. release.set()
  392. assert await reconcile == 1
  393. assert (await stop).status is BuildStatus.STOPPED
  394. @pytest.mark.asyncio
  395. @pytest.mark.parametrize(
  396. "initial_status",
  397. [BuildStatus.RUNNING, BuildStatus.STOPPING, BuildStatus.PARTIAL],
  398. )
  399. async def test_stop_closes_all_nonterminal_states_when_no_execution_remains(
  400. initial_status: BuildStatus,
  401. ) -> None:
  402. state = _StopState(initial_status)
  403. service = ScriptMissionService(
  404. runner=SimpleNamespace(
  405. stop=lambda _root: _async_value(True),
  406. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(None)),
  407. ),
  408. coordinator=SimpleNamespace(
  409. task_store=SimpleNamespace(
  410. load=lambda _root: _async_value(SimpleNamespace(operations={}))
  411. )
  412. ),
  413. factory=SimpleNamespace(),
  414. input_snapshots=SimpleNamespace(),
  415. bindings=SimpleNamespace(
  416. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  417. ),
  418. legacy_state=state,
  419. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  420. direction_reconciler=SimpleNamespace(),
  421. )
  422. result = await service.stop(7, Principal("owner"))
  423. assert result.status is BuildStatus.STOPPED
  424. assert state.status is BuildStatus.STOPPED
  425. @pytest.mark.asyncio
  426. async def test_stop_timeout_keeps_intent_and_reports_active_operation_ids() -> None:
  427. state = _StopState(BuildStatus.RUNNING)
  428. operation = SimpleNamespace(
  429. operation_id="operation-1",
  430. status=OperationStatus.RUNNING,
  431. )
  432. service = ScriptMissionService(
  433. runner=SimpleNamespace(
  434. stop=lambda _root: _async_value(True),
  435. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(None)),
  436. ),
  437. coordinator=SimpleNamespace(
  438. task_store=SimpleNamespace(
  439. load=lambda _root: _async_value(
  440. SimpleNamespace(operations={"operation-1": operation})
  441. )
  442. ),
  443. stop_operation=lambda *_args, **_kwargs: _async_value(None),
  444. get_operation=lambda *_args, **_kwargs: _async_value(operation),
  445. ),
  446. factory=SimpleNamespace(),
  447. input_snapshots=SimpleNamespace(),
  448. bindings=SimpleNamespace(
  449. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  450. ),
  451. legacy_state=state,
  452. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  453. direction_reconciler=SimpleNamespace(),
  454. stop_timeout_seconds=0.01,
  455. )
  456. result = await service.stop(7, Principal("owner"))
  457. assert result.status is BuildStatus.STOPPING
  458. assert result.stopped_operation_ids == ("operation-1",)
  459. assert state.status is BuildStatus.STOPPING
  460. @pytest.mark.asyncio
  461. async def test_stop_with_missing_binding_records_terminal_failure() -> None:
  462. class RecordingState(_StopState):
  463. error_summary: str | None = None
  464. async def set_status(
  465. self,
  466. _build: int,
  467. status: BuildStatus,
  468. *,
  469. error_summary: str | None = None,
  470. ) -> None:
  471. self.status = status
  472. self.error_summary = error_summary
  473. state = RecordingState(BuildStatus.PARTIAL)
  474. async def missing_binding(_build: int):
  475. raise LookupError("missing")
  476. service = ScriptMissionService(
  477. runner=SimpleNamespace(),
  478. coordinator=SimpleNamespace(),
  479. factory=SimpleNamespace(),
  480. input_snapshots=SimpleNamespace(),
  481. bindings=SimpleNamespace(get_by_build=missing_binding),
  482. legacy_state=state,
  483. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  484. direction_reconciler=SimpleNamespace(),
  485. )
  486. with pytest.raises(LookupError, match="missing"):
  487. await service.stop(7, Principal("owner"))
  488. assert state.status is BuildStatus.FAILED
  489. assert state.error_summary == "MISSION_BINDING_MISSING"
  490. async def _async_value(value):
  491. return value