from __future__ import annotations import json import logging from dataclasses import asdict from datetime import datetime, timedelta from typing import TYPE_CHECKING, Any from uuid import uuid4 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, PlannerDecisionLogRow, PluginRow, PooledDeviceRow, ScheduledTaskRow, TaskAttemptRow, AuthAuditRow, HostGovernancePolicyRow, LlmProviderProfileRow, LlmProviderSettingsRow, LoginThrottleRow, UserRow, UserSessionRow, UserSubmissionPolicyRow, TokenReservationRow, TokenUsageEventRow, CloudSkillRow, CloudSkillEntitlementRow, CloudSkillSyncStateRow, CloudSkillHostChangeRow, HostSkillInventoryRow, ) from cloud.observability import current_correlation_id from core.models import utc_now if TYPE_CHECKING: from cloud.repository import AssignmentProgressSnapshot 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) self._ensure_llm_provider_settings() def _ensure_llm_provider_settings(self) -> None: with self._sessions.begin() as session: if session.get(LlmProviderSettingsRow, "global") is None: session.add( LlmProviderSettingsRow( id="global", active_profile_id=None, revision=0, updated_at=_iso(utc_now()), ) ) 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, planner_transport: str = "direct", ) -> 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), planner_transport=planner_transport, ) ) return if address is not None: row.address = address row.last_seen_at = _iso(last_seen_at) row.planner_transport = planner_transport 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) # --------------------------------------------------------- LLM Providers def get_llm_provider_settings(self) -> Any: with self._sessions() as session: row = session.get(LlmProviderSettingsRow, "global") return _llm_provider_settings_from_row(row) def list_llm_provider_profiles(self) -> list[Any]: with self._sessions() as session: rows = session.scalars( select(LlmProviderProfileRow).order_by( LlmProviderProfileRow.name_normalized ) ).all() return [_llm_provider_profile_from_row(row) for row in rows] def get_llm_provider_profile(self, profile_id: str) -> Any | None: with self._sessions() as session: row = session.get(LlmProviderProfileRow, profile_id) return _llm_provider_profile_from_row(row) if row is not None else None def create_llm_provider_profile(self, profile: Any) -> Any: from cloud.llm_providers import LlmProviderConflictError try: with self._sessions.begin() as session: if ( session.scalars( select(LlmProviderProfileRow).where( LlmProviderProfileRow.name_normalized == profile.name_normalized ) ).first() is not None ): raise LlmProviderConflictError( "Provider profile name already exists" ) row = LlmProviderProfileRow( id=profile.id, name=profile.name, name_normalized=profile.name_normalized, provider_type=profile.provider_type, model=profile.model, base_url=profile.base_url, timeout_seconds=profile.timeout_seconds, api_key_ciphertext=profile.api_key_ciphertext, key_last_rotated_at=_iso(profile.key_last_rotated_at), enabled=1 if profile.enabled else 0, revision=profile.revision, created_at=_iso(profile.created_at), updated_at=_iso(profile.updated_at), ) session.add(row) session.flush() return _llm_provider_profile_from_row(row) except IntegrityError as exc: raise LlmProviderConflictError( "Provider profile name already exists" ) from exc def update_llm_provider_profile( self, profile_id: str, *, name: str, name_normalized: str, provider_type: str, model: str, base_url: str | None, timeout_seconds: float, api_key_ciphertext: str, key_last_rotated_at: datetime, enabled: bool, expected_revision: int | None, updated_at: datetime, ) -> Any: from cloud.llm_providers import ( LlmProviderActiveConflictError, LlmProviderConflictError, ) try: with self._sessions.begin() as session: row = session.get( LlmProviderProfileRow, profile_id, with_for_update=self.engine.dialect.name == "postgresql", ) if row is None: raise KeyError(profile_id) if expected_revision is not None and expected_revision != row.revision: raise LlmProviderConflictError("Provider profile revision changed") existing_name = session.scalars( select(LlmProviderProfileRow).where( LlmProviderProfileRow.name_normalized == name_normalized, LlmProviderProfileRow.id != profile_id, ) ).first() if existing_name is not None: raise LlmProviderConflictError( "Provider profile name already exists" ) settings = session.get( LlmProviderSettingsRow, "global", with_for_update=self.engine.dialect.name == "postgresql", ) if ( not enabled and settings is not None and settings.active_profile_id == profile_id ): raise LlmProviderActiveConflictError( "activate another Provider profile before disabling this one" ) row.name = name row.name_normalized = name_normalized row.provider_type = provider_type row.model = model row.base_url = base_url row.timeout_seconds = timeout_seconds row.api_key_ciphertext = api_key_ciphertext row.key_last_rotated_at = _iso(key_last_rotated_at) row.enabled = 1 if enabled else 0 row.revision += 1 row.updated_at = _iso(updated_at) session.flush() return _llm_provider_profile_from_row(row) except IntegrityError as exc: raise LlmProviderConflictError( "Provider profile name already exists" ) from exc def activate_llm_provider_profile( self, profile_id: str, *, expected_settings_revision: int | None, updated_at: datetime, ) -> Any: from cloud.llm_providers import LlmProviderConflictError with self._sessions.begin() as session: profile = session.get( LlmProviderProfileRow, profile_id, with_for_update=self.engine.dialect.name == "postgresql", ) if profile is None: raise KeyError(profile_id) if not profile.enabled: raise LlmProviderConflictError( "Provider profile must be enabled before activation" ) settings = session.get( LlmProviderSettingsRow, "global", with_for_update=self.engine.dialect.name == "postgresql", ) if settings is None: if expected_settings_revision not in {None, 0}: raise LlmProviderConflictError("Provider settings revision changed") settings = LlmProviderSettingsRow( id="global", active_profile_id=profile_id, revision=1, updated_at=_iso(updated_at), ) session.add(settings) else: if ( expected_settings_revision is not None and expected_settings_revision != settings.revision ): raise LlmProviderConflictError("Provider settings revision changed") settings.active_profile_id = profile_id settings.revision += 1 settings.updated_at = _iso(updated_at) session.flush() return _llm_provider_settings_from_row(settings) def delete_llm_provider_profile( self, profile_id: str, *, expected_revision: int | None, ) -> None: from cloud.llm_providers import ( LlmProviderActiveConflictError, LlmProviderConflictError, ) with self._sessions.begin() as session: row = session.get( LlmProviderProfileRow, profile_id, with_for_update=self.engine.dialect.name == "postgresql", ) if row is None: raise KeyError(profile_id) if expected_revision is not None and expected_revision != row.revision: raise LlmProviderConflictError("Provider profile revision changed") settings = session.get( LlmProviderSettingsRow, "global", with_for_update=self.engine.dialect.name == "postgresql", ) if settings is not None and settings.active_profile_id == profile_id: raise LlmProviderActiveConflictError( "activate another Provider profile before deleting this one" ) session.delete(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_host_token_usage_events( self, *, host_id: str, limit: int, offset: int, ) -> list[Any]: with self._sessions() as session: rows = session.scalars( select(TokenUsageEventRow) .where(TokenUsageEventRow.host_id == host_id) .order_by( TokenUsageEventRow.occurred_at.desc(), TokenUsageEventRow.id.desc() ) .limit(limit) .offset(offset) ).all() return [_token_usage_event_from_row(row) for row in rows] 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, progress: AssignmentProgressSnapshot | None = None, ) -> 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 if progress is not None: task.progress_step_index = progress.step_index task.progress_step_status = progress.step_status task.progress_summary = progress.summary task.progress_updated_at = _iso(progress.updated_at) _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 task.progress_step_index = None task.progress_step_status = None task.progress_summary = None task.progress_updated_at = None 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 record_planner_decision( self, *, host_id: str, task_id: str, attempt: int, system_prompt: str, user_prompt: str, tool_name: str, arguments_json: str, now: datetime, rationale: str | None = None, thinking: str | None = None, purpose: str | None = None, expected_outcome: str | None = None, ) -> int: with self._sessions.begin() as session: current_max = session.scalars( select(func.coalesce(func.max(PlannerDecisionLogRow.step_index), 0)) .where(PlannerDecisionLogRow.task_id == task_id) .where(PlannerDecisionLogRow.attempt == attempt) ).one() next_step = current_max + 1 session.add( PlannerDecisionLogRow( id=uuid4().hex, host_id=host_id, task_id=task_id, attempt=attempt, step_index=next_step, system_prompt=system_prompt, user_prompt=user_prompt, tool_name=tool_name, arguments_json=arguments_json, created_at=_iso(now), rationale=rationale, thinking=thinking, purpose=purpose, expected_outcome=expected_outcome, ) ) session.flush() return next_step def prune_planner_decision_log( self, *, now: datetime, prune_after_terminal_seconds: int, ) -> int: cutoff_iso = _iso(now - timedelta(seconds=prune_after_terminal_seconds)) with self._sessions.begin() as session: terminal_task_ids = select(TaskAttemptRow.task_id).where( TaskAttemptRow.status.in_(("done", "failed")), TaskAttemptRow.completed_at.is_not(None), TaskAttemptRow.completed_at <= cutoff_iso, ) result = session.execute( delete(PlannerDecisionLogRow).where( PlannerDecisionLogRow.task_id.in_(terminal_task_ids) ) ) return result.rowcount or 0 def list_planner_decisions( self, *, task_id: str, attempt: int, ) -> list[Any]: from cloud.repository import PlannerDecisionRecord with self._sessions() as session: rows = session.scalars( select(PlannerDecisionLogRow) .where(PlannerDecisionLogRow.task_id == task_id) .where(PlannerDecisionLogRow.attempt == attempt) .order_by(PlannerDecisionLogRow.step_index) ).all() return [ PlannerDecisionRecord( id=row.id, host_id=row.host_id, task_id=row.task_id, attempt=row.attempt, step_index=row.step_index, system_prompt=row.system_prompt, user_prompt=row.user_prompt, tool_name=row.tool_name, arguments_json=row.arguments_json, created_at=_parse_dt(row.created_at), # type: ignore[arg-type] rationale=row.rationale, thinking=row.thinking, purpose=row.purpose, expected_outcome=row.expected_outcome, ) 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() # ------------------------------------------------------------------ # Cloud-managed skills + per-host entitlement + sync versioning # ------------------------------------------------------------------ def list_cloud_skills(self) -> list[Any]: with self._sessions() as session: rows = session.scalars( select(CloudSkillRow).order_by(CloudSkillRow.name_normalized) ).all() return [_cloud_skill_from_row(row) for row in rows] def get_cloud_skill(self, skill_id: str) -> Any | None: with self._sessions() as session: row = session.get(CloudSkillRow, skill_id) return _cloud_skill_from_row(row) if row is not None else None def create_cloud_skill(self, skill: Any) -> Any: from cloud.skills import CloudSkillConflictError try: with self._sessions.begin() as session: if ( session.scalars( select(CloudSkillRow).where( CloudSkillRow.name_normalized == skill.name_normalized ) ).first() is not None ): raise CloudSkillConflictError("Skill name already exists") row = _cloud_skill_to_row(skill) session.add(row) session.flush() return _cloud_skill_from_row(row) except IntegrityError as exc: raise CloudSkillConflictError("Skill name already exists") from exc def update_cloud_skill( self, skill_id: str, *, name: str, name_normalized: str, kind: str, description: str, tags: list[str], content: str, steps_json: str, parameters_json: str, updated_at: datetime, ) -> Any: from cloud.skills import CloudSkillConflictError, CloudSkillValidationError try: with self._sessions.begin() as session: row = session.get(CloudSkillRow, skill_id) if row is None: raise CloudSkillValidationError("Unknown skill") clash = session.scalars( select(CloudSkillRow).where( CloudSkillRow.name_normalized == name_normalized, CloudSkillRow.id != skill_id, ) ).first() if clash is not None: raise CloudSkillConflictError("Skill name already exists") row.name = name row.name_normalized = name_normalized row.kind = kind row.description = description row.tags_json = json.dumps(list(tags)) row.content = content row.steps_json = steps_json row.parameters_json = parameters_json row.revision = row.revision + 1 row.updated_at = _iso(updated_at) session.flush() updated = _cloud_skill_from_row(row) # Notify every entitled host that the skill changed. host_ids = [ row.host_id for row in session.scalars( select(CloudSkillEntitlementRow).where( CloudSkillEntitlementRow.skill_id == skill_id ) ).all() ] for host_id in host_ids: self._bump_host(session, host_id, skill_id, "upsert", updated_at) return updated except IntegrityError as exc: raise CloudSkillConflictError("Skill name already exists") from exc def delete_cloud_skill(self, skill_id: str) -> None: with self._sessions.begin() as session: host_ids = [ row.host_id for row in session.scalars( select(CloudSkillEntitlementRow).where( CloudSkillEntitlementRow.skill_id == skill_id ) ).all() ] now = utc_now() for host_id in host_ids: self._bump_host(session, host_id, skill_id, "remove", now) session.execute( delete(CloudSkillEntitlementRow).where( CloudSkillEntitlementRow.skill_id == skill_id ) ) session.execute(delete(CloudSkillRow).where(CloudSkillRow.id == skill_id)) def list_entitlements_for_skill(self, skill_id: str) -> list[str]: with self._sessions() as session: rows = session.scalars( select(CloudSkillEntitlementRow).where( CloudSkillEntitlementRow.skill_id == skill_id ) ).all() return [row.host_id for row in rows] def list_entitled_skills_for_host(self, host_id: str) -> list[Any]: with self._sessions() as session: entitlements = session.scalars( select(CloudSkillEntitlementRow).where( CloudSkillEntitlementRow.host_id == host_id ) ).all() skills: list[Any] = [] for ent in entitlements: row = session.get(CloudSkillRow, ent.skill_id) if row is not None: skills.append(_cloud_skill_from_row(row)) skills.sort(key=lambda s: s.name_normalized) return skills def grant_entitlement(self, skill_id: str, host_id: str, *, now: datetime) -> None: with self._sessions.begin() as session: existing = session.get(CloudSkillEntitlementRow, (skill_id, host_id)) if existing is not None: return # idempotent session.add( CloudSkillEntitlementRow( skill_id=skill_id, host_id=host_id, granted_at=_iso(now), ) ) self._bump_host(session, host_id, skill_id, "upsert", now) def revoke_entitlement(self, skill_id: str, host_id: str, *, now: datetime) -> None: with self._sessions.begin() as session: existing = session.get(CloudSkillEntitlementRow, (skill_id, host_id)) if existing is None: return session.delete(existing) self._bump_host(session, host_id, skill_id, "remove", now) def fetch_host_delta(self, host_id: str, since_version: int | None) -> Any: from cloud.skills import HostSkillDelta with self._sessions.begin() as session: state = self._ensure_sync_state(session, host_id) latest = state.last_version if since_version is None: skills = self._host_skills_in_session(session, host_id) return HostSkillDelta( skills=skills, removed_ids=[], latest_version=latest, is_full_replace=True, ) oldest = self._oldest_changelog_version(session, host_id) if oldest is None or since_version < oldest - 1: skills = self._host_skills_in_session(session, host_id) return HostSkillDelta( skills=skills, removed_ids=[], latest_version=latest, is_full_replace=True, ) changes = ( session.scalars( select(CloudSkillHostChangeRow) .where( CloudSkillHostChangeRow.host_id == host_id, CloudSkillHostChangeRow.version > since_version, ) .order_by(CloudSkillHostChangeRow.version) ) .unique() .all() ) upsert_ids: list[str] = [] removed_ids: list[str] = [] seen: set[str] = set() for change in changes: if change.skill_id in seen: continue seen.add(change.skill_id) if change.change_type == "remove": removed_ids.append(change.skill_id) else: upsert_ids.append(change.skill_id) skills: list[Any] = [] entitled_ids = { ent.skill_id for ent in session.scalars( select(CloudSkillEntitlementRow).where( CloudSkillEntitlementRow.host_id == host_id ) ).all() } for sid in upsert_ids: if sid in entitled_ids: row = session.get(CloudSkillRow, sid) if row is not None: skills.append(_cloud_skill_from_row(row)) skills.sort(key=lambda s: s.name_normalized) return HostSkillDelta( skills=skills, removed_ids=removed_ids, latest_version=latest, is_full_replace=False, ) def record_host_inventory( self, host_id: str, payload_json: str, *, now: datetime ) -> None: with self._sessions.begin() as session: row = session.get(HostSkillInventoryRow, host_id) if row is None: session.add( HostSkillInventoryRow( host_id=host_id, payload_json=payload_json, reported_at=_iso(now), ) ) else: row.payload_json = payload_json row.reported_at = _iso(now) def get_host_inventory(self, host_id: str) -> Any | None: from cloud.skills import HostInventoryEntry with self._sessions() as session: row = session.get(HostSkillInventoryRow, host_id) if row is None: return None return HostInventoryEntry( host_id=row.host_id, payload_json=row.payload_json, reported_at=_parse_dt(row.reported_at), ) # -- internals -------------------------------------------------------- def _ensure_sync_state(self, session: Any, host_id: str) -> Any: row = session.get(CloudSkillSyncStateRow, host_id) if row is None: row = CloudSkillSyncStateRow( host_id=host_id, last_version=0, updated_at=_iso(utc_now()), ) session.add(row) session.flush() return row def _bump_host( self, session: Any, host_id: str, skill_id: str, change_type: str, now: datetime, ) -> None: state = self._ensure_sync_state(session, host_id) state.last_version = state.last_version + 1 state.updated_at = _iso(now) session.add( CloudSkillHostChangeRow( host_id=host_id, version=state.last_version, skill_id=skill_id, change_type=change_type, recorded_at=_iso(now), ) ) def _oldest_changelog_version(self, session: Any, host_id: str) -> int | None: row = session.scalars( select(func.min(CloudSkillHostChangeRow.version)).where( CloudSkillHostChangeRow.host_id == host_id ) ).first() return int(row) if row is not None else None def _host_skills_in_session(self, session: Any, host_id: str) -> list[Any]: entitlements = session.scalars( select(CloudSkillEntitlementRow).where( CloudSkillEntitlementRow.host_id == host_id ) ).all() skills = [ _cloud_skill_from_row(row) for row in ( session.get(CloudSkillRow, ent.skill_id) for ent in entitlements ) if row is not None ] skills.sort(key=lambda s: s.name_normalized) return skills 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(), planner_transport=( row.planner_transport if row.planner_transport in {"direct", "cloud"} else "direct" ), ) 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(), progress_step_index=row.progress_step_index, progress_step_status=row.progress_step_status, progress_summary=row.progress_summary, progress_updated_at=_parse_dt(row.progress_updated_at), ) 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), ) def _llm_provider_profile_from_row(row: LlmProviderProfileRow) -> Any: from cloud.llm_providers import LlmProviderProfile now = utc_now() return LlmProviderProfile( id=row.id, name=row.name, name_normalized=row.name_normalized, provider_type=row.provider_type, model=row.model, base_url=row.base_url, timeout_seconds=row.timeout_seconds, api_key_ciphertext=row.api_key_ciphertext, key_last_rotated_at=_parse_dt(row.key_last_rotated_at) or now, enabled=bool(row.enabled), revision=row.revision, created_at=_parse_dt(row.created_at) or now, updated_at=_parse_dt(row.updated_at) or now, ) def _llm_provider_settings_from_row(row: LlmProviderSettingsRow | None) -> Any: from cloud.llm_providers import LlmProviderSettings if row is None: return LlmProviderSettings(active_profile_id=None) return LlmProviderSettings( active_profile_id=row.active_profile_id, revision=row.revision, updated_at=_parse_dt(row.updated_at), ) def _cloud_skill_from_row(row: CloudSkillRow) -> Any: from cloud.skills import CloudSkill now = utc_now() return CloudSkill( id=row.id, name=row.name, name_normalized=row.name_normalized, kind=row.kind, # type: ignore[arg-type] description=row.description, tags=json.loads(row.tags_json or "[]"), content=row.content, steps_json=row.steps_json, parameters_json=row.parameters_json, revision=row.revision, created_at=_parse_dt(row.created_at) or now, updated_at=_parse_dt(row.updated_at) or now, ) def _cloud_skill_to_row(skill: Any) -> CloudSkillRow: return CloudSkillRow( id=skill.id, name=skill.name, name_normalized=skill.name_normalized, kind=skill.kind, description=skill.description, tags_json=json.dumps(list(skill.tags)), content=skill.content, steps_json=skill.steps_json, parameters_json=skill.parameters_json, revision=skill.revision, created_at=_iso(skill.created_at), updated_at=_iso(skill.updated_at), )