Files
agentic-mobile-control/packages/cloud-platform/cloud/sql_repository.py
T
q792602257andClaude Opus 4.6 52e442790a
Tests / Test passed: 819
feat(cloud): cloud-managed skill store + per-host entitlement + sync versioning
Adds cloud/skills.py (domain + service), SQLAlchemy models and Alembic
migration 0010_skill_management (cloud_skills, cloud_skill_entitlements,
cloud_skill_sync_state, a per-host changelog, and host_skill_inventory),
and repository methods with a monotonic per-host entitlement_version that
drives correct incremental fetch_host_delta. cloud-api suite green (41
passed); HEAD_REVISION bumped to 0010.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-07-15 07:45:22 +08:00

2527 lines
92 KiB
Python

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,
) -> 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),
)
)
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]
)
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),
)