from __future__ import annotations import asyncio from collections.abc import Callable from datetime import UTC, datetime from threading import Event, Lock from typing import Protocol import httpx from cloud.internal_api.models import AssignmentModel from host_agent.assignment import AssignmentExecutionResult from host_agent.client import HostAgentAPIError, HostAgentClient, StaleLeaseError class InterruptibleAssignmentExecutor(Protocol): def execute( self, assignment: AssignmentModel, *, should_stop: Callable[[], bool] | None = None, ) -> AssignmentExecutionResult: ... class LeaseGuard: def __init__(self) -> None: self._lost = Event() self._lock = Lock() self._reason: str | None = None @property def reason(self) -> str | None: with self._lock: return self._reason def is_lost(self) -> bool: return self._lost.is_set() def mark_lost(self, reason: str) -> None: with self._lock: if self._reason is None: self._reason = reason self._lost.set() class ActiveAssignmentRunner: def __init__( self, client: HostAgentClient, executor: InterruptibleAssignmentExecutor, *, now: Callable[[], datetime] | None = None, ) -> None: self.client = client self.executor = executor self._now = now or (lambda: datetime.now(UTC)) self._stop_requested = Event() def request_stop(self) -> None: self._stop_requested.set() async def run(self, assignment: AssignmentModel) -> AssignmentExecutionResult: guard = LeaseGuard() execution = asyncio.create_task( asyncio.to_thread( self.executor.execute, assignment, should_stop=lambda: guard.is_lost() or self._stop_requested.is_set(), ) ) renewal = asyncio.create_task( self._renew_while_running(assignment, execution, guard) ) try: return await execution finally: await renewal async def _renew_while_running( self, assignment: AssignmentModel, execution: asyncio.Task[AssignmentExecutionResult], guard: LeaseGuard, ) -> None: lease_expires_at = assignment.lease_expires_at while not execution.done(): delay = max( 0.0, (lease_expires_at - self._now()).total_seconds() / 3, ) done, _ = await asyncio.wait({execution}, timeout=delay) if done: return try: response = await self.client.renew(assignment) except StaleLeaseError: guard.mark_lost("lease rejected by control plane") return except HostAgentAPIError as exc: guard.mark_lost(f"lease renewal rejected: {exc.detail}") return except httpx.TransportError: guard.mark_lost("lease renewal failed after transport retries") return else: lease_expires_at = response.lease_expires_at