108 lines
3.2 KiB
Python
108 lines
3.2 KiB
Python
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
|