from __future__ import annotations import json import logging from dataclasses import asdict from datetime import datetime from typing import Any from sqlalchemy import Engine, delete, func, select from sqlalchemy.orm import sessionmaker from cloud.db_models import ( Base, HostRow, PluginRow, PooledDeviceRow, ScheduledTaskRow, TaskAttemptRow, ) from cloud.observability import current_correlation_id from core.models import utc_now logger = logging.getLogger(__name__) class SQLAlchemyCloudRepository: """SQLAlchemy adapter preserving the existing CloudStore CRUD surface.""" def __init__(self, engine: Engine, *, create_schema: bool = True) -> None: self.engine = engine self._sessions = sessionmaker(bind=engine, expire_on_commit=False) if create_schema: Base.metadata.create_all(engine) def upsert_host( self, host_id: str, *, address: str | None, last_seen_at: datetime, ) -> None: with self._sessions.begin() as session: row = session.get(HostRow, host_id) if row is None: session.add( HostRow( host_id=host_id, address=address, last_seen_at=_iso(last_seen_at), ) ) return if address is not None: row.address = address row.last_seen_at = _iso(last_seen_at) def replace_host_devices( self, host_id: str, devices: list[Any], ) -> None: with self._sessions.begin() as session: session.execute( delete(PooledDeviceRow).where(PooledDeviceRow.host_id == host_id) ) session.add_all( [ PooledDeviceRow( device_id=device.device_id, host_id=device.host_id, driver_type=device.driver_type, status=device.status, capability_tags_json=json.dumps( list(device.capability_tags), ensure_ascii=False, ), synced_at=_iso(device.synced_at) if device.synced_at else None, ) for device in devices ] ) def list_hosts(self) -> list[Any]: with self._sessions() as session: rows = session.scalars(select(HostRow).order_by(HostRow.host_id)).all() return [_host_from_row(row) for row in rows] def get_host(self, host_id: str) -> Any | None: with self._sessions() as session: row = session.get(HostRow, host_id) return _host_from_row(row) if row else None def list_devices(self) -> list[Any]: with self._sessions() as session: rows = session.scalars( select(PooledDeviceRow).order_by( PooledDeviceRow.host_id, PooledDeviceRow.device_id, ) ).all() return [_device_from_row(row) for row in rows] def get_device(self, device_id: str) -> Any | None: with self._sessions() as session: row = session.scalars( select(PooledDeviceRow) .where(PooledDeviceRow.device_id == device_id) .order_by(PooledDeviceRow.host_id) .limit(1) ).first() return _device_from_row(row) if row else None def enqueue_task(self, task: Any) -> None: with self._sessions.begin() as session: session.add( ScheduledTaskRow( id=task.id, goal=task.goal, workflow_definition_id=task.workflow_definition_id, constraints_json=json.dumps( asdict(task.constraints), ensure_ascii=False, ), status=task.status, assigned_device_id=task.assigned_device_id, assigned_host_id=task.assigned_host_id, attempt_count=task.attempt_count, lease_id=task.lease_id, lease_expires_at=( _iso(task.lease_expires_at) if task.lease_expires_at else None ), failure_reason=task.failure_reason, result_json=( json.dumps(task.terminal_result, ensure_ascii=False) if task.terminal_result is not None else None ), updated_at=_iso(task.updated_at) if task.updated_at else None, created_at=_iso(task.created_at), ) ) def list_queued_tasks(self) -> list[Any]: with self._sessions() as session: rows = session.scalars( select(ScheduledTaskRow) .where(ScheduledTaskRow.status == "queued") .order_by(ScheduledTaskRow.created_at, ScheduledTaskRow.id) ).all() return [_task_from_row(row) for row in rows] def get_task(self, task_id: str) -> Any | None: with self._sessions() as session: row = session.get(ScheduledTaskRow, task_id) return _task_from_row(row) if row else None def update_task( self, task_id: str, *, status: str | None = None, assigned_device_id: str | None = None, assigned_host_id: str | None = None, ) -> None: with self._sessions.begin() as session: row = session.get(ScheduledTaskRow, task_id) if row is None: return if status is not None: row.status = status if assigned_device_id is not None: row.assigned_device_id = assigned_device_id if assigned_host_id is not None: row.assigned_host_id = assigned_host_id def count_queued_tasks(self) -> int: with self._sessions() as session: count = session.scalar( select(func.count()) .select_from(ScheduledTaskRow) .where(ScheduledTaskRow.status == "queued") ) return int(count or 0) def save_plugin(self, manifest: Any, *, wired: bool) -> None: with self._sessions.begin() as session: row = session.get(PluginRow, manifest.name) if row is None: session.add( PluginRow( name=manifest.name, version=manifest.version, entry_point_kind=manifest.entry_point_kind, target=manifest.target, wired=1 if wired else 0, ) ) return row.version = manifest.version row.entry_point_kind = manifest.entry_point_kind row.target = manifest.target row.wired = 1 if wired else 0 def list_plugins(self) -> list[tuple[Any, bool]]: with self._sessions() as session: rows = session.scalars(select(PluginRow).order_by(PluginRow.name)).all() return [_plugin_from_row(row) for row in rows] def get_plugin(self, name: str) -> tuple[Any, bool] | None: with self._sessions() as session: row = session.get(PluginRow, name) return _plugin_from_row(row) if row else None def list_reserved_device_ids(self, *, now: datetime) -> set[str]: with self._sessions() as session: device_ids = session.scalars( select(ScheduledTaskRow.assigned_device_id).where( ScheduledTaskRow.status.in_(("assigned", "dispatched")), ScheduledTaskRow.assigned_device_id.is_not(None), ScheduledTaskRow.lease_expires_at.is_not(None), ScheduledTaskRow.lease_expires_at > _iso(now), ) ).all() return {device_id for device_id in device_ids if device_id is not None} def assign_task( self, *, task_id: str, host_id: str, device_id: str, lease_id: str, lease_expires_at: datetime, now: datetime, ) -> Any | None: with self._sessions.begin() as session: task_statement = select(ScheduledTaskRow).where( ScheduledTaskRow.id == task_id ) device_statement = select(PooledDeviceRow).where( PooledDeviceRow.host_id == host_id, PooledDeviceRow.device_id == device_id, ) if self.engine.dialect.name == "postgresql": task_statement = task_statement.with_for_update(skip_locked=True) device_statement = device_statement.with_for_update() task = session.scalars(task_statement).first() if task is None or task.status != "queued": return None device = session.scalars(device_statement).first() if device is None or device.status != "idle": return None active_reservation = session.scalar( select(ScheduledTaskRow.id) .where( ScheduledTaskRow.id != task_id, ScheduledTaskRow.assigned_device_id == device_id, ScheduledTaskRow.status.in_(("assigned", "dispatched")), ScheduledTaskRow.lease_expires_at.is_not(None), ScheduledTaskRow.lease_expires_at > _iso(now), ) .limit(1) ) if active_reservation is not None: return None attempt = task.attempt_count + 1 task.status = "assigned" task.assigned_host_id = host_id task.assigned_device_id = device_id task.attempt_count = attempt task.lease_id = lease_id task.lease_expires_at = _iso(lease_expires_at) task.failure_reason = None task.result_json = None task.updated_at = _iso(now) session.add( TaskAttemptRow( task_id=task.id, attempt=attempt, lease_id=lease_id, host_id=host_id, device_id=device_id, status="assigned", lease_expires_at=_iso(lease_expires_at), created_at=_iso(now), completed_at=None, failure_reason=None, result_json=None, ) ) session.flush() _log_task_lifecycle("assigned", task) return _leased_assignment_from_row(task) def claim_assignment( self, *, host_id: str, now: datetime, ) -> Any | None: with self._sessions.begin() as session: statement = ( select(ScheduledTaskRow) .where( ScheduledTaskRow.status == "assigned", ScheduledTaskRow.assigned_host_id == host_id, ScheduledTaskRow.lease_id.is_not(None), ScheduledTaskRow.lease_expires_at.is_not(None), ScheduledTaskRow.lease_expires_at > _iso(now), ) .order_by(ScheduledTaskRow.created_at, ScheduledTaskRow.id) .limit(1) ) if self.engine.dialect.name == "postgresql": statement = statement.with_for_update(skip_locked=True) task = session.scalars(statement).first() if task is None: return None attempt = session.get( TaskAttemptRow, (task.id, task.attempt_count), with_for_update=self.engine.dialect.name == "postgresql", ) if ( attempt is None or attempt.status != "assigned" or attempt.lease_id != task.lease_id ): return None task.status = "dispatched" task.updated_at = _iso(now) attempt.status = "dispatched" session.flush() _log_task_lifecycle("claimed", task) return _leased_assignment_from_row(task) def renew_lease( self, *, task_id: str, attempt: int, lease_id: str, host_id: str, lease_expires_at: datetime, now: datetime, ) -> str: with self._sessions.begin() as session: task = session.get( ScheduledTaskRow, task_id, with_for_update=self.engine.dialect.name == "postgresql", ) if task is None: return "not_found" if ( task.status not in {"assigned", "dispatched"} or task.attempt_count != attempt or task.lease_id != lease_id or task.assigned_host_id != host_id ): return "conflict" current_expiry = _parse_dt(task.lease_expires_at) if current_expiry is None or current_expiry <= now: return "expired" if lease_expires_at <= now: return "conflict" attempt_row = session.get( TaskAttemptRow, (task_id, attempt), with_for_update=self.engine.dialect.name == "postgresql", ) if ( attempt_row is None or attempt_row.status not in {"assigned", "dispatched"} or attempt_row.lease_id != lease_id or attempt_row.host_id != host_id ): return "conflict" renewed_until = _iso(lease_expires_at) task.lease_expires_at = renewed_until task.updated_at = _iso(now) attempt_row.lease_expires_at = renewed_until _log_task_lifecycle("renewed", task) return "renewed" def record_task_result( self, *, task_id: str, attempt: int, lease_id: str, host_id: str, status: str, failure_reason: str | None, terminal_result: dict[str, Any] | None, completed_at: datetime, ) -> str: with self._sessions.begin() as session: task = session.get( ScheduledTaskRow, task_id, with_for_update=self.engine.dialect.name == "postgresql", ) if task is None: return "conflict" attempt_row = session.get( TaskAttemptRow, (task_id, attempt), with_for_update=self.engine.dialect.name == "postgresql", ) if ( attempt_row is None or task.attempt_count != attempt or task.lease_id != lease_id or task.assigned_host_id != host_id or attempt_row.lease_id != lease_id or attempt_row.host_id != host_id ): return "conflict" if task.status in {"done", "failed"}: if ( task.status == status and task.failure_reason == failure_reason and _parse_json_object(task.result_json) == terminal_result and attempt_row.status == status ): return "already_recorded" return "conflict" if task.status not in {"assigned", "dispatched"}: return "conflict" current_expiry = _parse_dt(task.lease_expires_at) if current_expiry is None or current_expiry <= completed_at: return "conflict" if status not in {"done", "failed"}: return "conflict" result_json = ( json.dumps(terminal_result, ensure_ascii=False) if terminal_result is not None else None ) completed_at_iso = _iso(completed_at) task.status = status task.failure_reason = failure_reason task.result_json = result_json task.updated_at = completed_at_iso attempt_row.status = status attempt_row.completed_at = completed_at_iso attempt_row.failure_reason = failure_reason attempt_row.result_json = result_json _log_task_lifecycle("completed" if status == "done" else "failed", task) return "recorded" def reap_expired_leases( self, *, now: datetime, max_attempts: int, ) -> list[str]: if max_attempts < 1: raise ValueError("max_attempts must be at least 1") with self._sessions.begin() as session: statement = ( select(ScheduledTaskRow) .where( ScheduledTaskRow.status.in_(("assigned", "dispatched")), ScheduledTaskRow.lease_expires_at.is_not(None), ScheduledTaskRow.lease_expires_at <= _iso(now), ) .order_by(ScheduledTaskRow.lease_expires_at, ScheduledTaskRow.id) ) if self.engine.dialect.name == "postgresql": statement = statement.with_for_update(skip_locked=True) tasks = session.scalars(statement).all() reaped_task_ids: list[str] = [] for task in tasks: attempt_row = session.get( TaskAttemptRow, (task.id, task.attempt_count), with_for_update=self.engine.dialect.name == "postgresql", ) if attempt_row is None: continue attempt_row.status = "expired" attempt_row.completed_at = _iso(now) attempt_row.failure_reason = "lease expired" task.updated_at = _iso(now) task.lease_id = None task.lease_expires_at = None task.result_json = None if task.attempt_count < max_attempts: task.status = "queued" task.assigned_host_id = None task.assigned_device_id = None task.failure_reason = None else: task.status = "failed" task.failure_reason = ( f"lease expired after {task.attempt_count} attempts" ) _log_task_lifecycle( "retried" if task.status == "queued" else "failed", task, ) reaped_task_ids.append(task.id) return reaped_task_ids def list_task_attempts(self, task_id: str) -> list[Any]: with self._sessions() as session: rows = session.scalars( select(TaskAttemptRow) .where(TaskAttemptRow.task_id == task_id) .order_by(TaskAttemptRow.attempt) ).all() return [_task_attempt_from_row(row) for row in rows] def health_check(self) -> None: with self._sessions() as session: session.execute(select(1)) def close(self) -> None: self.engine.dispose() def _iso(value: datetime) -> str: return value.isoformat() def _parse_dt(value: str | None) -> datetime | None: if not value: return None try: return datetime.fromisoformat(value) except ValueError: return None def _host_from_row(row: HostRow) -> Any: from cloud.pool import HostRegistration return HostRegistration( host_id=row.host_id, address=row.address, last_seen_at=_parse_dt(row.last_seen_at) or utc_now(), ) def _device_from_row(row: PooledDeviceRow) -> Any: from cloud.pool import PooledDevice try: tags = list(json.loads(row.capability_tags_json)) except TypeError, ValueError: tags = [] return PooledDevice( device_id=row.device_id, host_id=row.host_id, driver_type=row.driver_type, status=row.status, capability_tags=tags, synced_at=_parse_dt(row.synced_at), ) def _task_from_row(row: ScheduledTaskRow) -> Any: from cloud.scheduler import ScheduledTask, TaskConstraints try: constraints_data = json.loads(row.constraints_json) except TypeError, ValueError: constraints_data = {} terminal_result = _parse_json_object(row.result_json) return ScheduledTask( id=row.id, goal=row.goal, workflow_definition_id=row.workflow_definition_id, constraints=TaskConstraints( driver_type=constraints_data.get("driver_type"), capability_tags=list(constraints_data.get("capability_tags") or []), ), status=row.status, assigned_device_id=row.assigned_device_id, assigned_host_id=row.assigned_host_id, attempt_count=row.attempt_count, lease_id=row.lease_id, lease_expires_at=_parse_dt(row.lease_expires_at), terminal_result=terminal_result, failure_reason=row.failure_reason, updated_at=_parse_dt(row.updated_at), created_at=_parse_dt(row.created_at) or utc_now(), ) def _task_attempt_from_row(row: TaskAttemptRow) -> Any: from cloud.repository import TaskAttemptRecord return TaskAttemptRecord( task_id=row.task_id, attempt=row.attempt, lease_id=row.lease_id, host_id=row.host_id, device_id=row.device_id, status=row.status, lease_expires_at=_parse_dt(row.lease_expires_at) or utc_now(), created_at=_parse_dt(row.created_at) or utc_now(), completed_at=_parse_dt(row.completed_at), failure_reason=row.failure_reason, terminal_result=_parse_json_object(row.result_json), ) def _leased_assignment_from_row(row: ScheduledTaskRow) -> Any: from cloud.repository import LeasedAssignment lease_expires_at = _parse_dt(row.lease_expires_at) if ( row.lease_id is None or lease_expires_at is None or row.assigned_host_id is None or row.assigned_device_id is None ): raise ValueError(f"task {row.id!r} does not contain a complete lease") return LeasedAssignment( task_id=row.id, attempt=row.attempt_count, lease_id=row.lease_id, lease_expires_at=lease_expires_at, host_id=row.assigned_host_id, device_id=row.assigned_device_id, goal=row.goal, workflow_definition_id=row.workflow_definition_id, ) def _parse_json_object(value: str | None) -> dict[str, Any] | None: if value is None: return None try: parsed = json.loads(value) except TypeError, ValueError: return None return parsed if isinstance(parsed, dict) else None def _log_task_lifecycle(event: str, task: ScheduledTaskRow) -> None: logger.info( "cloud task lifecycle", extra={ "event": event, "correlation_id": current_correlation_id(), "task_id": task.id, "host_id": task.assigned_host_id, "device_id": task.assigned_device_id, "attempt": task.attempt_count, "lease_id": task.lease_id, }, ) def _plugin_from_row(row: PluginRow) -> tuple[Any, bool]: from cloud.plugins import PluginManifest return ( PluginManifest( name=row.name, version=row.version, entry_point_kind=row.entry_point_kind, target=row.target, ), bool(row.wired), )