from __future__ import annotations import json import logging from dataclasses import asdict from datetime import datetime, timedelta from typing import Any from sqlalchemy import Engine, delete, func, select from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import sessionmaker from cloud.db_models import ( Base, DeviceEnrollmentRow, HostRow, PluginRow, PooledDeviceRow, ScheduledTaskRow, TaskAttemptRow, AuthAuditRow, HostGovernancePolicyRow, LoginThrottleRow, UserRow, UserSessionRow, UserSubmissionPolicyRow, TokenReservationRow, TokenUsageEventRow, ) 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 enroll_host( self, *, host_id: str, agent_instance_id: str, credential_digest: str, enrollment_token_digest: str | None, display_name: str | None, enrolled_at: datetime, ) -> Any: from cloud.repository import ( EnrollmentTokenConflictError, HostEnrollmentConflictError, ) try: with self._sessions.begin() as session: statement = select(HostRow).where( HostRow.agent_instance_id == agent_instance_id ) if self.engine.dialect.name == "postgresql": statement = statement.with_for_update() existing = session.scalars(statement).first() if existing is not None: if ( existing.credential_digest != credential_digest or existing.enrollment_token_digest != enrollment_token_digest ): raise HostEnrollmentConflictError( "Host enrollment identity does not match existing binding" ) return _host_enrollment_from_row(existing) # A NULL enrollment_token_digest marks self-service Hosts, which # never bind a shared token; comparing `== NULL` would otherwise # match every other self-service Host via SQL's `IS NULL` and # falsely report a token conflict. token_owner = ( session.scalars( select(HostRow).where( HostRow.enrollment_token_digest == enrollment_token_digest ) ).first() if enrollment_token_digest is not None else None ) if token_owner is not None: raise EnrollmentTokenConflictError( "enrollment token is already bound to another Host" ) row = HostRow( host_id=host_id, address=None, last_seen_at=_iso(enrolled_at), agent_instance_id=agent_instance_id, credential_digest=credential_digest, enrollment_token_digest=enrollment_token_digest, display_name=display_name, enrolled_at=_iso(enrolled_at), revoked_at=None, ) session.add(row) session.flush() return _host_enrollment_from_row(row) except IntegrityError as exc: with self._sessions() as session: existing = session.scalars( select(HostRow).where( HostRow.agent_instance_id == agent_instance_id ) ).first() if ( existing is not None and existing.credential_digest == credential_digest and existing.enrollment_token_digest == enrollment_token_digest ): return _host_enrollment_from_row(existing) token_owner = ( session.scalars( select(HostRow).where( HostRow.enrollment_token_digest == enrollment_token_digest ) ).first() if enrollment_token_digest is not None else None ) if token_owner is not None: raise EnrollmentTokenConflictError( "enrollment token is already bound to another Host" ) from exc raise HostEnrollmentConflictError( "Host enrollment conflicts with an existing identity" ) from exc def authenticate_enrolled_host(self, credential_digest: str) -> str | None: with self._sessions() as session: host_id = session.scalar( select(HostRow.host_id) .where( HostRow.credential_digest == credential_digest, HostRow.revoked_at.is_(None), ) .limit(1) ) return str(host_id) if host_id is not None else None def revoke_enrolled_host( self, host_id: str, *, revoked_at: datetime, ) -> bool: with self._sessions.begin() as session: row = session.get( HostRow, host_id, with_for_update=self.engine.dialect.name == "postgresql", ) if row is None or row.credential_digest is None: return False row.revoked_at = _iso(revoked_at) return True def is_enrollment_managed_host(self, host_id: str) -> bool: with self._sessions() as session: credential_digest = session.scalar( select(HostRow.credential_digest).where(HostRow.host_id == host_id) ) return credential_digest is not None def enroll_device( self, *, device_id: str, host_id: str, local_device_id: str, driver_type: str, name: str | None, capability_tags: list[str], enrolled_at: datetime, ) -> Any: from cloud.repository import DeviceEnrollmentConflictError tags_json = json.dumps(list(capability_tags), ensure_ascii=False) try: with self._sessions.begin() as session: statement = select(DeviceEnrollmentRow).where( DeviceEnrollmentRow.host_id == host_id, DeviceEnrollmentRow.local_device_id == local_device_id, ) if self.engine.dialect.name == "postgresql": statement = statement.with_for_update() existing = session.scalars(statement).first() if existing is not None: if existing.driver_type != driver_type: raise DeviceEnrollmentConflictError( "device driver type does not match existing enrollment" ) if existing.revoked_at is not None: raise DeviceEnrollmentConflictError( "device enrollment has been revoked" ) existing.name = name existing.capability_tags_json = tags_json return _device_enrollment_from_row(existing) host = session.get(HostRow, host_id) if host is None: session.add( HostRow( host_id=host_id, address=None, last_seen_at=_iso(enrolled_at), ) ) row = DeviceEnrollmentRow( device_id=device_id, host_id=host_id, local_device_id=local_device_id, driver_type=driver_type, name=name, capability_tags_json=tags_json, enrolled_at=_iso(enrolled_at), revoked_at=None, ) session.add(row) session.flush() return _device_enrollment_from_row(row) except IntegrityError as exc: with self._sessions() as session: existing = session.scalars( select(DeviceEnrollmentRow).where( DeviceEnrollmentRow.host_id == host_id, DeviceEnrollmentRow.local_device_id == local_device_id, ) ).first() if existing is not None and existing.driver_type == driver_type: return _device_enrollment_from_row(existing) raise DeviceEnrollmentConflictError( "device enrollment conflicts with an existing identity" ) from exc def get_device_enrollment(self, device_id: str) -> Any | None: with self._sessions() as session: row = session.get(DeviceEnrollmentRow, device_id) return _device_enrollment_from_row(row) if row else None def list_device_enrollments(self, host_id: str) -> list[Any]: with self._sessions() as session: rows = session.scalars( select(DeviceEnrollmentRow) .where(DeviceEnrollmentRow.host_id == host_id) .order_by(DeviceEnrollmentRow.device_id) ).all() return [_device_enrollment_from_row(row) for row in rows] 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], *, allow_device_takeover: bool = False, ) -> None: with self._sessions.begin() as session: session.execute( delete(PooledDeviceRow).where(PooledDeviceRow.host_id == host_id) ) if allow_device_takeover: device_ids = [device.device_id for device in devices] if device_ids: session.execute( delete(PooledDeviceRow).where( PooledDeviceRow.host_id != host_id, PooledDeviceRow.device_id.in_(device_ids), ) ) 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 list_tasks( self, *, status: str | None = None, limit: int = 50, offset: int = 0, ) -> list[Any]: with self._sessions() as session: statement = select(ScheduledTaskRow) if status is not None: statement = statement.where(ScheduledTaskRow.status == status) statement = ( statement.order_by( ScheduledTaskRow.created_at.desc(), ScheduledTaskRow.id.desc(), ) .limit(limit) .offset(offset) ) rows = session.scalars(statement).all() return [_task_from_row(row) for row in rows] def count_tasks(self, status: str | None = None) -> int: with self._sessions() as session: statement = select(func.count()).select_from(ScheduledTaskRow) if status is not None: statement = statement.where(ScheduledTaskRow.status == status) count = session.scalar(statement) return int(count or 0) 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 # --------------------------------------------------------------- user auth def create_user(self, user: Any) -> Any: from cloud.repository import UserConflictError try: with self._sessions.begin() as session: row = UserRow( id=user.id, username=user.username, username_normalized=user.username_normalized, display_name=user.display_name, password_hash=user.password_hash, role=user.role, enabled=1 if user.enabled else 0, must_change_password=1 if user.must_change_password else 0, authentication_version=user.authentication_version, created_at=_iso(user.created_at), updated_at=_iso(user.updated_at), last_login_at=( _iso(user.last_login_at) if user.last_login_at is not None else None ), ) session.add(row) session.flush() return _user_from_row(row) except IntegrityError as exc: raise UserConflictError("username is already in use") from exc def get_user(self, user_id: str) -> Any | None: with self._sessions() as session: row = session.get(UserRow, user_id) return _user_from_row(row) if row else None def get_user_by_normalized_username(self, username_normalized: str) -> Any | None: with self._sessions() as session: row = session.scalar( select(UserRow) .where(UserRow.username_normalized == username_normalized) .limit(1) ) return _user_from_row(row) if row else None def list_users(self, *, limit: int, offset: int) -> list[Any]: with self._sessions() as session: rows = session.scalars( select(UserRow) .order_by(UserRow.created_at, UserRow.id) .limit(limit) .offset(offset) ).all() return [_user_from_row(row) for row in rows] def update_user( self, user_id: str, *, display_name: str | None = None, role: str | None = None, enabled: bool | None = None, updated_at: datetime, ) -> Any: from cloud.repository import LastAdministratorConflictError with self._sessions.begin() as session: row = session.get( UserRow, user_id, with_for_update=self.engine.dialect.name == "postgresql", ) if row is None: raise KeyError(f"user {user_id!r} not found") next_role = role if role is not None else row.role next_enabled = enabled if enabled is not None else bool(row.enabled) removes_administrator = ( bool(row.enabled) and row.role == "admin" and (next_role != "admin" or not next_enabled) ) if removes_administrator: other_admins = session.scalar( select(func.count()) .select_from(UserRow) .where( UserRow.id != user_id, UserRow.enabled == 1, UserRow.role == "admin", ) ) if int(other_admins or 0) == 0: raise LastAdministratorConflictError( "cannot remove the last enabled administrator" ) security_changed = next_role != row.role or next_enabled != bool(row.enabled) if display_name is not None: row.display_name = display_name row.role = next_role row.enabled = 1 if next_enabled else 0 row.updated_at = _iso(updated_at) if security_changed: row.authentication_version += 1 _revoke_user_session_rows(session, user_id, updated_at) session.flush() return _user_from_row(row) def rehash_user_password( self, user_id: str, *, password_hash: str, updated_at: datetime, ) -> Any: with self._sessions.begin() as session: row = session.get(UserRow, user_id) if row is None: raise KeyError(f"user {user_id!r} not found") row.password_hash = password_hash row.updated_at = _iso(updated_at) session.flush() return _user_from_row(row) def update_user_password( self, user_id: str, *, password_hash: str, must_change_password: bool, updated_at: datetime, revoke_sessions: bool, ) -> Any: with self._sessions.begin() as session: row = session.get( UserRow, user_id, with_for_update=self.engine.dialect.name == "postgresql", ) if row is None: raise KeyError(f"user {user_id!r} not found") row.password_hash = password_hash row.must_change_password = 1 if must_change_password else 0 row.authentication_version += 1 row.updated_at = _iso(updated_at) if revoke_sessions: _revoke_user_session_rows(session, user_id, updated_at) session.flush() return _user_from_row(row) def mark_user_login(self, user_id: str, *, now: datetime) -> Any: with self._sessions.begin() as session: row = session.get(UserRow, user_id) if row is None: raise KeyError(f"user {user_id!r} not found") row.last_login_at = _iso(now) row.updated_at = _iso(now) session.flush() return _user_from_row(row) def create_user_session(self, user_session: Any) -> None: with self._sessions.begin() as session: session.add( UserSessionRow( id=user_session.id, user_id=user_session.user_id, token_digest=user_session.token_digest, csrf_digest=user_session.csrf_digest, authentication_version=user_session.authentication_version, issued_at=_iso(user_session.issued_at), last_seen_at=_iso(user_session.last_seen_at), idle_expires_at=_iso(user_session.idle_expires_at), absolute_expires_at=_iso(user_session.absolute_expires_at), revoked_at=( _iso(user_session.revoked_at) if user_session.revoked_at is not None else None ), ) ) def get_authenticated_user_session( self, token_digest: str, *, now: datetime, ) -> Any | None: from cloud.user_auth import AuthenticatedUserSession with self._sessions() as session: match = session.execute( select(UserSessionRow, UserRow) .join(UserRow, UserRow.id == UserSessionRow.user_id) .where( UserSessionRow.token_digest == token_digest, UserSessionRow.revoked_at.is_(None), UserRow.enabled == 1, ) .limit(1) ).first() if match is None: return None session_row, user_row = match user = _user_from_row(user_row) user_session = _user_session_from_row(session_row) if ( user.authentication_version != user_session.authentication_version or user_session.idle_expires_at <= now or user_session.absolute_expires_at <= now ): return None return AuthenticatedUserSession(user=user, session=user_session) def touch_user_session( self, session_id: str, *, last_seen_at: datetime, idle_expires_at: datetime, ) -> Any: with self._sessions.begin() as session: row = session.get(UserSessionRow, session_id) if row is None: raise KeyError(f"session {session_id!r} not found") row.last_seen_at = _iso(last_seen_at) row.idle_expires_at = _iso(idle_expires_at) session.flush() return _user_session_from_row(row) def revoke_user_session(self, session_id: str, *, revoked_at: datetime) -> bool: with self._sessions.begin() as session: row = session.get(UserSessionRow, session_id) if row is None or row.revoked_at is not None: return False row.revoked_at = _iso(revoked_at) return True def revoke_user_sessions(self, user_id: str, *, revoked_at: datetime) -> int: with self._sessions.begin() as session: return _revoke_user_session_rows(session, user_id, revoked_at) def get_login_throttle( self, username_normalized: str, client_bucket: str, ) -> Any | None: with self._sessions() as session: row = session.get(LoginThrottleRow, (username_normalized, client_bucket)) return _login_throttle_from_row(row) if row else None def record_login_failure( self, *, username_normalized: str, client_bucket: str, now: datetime, failure_limit: int, failure_window: timedelta, block_duration: timedelta, ) -> Any: with self._sessions.begin() as session: row = session.get(LoginThrottleRow, (username_normalized, client_bucket)) if row is None: row = LoginThrottleRow( username_normalized=username_normalized, client_bucket=client_bucket, failure_count=0, window_started_at=_iso(now), last_attempt_at=_iso(now), blocked_until=None, ) session.add(row) window_started = _parse_dt(row.window_started_at) or now if now - window_started > failure_window: row.failure_count = 0 row.window_started_at = _iso(now) row.blocked_until = None row.failure_count += 1 row.last_attempt_at = _iso(now) if row.failure_count >= failure_limit: row.blocked_until = _iso(now + block_duration) session.flush() return _login_throttle_from_row(row) def clear_login_throttle(self, username_normalized: str, client_bucket: str) -> None: with self._sessions.begin() as session: row = session.get(LoginThrottleRow, (username_normalized, client_bucket)) if row is not None: session.delete(row) def record_auth_audit(self, event: Any) -> None: with self._sessions.begin() as session: session.add( AuthAuditRow( id=event.id, occurred_at=_iso(event.occurred_at), actor_principal_id=event.actor_principal_id, target_user_id=event.target_user_id, action=event.action, outcome=event.outcome, correlation_id=event.correlation_id, metadata_json=json.dumps(event.metadata, ensure_ascii=False), ) ) def cleanup_auth_state(self, *, now: datetime, limit: int) -> int: removed = 0 with self._sessions.begin() as session: expired_sessions = session.scalars( select(UserSessionRow) .where( (UserSessionRow.idle_expires_at <= _iso(now)) | (UserSessionRow.absolute_expires_at <= _iso(now)) ) .order_by(UserSessionRow.absolute_expires_at) .limit(limit) ).all() for row in expired_sessions: session.delete(row) removed += len(expired_sessions) remaining = max(0, limit - removed) if remaining: stale_before = now - timedelta(days=1) stale_throttles = session.scalars( select(LoginThrottleRow) .where(LoginThrottleRow.last_attempt_at <= _iso(stale_before)) .order_by(LoginThrottleRow.last_attempt_at) .limit(remaining) ).all() for row in stale_throttles: session.delete(row) removed += len(stale_throttles) return removed # -------------------------------------------------------------- governance def get_user_submission_policy(self, user_id: str) -> Any | None: with self._sessions() as session: row = session.get(UserSubmissionPolicyRow, user_id) return _user_submission_policy_from_row(row) if row is not None else None def upsert_user_submission_policy( self, *, user_id: str, submission_enabled: bool, allowed_host_ids: tuple[str, ...] | None, allowed_device_targets: tuple[tuple[str, str], ...] | None, updated_at: datetime, expected_revision: int | None = None, ) -> Any: from cloud.governance import GovernancePolicyConflictError with self._sessions.begin() as session: if session.get(UserRow, user_id) is None: raise KeyError(f"unknown user {user_id!r}") row = session.get( UserSubmissionPolicyRow, user_id, with_for_update=self.engine.dialect.name == "postgresql", ) if row is None: if expected_revision not in {None, 0}: raise GovernancePolicyConflictError("submission policy revision changed") row = UserSubmissionPolicyRow( user_id=user_id, revision=1, submission_enabled=1 if submission_enabled else 0, allowed_host_ids_json=_dump_optional_list(allowed_host_ids), allowed_device_targets_json=_dump_optional_list( allowed_device_targets ), updated_at=_iso(updated_at), ) session.add(row) else: if ( expected_revision is not None and expected_revision != row.revision ): raise GovernancePolicyConflictError("submission policy revision changed") row.revision += 1 row.submission_enabled = 1 if submission_enabled else 0 row.allowed_host_ids_json = _dump_optional_list(allowed_host_ids) row.allowed_device_targets_json = _dump_optional_list( allowed_device_targets ) row.updated_at = _iso(updated_at) session.flush() return _user_submission_policy_from_row(row) def get_host_governance_policy(self, host_id: str) -> Any | None: with self._sessions() as session: row = session.get(HostGovernancePolicyRow, host_id) return _host_governance_policy_from_row(row) if row is not None else None def upsert_host_governance_policy( self, *, host_id: str, self_submission_enabled: bool, max_active_tasks: int | None, daily_token_budget: int | None, updated_at: datetime, expected_revision: int | None = None, ) -> Any: from cloud.governance import GovernancePolicyConflictError with self._sessions.begin() as session: if session.get(HostRow, host_id) is None: raise KeyError(f"unknown host {host_id!r}") row = session.get( HostGovernancePolicyRow, host_id, with_for_update=self.engine.dialect.name == "postgresql", ) if row is None: if expected_revision not in {None, 0}: raise GovernancePolicyConflictError("Host policy revision changed") row = HostGovernancePolicyRow( host_id=host_id, revision=1, self_submission_enabled=1 if self_submission_enabled else 0, max_active_tasks=max_active_tasks, daily_token_budget=daily_token_budget, updated_at=_iso(updated_at), ) session.add(row) else: if ( expected_revision is not None and expected_revision != row.revision ): raise GovernancePolicyConflictError("Host policy revision changed") row.revision += 1 row.self_submission_enabled = 1 if self_submission_enabled else 0 row.max_active_tasks = max_active_tasks row.daily_token_budget = daily_token_budget row.updated_at = _iso(updated_at) session.flush() return _host_governance_policy_from_row(row) def count_active_tasks_for_host(self, host_id: str) -> int: with self._sessions() as session: count = session.scalar( select(func.count()) .select_from(ScheduledTaskRow) .where( ScheduledTaskRow.assigned_host_id == host_id, ScheduledTaskRow.status.in_(("assigned", "dispatched")), ) ) return int(count or 0) def reserve_host_token_budget( self, *, reservation_id: str, host_id: str, usage_day: str, reserved_tokens: int, task_id: str | None, attempt: int | None, created_at: datetime, expires_at: datetime, ) -> Any | None: from cloud.governance import TokenBudgetExceededError with self._sessions.begin() as session: statement = select(HostGovernancePolicyRow).where( HostGovernancePolicyRow.host_id == host_id ) if self.engine.dialect.name == "postgresql": statement = statement.with_for_update() policy = session.scalars(statement).first() if policy is None or policy.daily_token_budget is None: return None used = session.scalar( select(func.coalesce(func.sum(TokenUsageEventRow.total_tokens), 0)).where( TokenUsageEventRow.host_id == host_id, TokenUsageEventRow.usage_day == usage_day, ) ) reserved = session.scalar( select(func.coalesce(func.sum(TokenReservationRow.reserved_tokens), 0)).where( TokenReservationRow.host_id == host_id, TokenReservationRow.usage_day == usage_day, TokenReservationRow.expires_at > _iso(created_at), ) ) if int(used or 0) + int(reserved or 0) + reserved_tokens > policy.daily_token_budget: raise TokenBudgetExceededError("Host daily token budget is exhausted") row = TokenReservationRow( id=reservation_id, host_id=host_id, usage_day=usage_day, reserved_tokens=reserved_tokens, task_id=task_id, attempt=attempt, created_at=_iso(created_at), expires_at=_iso(expires_at), ) session.add(row) session.flush() return _token_reservation_from_row(row) def settle_host_token_reservation( self, *, reservation_id: str, event_id: str, provider: str, model: str, input_tokens: int | None, output_tokens: int | None, total_tokens: int, occurred_at: datetime, ) -> Any | None: with self._sessions.begin() as session: row = session.get( TokenReservationRow, reservation_id, with_for_update=self.engine.dialect.name == "postgresql", ) if row is None: return None event = TokenUsageEventRow( id=event_id, host_id=row.host_id, usage_day=row.usage_day, task_id=row.task_id, attempt=row.attempt, provider=provider, model=model, input_tokens=input_tokens, output_tokens=output_tokens, total_tokens=total_tokens, occurred_at=_iso(occurred_at), ) session.add(event) session.delete(row) session.flush() return _token_usage_event_from_row(event) def cleanup_expired_token_reservations(self, *, now: datetime, limit: int) -> int: with self._sessions.begin() as session: rows = session.scalars( select(TokenReservationRow) .where(TokenReservationRow.expires_at <= _iso(now)) .order_by(TokenReservationRow.expires_at) .limit(limit) ).all() for row in rows: session.delete(row) return len(rows) def get_host_token_usage_summary( self, *, host_id: str, usage_day: str, now: datetime, ) -> Any: from cloud.governance import TokenUsageSummary with self._sessions() as session: policy = session.get(HostGovernancePolicyRow, host_id) used = session.scalar( select(func.coalesce(func.sum(TokenUsageEventRow.total_tokens), 0)).where( TokenUsageEventRow.host_id == host_id, TokenUsageEventRow.usage_day == usage_day, ) ) reserved = session.scalar( select(func.coalesce(func.sum(TokenReservationRow.reserved_tokens), 0)).where( TokenReservationRow.host_id == host_id, TokenReservationRow.usage_day == usage_day, TokenReservationRow.expires_at > _iso(now), ) ) return TokenUsageSummary( host_id=host_id, usage_day=usage_day, daily_token_budget=(policy.daily_token_budget if policy else None), used_tokens=int(used or 0), reserved_tokens=int(reserved or 0), ) 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 policy_statement = select(HostGovernancePolicyRow).where( HostGovernancePolicyRow.host_id == host_id ) if self.engine.dialect.name == "postgresql": policy_statement = policy_statement.with_for_update() policy = session.scalars(policy_statement).first() if policy is not None and policy.max_active_tasks is not None: active_count = session.scalar( select(func.count()) .select_from(ScheduledTaskRow) .where( ScheduledTaskRow.assigned_host_id == host_id, ScheduledTaskRow.status.in_(("assigned", "dispatched")), ) ) if int(active_count or 0) >= policy.max_active_tasks: 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 _host_enrollment_from_row(row: HostRow) -> Any: from cloud.repository import HostEnrollment enrolled_at = _parse_dt(row.enrolled_at) or _parse_dt(row.last_seen_at) or utc_now() if row.agent_instance_id is None: raise ValueError(f"Host {row.host_id!r} is not enrollment-managed") return HostEnrollment( host_id=row.host_id, agent_instance_id=row.agent_instance_id, display_name=row.display_name, enrolled_at=enrolled_at, revoked_at=_parse_dt(row.revoked_at), ) def _device_enrollment_from_row(row: DeviceEnrollmentRow) -> Any: from cloud.repository import DeviceEnrollment try: tags = list(json.loads(row.capability_tags_json)) except TypeError, ValueError: tags = [] return DeviceEnrollment( device_id=row.device_id, host_id=row.host_id, local_device_id=row.local_device_id, driver_type=row.driver_type, name=row.name, capability_tags=tags, enrolled_at=_parse_dt(row.enrolled_at) or utc_now(), revoked_at=_parse_dt(row.revoked_at), ) 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 []), target_host_id=constraints_data.get("target_host_id"), target_device_id=constraints_data.get("target_device_id"), ), 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 _dump_optional_list(value: tuple[Any, ...] | None) -> str | None: return json.dumps(value, ensure_ascii=False) if value is not None else None def _load_optional_string_tuple(value: str | None) -> tuple[str, ...] | None: if value is None: return None try: parsed = json.loads(value) except (TypeError, ValueError): parsed = [] return tuple(item for item in parsed if isinstance(item, str)) def _load_optional_device_targets( value: str | None, ) -> tuple[tuple[str, str], ...] | None: if value is None: return None try: parsed = json.loads(value) except (TypeError, ValueError): parsed = [] return tuple( (item[0], item[1]) for item in parsed if isinstance(item, list | tuple) and len(item) == 2 and all(isinstance(part, str) for part in item) ) def _user_submission_policy_from_row(row: UserSubmissionPolicyRow) -> Any: from cloud.governance import UserSubmissionPolicy return UserSubmissionPolicy( user_id=row.user_id, revision=row.revision, submission_enabled=bool(row.submission_enabled), allowed_host_ids=_load_optional_string_tuple(row.allowed_host_ids_json), allowed_device_targets=_load_optional_device_targets( row.allowed_device_targets_json ), updated_at=_parse_dt(row.updated_at) or utc_now(), ) def _host_governance_policy_from_row(row: HostGovernancePolicyRow) -> Any: from cloud.governance import HostGovernancePolicy return HostGovernancePolicy( host_id=row.host_id, revision=row.revision, self_submission_enabled=bool(row.self_submission_enabled), max_active_tasks=row.max_active_tasks, daily_token_budget=row.daily_token_budget, updated_at=_parse_dt(row.updated_at) or utc_now(), ) def _token_reservation_from_row(row: TokenReservationRow) -> Any: from cloud.governance import TokenReservation return TokenReservation( id=row.id, host_id=row.host_id, usage_day=row.usage_day, reserved_tokens=row.reserved_tokens, task_id=row.task_id, attempt=row.attempt, created_at=_parse_dt(row.created_at) or utc_now(), expires_at=_parse_dt(row.expires_at) or utc_now(), ) def _token_usage_event_from_row(row: TokenUsageEventRow) -> Any: from cloud.governance import TokenUsageEvent return TokenUsageEvent( id=row.id, host_id=row.host_id, usage_day=row.usage_day, task_id=row.task_id, attempt=row.attempt, provider=row.provider, model=row.model, input_tokens=row.input_tokens, output_tokens=row.output_tokens, total_tokens=row.total_tokens, occurred_at=_parse_dt(row.occurred_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 _revoke_user_session_rows(session: Any, user_id: str, revoked_at: datetime) -> int: rows = session.scalars( select(UserSessionRow).where( UserSessionRow.user_id == user_id, UserSessionRow.revoked_at.is_(None), ) ).all() for row in rows: row.revoked_at = _iso(revoked_at) return len(rows) def _user_from_row(row: UserRow) -> Any: from cloud.user_auth import UserAccount return UserAccount( id=row.id, username=row.username, username_normalized=row.username_normalized, display_name=row.display_name, role=row.role, enabled=bool(row.enabled), must_change_password=bool(row.must_change_password), authentication_version=row.authentication_version, created_at=_parse_dt(row.created_at) or utc_now(), updated_at=_parse_dt(row.updated_at) or utc_now(), last_login_at=_parse_dt(row.last_login_at), password_hash=row.password_hash, ) def _user_session_from_row(row: UserSessionRow) -> Any: from cloud.user_auth import UserSession now = utc_now() return UserSession( id=row.id, user_id=row.user_id, token_digest=row.token_digest, csrf_digest=row.csrf_digest, authentication_version=row.authentication_version, issued_at=_parse_dt(row.issued_at) or now, last_seen_at=_parse_dt(row.last_seen_at) or now, idle_expires_at=_parse_dt(row.idle_expires_at) or now, absolute_expires_at=_parse_dt(row.absolute_expires_at) or now, revoked_at=_parse_dt(row.revoked_at), ) def _login_throttle_from_row(row: LoginThrottleRow) -> Any: from cloud.user_auth import LoginThrottle now = utc_now() return LoginThrottle( username_normalized=row.username_normalized, client_bucket=row.client_bucket, failure_count=row.failure_count, window_started_at=_parse_dt(row.window_started_at) or now, last_attempt_at=_parse_dt(row.last_attempt_at) or now, blocked_until=_parse_dt(row.blocked_until), ) 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), )