test_mission_service_boundaries.py 7.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213
  1. from __future__ import annotations
  2. import asyncio
  3. from types import SimpleNamespace
  4. import pytest
  5. from script_build_host.application.mission_service import (
  6. BuildTransitionGate,
  7. DirectionReconciler,
  8. ScriptMissionService,
  9. StartScriptBuildCommand,
  10. )
  11. from script_build_host.domain.errors import InputRelationMismatch, ProtocolViolation
  12. from script_build_host.domain.records import BuildStatus, Principal
  13. class _RejectingInputs:
  14. async def validate_source(self, **_: object) -> None:
  15. raise InputRelationMismatch()
  16. class _SourceAuthorizer:
  17. async def require_source_access(self, *_args: object, **_kwargs: object) -> None:
  18. return None
  19. async def require_access(self, *_args: object, **_kwargs: object) -> None:
  20. return None
  21. class _LegacyState:
  22. def __init__(self) -> None:
  23. self.created = 0
  24. async def create(self, **_: object) -> int:
  25. self.created += 1
  26. return 1
  27. @pytest.mark.asyncio
  28. async def test_relation_mismatch_produces_no_build_or_binding_write() -> None:
  29. legacy = _LegacyState()
  30. bindings = SimpleNamespace(created=0)
  31. service = ScriptMissionService(
  32. runner=SimpleNamespace(),
  33. coordinator=SimpleNamespace(),
  34. factory=SimpleNamespace(),
  35. input_snapshots=_RejectingInputs(), # type: ignore[arg-type]
  36. bindings=bindings,
  37. legacy_state=legacy, # type: ignore[arg-type]
  38. authorizer=_SourceAuthorizer(),
  39. direction_reconciler=SimpleNamespace(),
  40. )
  41. with pytest.raises(InputRelationMismatch):
  42. await service.start(
  43. StartScriptBuildCommand(
  44. execution_id=1,
  45. topic_build_id=2,
  46. topic_id=3,
  47. principal=Principal("tester"),
  48. )
  49. )
  50. assert legacy.created == 0
  51. assert bindings.created == 0
  52. @pytest.mark.asyncio
  53. async def test_source_authorization_runs_before_input_or_build_access() -> None:
  54. legacy = _LegacyState()
  55. inputs = SimpleNamespace(validated=False)
  56. class Denied:
  57. async def require_source_access(self, *_args: object, **_kwargs: object) -> None:
  58. raise PermissionError("source forbidden")
  59. service = ScriptMissionService(
  60. runner=SimpleNamespace(),
  61. coordinator=SimpleNamespace(),
  62. factory=SimpleNamespace(),
  63. input_snapshots=inputs,
  64. bindings=SimpleNamespace(),
  65. legacy_state=legacy, # type: ignore[arg-type]
  66. authorizer=Denied(), # type: ignore[arg-type]
  67. direction_reconciler=SimpleNamespace(),
  68. )
  69. with pytest.raises(PermissionError, match="source forbidden"):
  70. await service.start(StartScriptBuildCommand(1, 2, 3, Principal("intruder")))
  71. assert legacy.created == 0
  72. assert inputs.validated is False
  73. class _StopState:
  74. def __init__(self, status: BuildStatus) -> None:
  75. self.status = status
  76. async def get_status(self, _build: int) -> BuildStatus:
  77. return self.status
  78. async def set_status(self, _build: int, status: BuildStatus, **_: object) -> None:
  79. self.status = status
  80. @pytest.mark.asyncio
  81. async def test_stop_waits_for_planner_even_when_there_are_no_operations() -> None:
  82. state = _StopState(BuildStatus.RUNNING)
  83. pending: asyncio.Future[None] = asyncio.get_running_loop().create_future()
  84. trace = SimpleNamespace(status="running")
  85. runner = SimpleNamespace(
  86. stop=lambda _root: _async_value(True),
  87. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(trace)),
  88. )
  89. coordinator = SimpleNamespace(
  90. task_store=SimpleNamespace(load=lambda _root: _async_value(SimpleNamespace(operations={})))
  91. )
  92. service = ScriptMissionService(
  93. runner=runner,
  94. coordinator=coordinator,
  95. factory=SimpleNamespace(),
  96. input_snapshots=SimpleNamespace(),
  97. bindings=SimpleNamespace(
  98. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  99. ),
  100. legacy_state=state,
  101. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  102. direction_reconciler=SimpleNamespace(),
  103. stop_timeout_seconds=0.01,
  104. )
  105. service._runs[7] = pending # type: ignore[assignment]
  106. result = await service.stop(7, Principal("owner"))
  107. assert result.status == BuildStatus.STOPPING
  108. assert state.status == BuildStatus.STOPPING
  109. pending.cancel()
  110. @pytest.mark.asyncio
  111. async def test_stop_is_idempotent_and_publication_is_gated_by_durable_status() -> None:
  112. state = _StopState(BuildStatus.STOPPED)
  113. service = ScriptMissionService(
  114. runner=SimpleNamespace(),
  115. coordinator=SimpleNamespace(),
  116. factory=SimpleNamespace(),
  117. input_snapshots=SimpleNamespace(),
  118. bindings=SimpleNamespace(),
  119. legacy_state=state,
  120. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  121. direction_reconciler=SimpleNamespace(),
  122. )
  123. result = await service.stop(7, Principal("owner"))
  124. assert result.status == BuildStatus.STOPPED
  125. state.status = BuildStatus.STOPPING
  126. publications = SimpleNamespace(prepared=False)
  127. reconciler = DirectionReconciler(
  128. coordinator=SimpleNamespace(),
  129. bindings=SimpleNamespace(),
  130. artifacts=SimpleNamespace(),
  131. publications=publications,
  132. legacy_state=state,
  133. )
  134. with pytest.raises(ProtocolViolation, match="forbidden while stopping"):
  135. await reconciler.reconcile(7, "root")
  136. assert publications.prepared is False
  137. @pytest.mark.asyncio
  138. async def test_stop_intent_and_direction_reconcile_are_serialized() -> None:
  139. gate = BuildTransitionGate()
  140. entered = asyncio.Event()
  141. release = asyncio.Event()
  142. state = _StopState(BuildStatus.RUNNING)
  143. class BlockingReconciler:
  144. transition_gate = gate
  145. async def reconcile(self, script_build_id: int, _root: str) -> int:
  146. async with gate.hold(script_build_id):
  147. entered.set()
  148. await release.wait()
  149. assert state.status is BuildStatus.RUNNING
  150. return 1
  151. runner = SimpleNamespace(
  152. stop=lambda _root: _async_value(True),
  153. trace_store=SimpleNamespace(get_trace=lambda _root: _async_value(None)),
  154. )
  155. coordinator = SimpleNamespace(
  156. task_store=SimpleNamespace(load=lambda _root: _async_value(SimpleNamespace(operations={})))
  157. )
  158. service = ScriptMissionService(
  159. runner=runner,
  160. coordinator=coordinator,
  161. factory=SimpleNamespace(),
  162. input_snapshots=SimpleNamespace(),
  163. bindings=SimpleNamespace(
  164. get_by_build=lambda _build: _async_value(SimpleNamespace(root_trace_id="root"))
  165. ),
  166. legacy_state=state,
  167. authorizer=SimpleNamespace(require_access=lambda _principal, _build: _async_value(None)),
  168. direction_reconciler=BlockingReconciler(), # type: ignore[arg-type]
  169. transition_gate=gate,
  170. )
  171. reconcile = asyncio.create_task(service.direction_reconciler.reconcile(7, "root"))
  172. await entered.wait()
  173. stop = asyncio.create_task(service.stop(7, Principal("owner")))
  174. await asyncio.sleep(0)
  175. assert state.status is BuildStatus.RUNNING
  176. release.set()
  177. assert await reconcile == 1
  178. assert (await stop).status is BuildStatus.STOPPED
  179. async def _async_value(value):
  180. return value