ownership.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361
  1. """Single-host mission ownership and database fencing.
  2. The file descriptor lock remains held for the lifetime of an owner. Database
  3. epochs make every later command independently reject a stale process.
  4. """
  5. from __future__ import annotations
  6. import fcntl
  7. import os
  8. from collections.abc import Awaitable, Callable
  9. from contextlib import AbstractAsyncContextManager
  10. from datetime import UTC, datetime
  11. from pathlib import Path
  12. from types import TracebackType
  13. from typing import Any, cast
  14. from uuid import uuid4
  15. from sqlalchemy import select, update
  16. from sqlalchemy.engine import CursorResult
  17. from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
  18. from script_build_host.domain.errors import (
  19. BuildNotFound,
  20. MissionAlreadyOwned,
  21. MissionFencingTokenStale,
  22. MissionStopRequested,
  23. )
  24. from script_build_host.domain.records import BuildStatus, MissionOwnerToken
  25. from script_build_host.infrastructure.legacy_tables import (
  26. script_build_record,
  27. script_build_runtime_record,
  28. )
  29. from script_build_host.infrastructure.tables import mission_binding_table
  30. class OwnerLease(AbstractAsyncContextManager[MissionOwnerToken]):
  31. def __init__(
  32. self,
  33. sessions: async_sessionmaker[AsyncSession],
  34. lock_root: Path,
  35. *,
  36. owner_instance_id: str | None = None,
  37. ) -> None:
  38. self._sessions = sessions
  39. self._lock_root = lock_root.resolve()
  40. self._owner_instance_id = owner_instance_id or str(uuid4())
  41. self._fd: int | None = None
  42. self._token: MissionOwnerToken | None = None
  43. async def acquire(self, script_build_id: int) -> MissionOwnerToken:
  44. if self._fd is not None:
  45. if self._token is None or self._token.script_build_id != script_build_id:
  46. raise MissionAlreadyOwned()
  47. return self._token
  48. self._lock_root.mkdir(parents=True, exist_ok=True)
  49. path = self._lock_root / f"script-build-{script_build_id}.owner.lock"
  50. fd = os.open(path, os.O_RDWR | os.O_CREAT, 0o600)
  51. try:
  52. fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
  53. except BlockingIOError as exc:
  54. os.close(fd)
  55. raise MissionAlreadyOwned() from exc
  56. try:
  57. async with self._sessions() as session, session.begin():
  58. row = (
  59. (
  60. await session.execute(
  61. select(mission_binding_table)
  62. .where(mission_binding_table.c.script_build_id == script_build_id)
  63. .with_for_update()
  64. )
  65. )
  66. .mappings()
  67. .one_or_none()
  68. )
  69. if row is None:
  70. raise BuildNotFound()
  71. epoch = int(row["owner_epoch"] or 0) + 1
  72. now = datetime.now(UTC)
  73. await session.execute(
  74. update(mission_binding_table)
  75. .where(mission_binding_table.c.id == row["id"])
  76. .values(
  77. owner_epoch=epoch,
  78. owner_instance_id=self._owner_instance_id,
  79. owner_acquired_at=now,
  80. updated_at=now,
  81. )
  82. )
  83. token = MissionOwnerToken(
  84. script_build_id=script_build_id,
  85. root_trace_id=str(row["root_trace_id"]),
  86. owner_epoch=epoch,
  87. stop_epoch=int(row["stop_epoch"] or 0),
  88. owner_instance_id=self._owner_instance_id,
  89. )
  90. except BaseException:
  91. fcntl.flock(fd, fcntl.LOCK_UN)
  92. os.close(fd)
  93. raise
  94. self._fd = fd
  95. self._token = token
  96. return token
  97. async def release(self) -> None:
  98. fd, token = self._fd, self._token
  99. if fd is None:
  100. return
  101. try:
  102. if token is not None:
  103. async with self._sessions() as session, session.begin():
  104. await session.execute(
  105. update(mission_binding_table)
  106. .where(
  107. mission_binding_table.c.script_build_id == token.script_build_id,
  108. mission_binding_table.c.owner_epoch == token.owner_epoch,
  109. mission_binding_table.c.owner_instance_id == token.owner_instance_id,
  110. )
  111. .values(
  112. owner_instance_id=None,
  113. owner_acquired_at=None,
  114. updated_at=datetime.now(UTC),
  115. )
  116. )
  117. finally:
  118. fcntl.flock(fd, fcntl.LOCK_UN)
  119. os.close(fd)
  120. self._fd = None
  121. self._token = None
  122. async def __aenter__(self) -> MissionOwnerToken:
  123. if self._token is None:
  124. raise RuntimeError("OwnerLease.acquire() must be called before entering the lease")
  125. return self._token
  126. async def __aexit__(
  127. self,
  128. exc_type: type[BaseException] | None,
  129. exc: BaseException | None,
  130. traceback: TracebackType | None,
  131. ) -> None:
  132. await self.release()
  133. class FencedCommandGate:
  134. """Validate one owner token while holding binding/build row locks."""
  135. def __init__(self, sessions: async_sessionmaker[AsyncSession]) -> None:
  136. self._sessions = sessions
  137. bind = sessions.kw.get("bind")
  138. self._runtime_record = (
  139. script_build_runtime_record
  140. if getattr(getattr(bind, "dialect", None), "name", None) == "mysql"
  141. else script_build_record
  142. )
  143. async def verify(self, token: MissionOwnerToken) -> None:
  144. async with self._sessions() as session, session.begin():
  145. await self.verify_in_session(session, token)
  146. async def execute(
  147. self,
  148. token: MissionOwnerToken,
  149. mutation: Callable[[AsyncSession], Awaitable[Any]],
  150. ) -> Any:
  151. async with self._sessions() as session, session.begin():
  152. await self.verify_in_session(session, token)
  153. return await mutation(session)
  154. async def compare_and_set_input_snapshot(
  155. self,
  156. token: MissionOwnerToken,
  157. *,
  158. expected_snapshot_id: int,
  159. new_snapshot_id: int,
  160. ) -> None:
  161. """Move the binding to one immutable snapshot under the owner fence."""
  162. async with self._sessions() as session, session.begin():
  163. await self.verify_in_session(session, token)
  164. result = cast(
  165. CursorResult[Any],
  166. await session.execute(
  167. update(mission_binding_table)
  168. .where(
  169. mission_binding_table.c.script_build_id == token.script_build_id,
  170. mission_binding_table.c.input_snapshot_id == expected_snapshot_id,
  171. )
  172. .values(input_snapshot_id=new_snapshot_id, updated_at=datetime.now(UTC))
  173. ),
  174. )
  175. if result.rowcount != 1:
  176. current = await session.scalar(
  177. select(mission_binding_table.c.input_snapshot_id).where(
  178. mission_binding_table.c.script_build_id == token.script_build_id
  179. )
  180. )
  181. if current != new_snapshot_id:
  182. raise MissionFencingTokenStale()
  183. async def begin_phase(self, token: MissionOwnerToken) -> None:
  184. """Project the legacy build to running without a stop-overwrite window."""
  185. await self._update_runtime_state(
  186. token,
  187. expected_statuses=(BuildStatus.PARTIAL.value, BuildStatus.RUNNING.value),
  188. status=BuildStatus.RUNNING.value,
  189. end_time=None,
  190. error_message=None,
  191. summary=None,
  192. )
  193. async def set_checkpoint(
  194. self,
  195. token: MissionOwnerToken,
  196. *,
  197. checkpoint_code: str,
  198. summary: str,
  199. ) -> None:
  200. """Persist the Phase 3 checkpoint in the same transaction as fencing."""
  201. await self._update_runtime_state(
  202. token,
  203. expected_statuses=(BuildStatus.RUNNING.value,),
  204. idempotent_checkpoint=checkpoint_code[:255],
  205. status=BuildStatus.PARTIAL.value,
  206. error_message=checkpoint_code[:255],
  207. summary=summary[:2000],
  208. reson_trace_id=token.root_trace_id[:200],
  209. end_time=datetime.now(UTC),
  210. )
  211. async def _update_runtime_state(
  212. self,
  213. token: MissionOwnerToken,
  214. *,
  215. expected_statuses: tuple[str, ...],
  216. idempotent_checkpoint: str | None = None,
  217. **values: Any,
  218. ) -> None:
  219. async with self._sessions() as session, session.begin():
  220. await self.verify_in_session(session, token)
  221. result = cast(
  222. CursorResult[Any],
  223. await session.execute(
  224. update(self._runtime_record)
  225. .where(
  226. self._runtime_record.c.id == token.script_build_id,
  227. self._runtime_record.c.is_deleted.is_(False),
  228. self._runtime_record.c.status.in_(expected_statuses),
  229. )
  230. .values(**values)
  231. ),
  232. )
  233. if result.rowcount != 1:
  234. current = (
  235. (
  236. await session.execute(
  237. select(
  238. script_build_record.c.status,
  239. script_build_record.c.error_message,
  240. ).where(script_build_record.c.id == token.script_build_id)
  241. )
  242. )
  243. .mappings()
  244. .one_or_none()
  245. )
  246. if current is None:
  247. raise BuildNotFound()
  248. if (
  249. idempotent_checkpoint is not None
  250. and current["status"] == BuildStatus.PARTIAL.value
  251. and current["error_message"] == idempotent_checkpoint
  252. ):
  253. return
  254. raise MissionFencingTokenStale()
  255. async def verify_in_session(
  256. self,
  257. session: AsyncSession,
  258. token: MissionOwnerToken,
  259. *,
  260. allow_success: bool = False,
  261. ) -> None:
  262. binding = (
  263. (
  264. await session.execute(
  265. select(mission_binding_table)
  266. .where(mission_binding_table.c.script_build_id == token.script_build_id)
  267. .with_for_update()
  268. )
  269. )
  270. .mappings()
  271. .one_or_none()
  272. )
  273. if binding is None:
  274. raise BuildNotFound()
  275. if (
  276. str(binding["root_trace_id"]) != token.root_trace_id
  277. or int(binding["owner_epoch"] or 0) != token.owner_epoch
  278. or binding["owner_instance_id"] != token.owner_instance_id
  279. ):
  280. raise MissionFencingTokenStale()
  281. if int(binding["stop_epoch"] or 0) != token.stop_epoch:
  282. raise MissionStopRequested()
  283. status = await session.scalar(
  284. select(script_build_record.c.status)
  285. .where(script_build_record.c.id == token.script_build_id)
  286. .with_for_update()
  287. )
  288. if status is None:
  289. raise BuildNotFound()
  290. if status in {BuildStatus.STOPPING.value, BuildStatus.STOPPED.value}:
  291. raise MissionStopRequested()
  292. if status == BuildStatus.FAILED.value:
  293. raise MissionFencingTokenStale()
  294. if status == BuildStatus.SUCCESS.value and not allow_success:
  295. raise MissionFencingTokenStale()
  296. async def request_stop(self, script_build_id: int) -> int:
  297. """Persist stop intent without acquiring the process OwnerLease."""
  298. async with self._sessions() as session, session.begin():
  299. binding = (
  300. (
  301. await session.execute(
  302. select(mission_binding_table)
  303. .where(mission_binding_table.c.script_build_id == script_build_id)
  304. .with_for_update()
  305. )
  306. )
  307. .mappings()
  308. .one_or_none()
  309. )
  310. if binding is None:
  311. raise BuildNotFound()
  312. status = await session.scalar(
  313. select(script_build_record.c.status)
  314. .where(script_build_record.c.id == script_build_id)
  315. .with_for_update()
  316. )
  317. if status is None:
  318. raise BuildNotFound()
  319. if status in {BuildStatus.SUCCESS.value, BuildStatus.STOPPED.value}:
  320. return int(binding["stop_epoch"] or 0)
  321. stop_epoch = int(binding["stop_epoch"] or 0) + 1
  322. await session.execute(
  323. update(mission_binding_table)
  324. .where(mission_binding_table.c.id == binding["id"])
  325. .values(stop_epoch=stop_epoch, updated_at=datetime.now(UTC))
  326. )
  327. await session.execute(
  328. update(self._runtime_record)
  329. .where(self._runtime_record.c.id == script_build_id)
  330. .values(status=BuildStatus.STOPPING.value, end_time=None)
  331. )
  332. return stop_epoch
  333. __all__ = ["FencedCommandGate", "OwnerLease"]