Files
agentic-mobile-control/packages/cloud-platform/cloud/sql_repository.py
T
2026-07-13 22:56:31 +08:00

1717 lines
63 KiB
Python

from __future__ import annotations
import json
import logging
from dataclasses import asdict
from datetime import datetime, timedelta
from typing import Any
from sqlalchemy import Engine, delete, func, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import sessionmaker
from cloud.db_models import (
Base,
DeviceEnrollmentRow,
HostRow,
PluginRow,
PooledDeviceRow,
ScheduledTaskRow,
TaskAttemptRow,
AuthAuditRow,
HostGovernancePolicyRow,
LoginThrottleRow,
UserRow,
UserSessionRow,
UserSubmissionPolicyRow,
TokenReservationRow,
TokenUsageEventRow,
)
from cloud.observability import current_correlation_id
from core.models import utc_now
logger = logging.getLogger(__name__)
class SQLAlchemyCloudRepository:
"""SQLAlchemy adapter preserving the existing CloudStore CRUD surface."""
def __init__(self, engine: Engine, *, create_schema: bool = True) -> None:
self.engine = engine
self._sessions = sessionmaker(bind=engine, expire_on_commit=False)
if create_schema:
Base.metadata.create_all(engine)
def enroll_host(
self,
*,
host_id: str,
agent_instance_id: str,
credential_digest: str,
enrollment_token_digest: str | None,
display_name: str | None,
enrolled_at: datetime,
) -> Any:
from cloud.repository import (
EnrollmentTokenConflictError,
HostEnrollmentConflictError,
)
try:
with self._sessions.begin() as session:
statement = select(HostRow).where(
HostRow.agent_instance_id == agent_instance_id
)
if self.engine.dialect.name == "postgresql":
statement = statement.with_for_update()
existing = session.scalars(statement).first()
if existing is not None:
if (
existing.credential_digest != credential_digest
or existing.enrollment_token_digest != enrollment_token_digest
):
raise HostEnrollmentConflictError(
"Host enrollment identity does not match existing binding"
)
return _host_enrollment_from_row(existing)
# A NULL enrollment_token_digest marks self-service Hosts, which
# never bind a shared token; comparing `== NULL` would otherwise
# match every other self-service Host via SQL's `IS NULL` and
# falsely report a token conflict.
token_owner = (
session.scalars(
select(HostRow).where(
HostRow.enrollment_token_digest == enrollment_token_digest
)
).first()
if enrollment_token_digest is not None
else None
)
if token_owner is not None:
raise EnrollmentTokenConflictError(
"enrollment token is already bound to another Host"
)
row = HostRow(
host_id=host_id,
address=None,
last_seen_at=_iso(enrolled_at),
agent_instance_id=agent_instance_id,
credential_digest=credential_digest,
enrollment_token_digest=enrollment_token_digest,
display_name=display_name,
enrolled_at=_iso(enrolled_at),
revoked_at=None,
)
session.add(row)
session.flush()
return _host_enrollment_from_row(row)
except IntegrityError as exc:
with self._sessions() as session:
existing = session.scalars(
select(HostRow).where(
HostRow.agent_instance_id == agent_instance_id
)
).first()
if (
existing is not None
and existing.credential_digest == credential_digest
and existing.enrollment_token_digest == enrollment_token_digest
):
return _host_enrollment_from_row(existing)
token_owner = (
session.scalars(
select(HostRow).where(
HostRow.enrollment_token_digest == enrollment_token_digest
)
).first()
if enrollment_token_digest is not None
else None
)
if token_owner is not None:
raise EnrollmentTokenConflictError(
"enrollment token is already bound to another Host"
) from exc
raise HostEnrollmentConflictError(
"Host enrollment conflicts with an existing identity"
) from exc
def authenticate_enrolled_host(self, credential_digest: str) -> str | None:
with self._sessions() as session:
host_id = session.scalar(
select(HostRow.host_id)
.where(
HostRow.credential_digest == credential_digest,
HostRow.revoked_at.is_(None),
)
.limit(1)
)
return str(host_id) if host_id is not None else None
def revoke_enrolled_host(
self,
host_id: str,
*,
revoked_at: datetime,
) -> bool:
with self._sessions.begin() as session:
row = session.get(
HostRow,
host_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None or row.credential_digest is None:
return False
row.revoked_at = _iso(revoked_at)
return True
def is_enrollment_managed_host(self, host_id: str) -> bool:
with self._sessions() as session:
credential_digest = session.scalar(
select(HostRow.credential_digest).where(HostRow.host_id == host_id)
)
return credential_digest is not None
def enroll_device(
self,
*,
device_id: str,
host_id: str,
local_device_id: str,
driver_type: str,
name: str | None,
capability_tags: list[str],
enrolled_at: datetime,
) -> Any:
from cloud.repository import DeviceEnrollmentConflictError
tags_json = json.dumps(list(capability_tags), ensure_ascii=False)
try:
with self._sessions.begin() as session:
statement = select(DeviceEnrollmentRow).where(
DeviceEnrollmentRow.host_id == host_id,
DeviceEnrollmentRow.local_device_id == local_device_id,
)
if self.engine.dialect.name == "postgresql":
statement = statement.with_for_update()
existing = session.scalars(statement).first()
if existing is not None:
if existing.driver_type != driver_type:
raise DeviceEnrollmentConflictError(
"device driver type does not match existing enrollment"
)
if existing.revoked_at is not None:
raise DeviceEnrollmentConflictError(
"device enrollment has been revoked"
)
existing.name = name
existing.capability_tags_json = tags_json
return _device_enrollment_from_row(existing)
host = session.get(HostRow, host_id)
if host is None:
session.add(
HostRow(
host_id=host_id,
address=None,
last_seen_at=_iso(enrolled_at),
)
)
row = DeviceEnrollmentRow(
device_id=device_id,
host_id=host_id,
local_device_id=local_device_id,
driver_type=driver_type,
name=name,
capability_tags_json=tags_json,
enrolled_at=_iso(enrolled_at),
revoked_at=None,
)
session.add(row)
session.flush()
return _device_enrollment_from_row(row)
except IntegrityError as exc:
with self._sessions() as session:
existing = session.scalars(
select(DeviceEnrollmentRow).where(
DeviceEnrollmentRow.host_id == host_id,
DeviceEnrollmentRow.local_device_id == local_device_id,
)
).first()
if existing is not None and existing.driver_type == driver_type:
return _device_enrollment_from_row(existing)
raise DeviceEnrollmentConflictError(
"device enrollment conflicts with an existing identity"
) from exc
def get_device_enrollment(self, device_id: str) -> Any | None:
with self._sessions() as session:
row = session.get(DeviceEnrollmentRow, device_id)
return _device_enrollment_from_row(row) if row else None
def list_device_enrollments(self, host_id: str) -> list[Any]:
with self._sessions() as session:
rows = session.scalars(
select(DeviceEnrollmentRow)
.where(DeviceEnrollmentRow.host_id == host_id)
.order_by(DeviceEnrollmentRow.device_id)
).all()
return [_device_enrollment_from_row(row) for row in rows]
def upsert_host(
self,
host_id: str,
*,
address: str | None,
last_seen_at: datetime,
) -> None:
with self._sessions.begin() as session:
row = session.get(HostRow, host_id)
if row is None:
session.add(
HostRow(
host_id=host_id,
address=address,
last_seen_at=_iso(last_seen_at),
)
)
return
if address is not None:
row.address = address
row.last_seen_at = _iso(last_seen_at)
def replace_host_devices(
self,
host_id: str,
devices: list[Any],
*,
allow_device_takeover: bool = False,
) -> None:
with self._sessions.begin() as session:
session.execute(
delete(PooledDeviceRow).where(PooledDeviceRow.host_id == host_id)
)
if allow_device_takeover:
device_ids = [device.device_id for device in devices]
if device_ids:
session.execute(
delete(PooledDeviceRow).where(
PooledDeviceRow.host_id != host_id,
PooledDeviceRow.device_id.in_(device_ids),
)
)
session.add_all(
[
PooledDeviceRow(
device_id=device.device_id,
host_id=device.host_id,
driver_type=device.driver_type,
status=device.status,
capability_tags_json=json.dumps(
list(device.capability_tags),
ensure_ascii=False,
),
synced_at=_iso(device.synced_at) if device.synced_at else None,
)
for device in devices
]
)
def list_hosts(self) -> list[Any]:
with self._sessions() as session:
rows = session.scalars(select(HostRow).order_by(HostRow.host_id)).all()
return [_host_from_row(row) for row in rows]
def get_host(self, host_id: str) -> Any | None:
with self._sessions() as session:
row = session.get(HostRow, host_id)
return _host_from_row(row) if row else None
def list_devices(self) -> list[Any]:
with self._sessions() as session:
rows = session.scalars(
select(PooledDeviceRow).order_by(
PooledDeviceRow.host_id,
PooledDeviceRow.device_id,
)
).all()
return [_device_from_row(row) for row in rows]
def get_device(self, device_id: str) -> Any | None:
with self._sessions() as session:
row = session.scalars(
select(PooledDeviceRow)
.where(PooledDeviceRow.device_id == device_id)
.order_by(PooledDeviceRow.host_id)
.limit(1)
).first()
return _device_from_row(row) if row else None
def enqueue_task(self, task: Any) -> None:
with self._sessions.begin() as session:
session.add(
ScheduledTaskRow(
id=task.id,
goal=task.goal,
workflow_definition_id=task.workflow_definition_id,
constraints_json=json.dumps(
asdict(task.constraints),
ensure_ascii=False,
),
status=task.status,
assigned_device_id=task.assigned_device_id,
assigned_host_id=task.assigned_host_id,
attempt_count=task.attempt_count,
lease_id=task.lease_id,
lease_expires_at=(
_iso(task.lease_expires_at) if task.lease_expires_at else None
),
failure_reason=task.failure_reason,
result_json=(
json.dumps(task.terminal_result, ensure_ascii=False)
if task.terminal_result is not None
else None
),
updated_at=_iso(task.updated_at) if task.updated_at else None,
created_at=_iso(task.created_at),
)
)
def list_queued_tasks(self) -> list[Any]:
with self._sessions() as session:
rows = session.scalars(
select(ScheduledTaskRow)
.where(ScheduledTaskRow.status == "queued")
.order_by(ScheduledTaskRow.created_at, ScheduledTaskRow.id)
).all()
return [_task_from_row(row) for row in rows]
def list_tasks(
self,
*,
status: str | None = None,
limit: int = 50,
offset: int = 0,
) -> list[Any]:
with self._sessions() as session:
statement = select(ScheduledTaskRow)
if status is not None:
statement = statement.where(ScheduledTaskRow.status == status)
statement = (
statement.order_by(
ScheduledTaskRow.created_at.desc(),
ScheduledTaskRow.id.desc(),
)
.limit(limit)
.offset(offset)
)
rows = session.scalars(statement).all()
return [_task_from_row(row) for row in rows]
def count_tasks(self, status: str | None = None) -> int:
with self._sessions() as session:
statement = select(func.count()).select_from(ScheduledTaskRow)
if status is not None:
statement = statement.where(ScheduledTaskRow.status == status)
count = session.scalar(statement)
return int(count or 0)
def get_task(self, task_id: str) -> Any | None:
with self._sessions() as session:
row = session.get(ScheduledTaskRow, task_id)
return _task_from_row(row) if row else None
def update_task(
self,
task_id: str,
*,
status: str | None = None,
assigned_device_id: str | None = None,
assigned_host_id: str | None = None,
) -> None:
with self._sessions.begin() as session:
row = session.get(ScheduledTaskRow, task_id)
if row is None:
return
if status is not None:
row.status = status
if assigned_device_id is not None:
row.assigned_device_id = assigned_device_id
if assigned_host_id is not None:
row.assigned_host_id = assigned_host_id
def count_queued_tasks(self) -> int:
with self._sessions() as session:
count = session.scalar(
select(func.count())
.select_from(ScheduledTaskRow)
.where(ScheduledTaskRow.status == "queued")
)
return int(count or 0)
def save_plugin(self, manifest: Any, *, wired: bool) -> None:
with self._sessions.begin() as session:
row = session.get(PluginRow, manifest.name)
if row is None:
session.add(
PluginRow(
name=manifest.name,
version=manifest.version,
entry_point_kind=manifest.entry_point_kind,
target=manifest.target,
wired=1 if wired else 0,
)
)
return
row.version = manifest.version
row.entry_point_kind = manifest.entry_point_kind
row.target = manifest.target
row.wired = 1 if wired else 0
def list_plugins(self) -> list[tuple[Any, bool]]:
with self._sessions() as session:
rows = session.scalars(select(PluginRow).order_by(PluginRow.name)).all()
return [_plugin_from_row(row) for row in rows]
def get_plugin(self, name: str) -> tuple[Any, bool] | None:
with self._sessions() as session:
row = session.get(PluginRow, name)
return _plugin_from_row(row) if row else None
# --------------------------------------------------------------- user auth
def create_user(self, user: Any) -> Any:
from cloud.repository import UserConflictError
try:
with self._sessions.begin() as session:
row = UserRow(
id=user.id,
username=user.username,
username_normalized=user.username_normalized,
display_name=user.display_name,
password_hash=user.password_hash,
role=user.role,
enabled=1 if user.enabled else 0,
must_change_password=1 if user.must_change_password else 0,
authentication_version=user.authentication_version,
created_at=_iso(user.created_at),
updated_at=_iso(user.updated_at),
last_login_at=(
_iso(user.last_login_at) if user.last_login_at is not None else None
),
)
session.add(row)
session.flush()
return _user_from_row(row)
except IntegrityError as exc:
raise UserConflictError("username is already in use") from exc
def get_user(self, user_id: str) -> Any | None:
with self._sessions() as session:
row = session.get(UserRow, user_id)
return _user_from_row(row) if row else None
def get_user_by_normalized_username(self, username_normalized: str) -> Any | None:
with self._sessions() as session:
row = session.scalar(
select(UserRow)
.where(UserRow.username_normalized == username_normalized)
.limit(1)
)
return _user_from_row(row) if row else None
def list_users(self, *, limit: int, offset: int) -> list[Any]:
with self._sessions() as session:
rows = session.scalars(
select(UserRow)
.order_by(UserRow.created_at, UserRow.id)
.limit(limit)
.offset(offset)
).all()
return [_user_from_row(row) for row in rows]
def update_user(
self,
user_id: str,
*,
display_name: str | None = None,
role: str | None = None,
enabled: bool | None = None,
updated_at: datetime,
) -> Any:
from cloud.repository import LastAdministratorConflictError
with self._sessions.begin() as session:
row = session.get(
UserRow,
user_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None:
raise KeyError(f"user {user_id!r} not found")
next_role = role if role is not None else row.role
next_enabled = enabled if enabled is not None else bool(row.enabled)
removes_administrator = (
bool(row.enabled)
and row.role == "admin"
and (next_role != "admin" or not next_enabled)
)
if removes_administrator:
other_admins = session.scalar(
select(func.count())
.select_from(UserRow)
.where(
UserRow.id != user_id,
UserRow.enabled == 1,
UserRow.role == "admin",
)
)
if int(other_admins or 0) == 0:
raise LastAdministratorConflictError(
"cannot remove the last enabled administrator"
)
security_changed = next_role != row.role or next_enabled != bool(row.enabled)
if display_name is not None:
row.display_name = display_name
row.role = next_role
row.enabled = 1 if next_enabled else 0
row.updated_at = _iso(updated_at)
if security_changed:
row.authentication_version += 1
_revoke_user_session_rows(session, user_id, updated_at)
session.flush()
return _user_from_row(row)
def rehash_user_password(
self,
user_id: str,
*,
password_hash: str,
updated_at: datetime,
) -> Any:
with self._sessions.begin() as session:
row = session.get(UserRow, user_id)
if row is None:
raise KeyError(f"user {user_id!r} not found")
row.password_hash = password_hash
row.updated_at = _iso(updated_at)
session.flush()
return _user_from_row(row)
def update_user_password(
self,
user_id: str,
*,
password_hash: str,
must_change_password: bool,
updated_at: datetime,
revoke_sessions: bool,
) -> Any:
with self._sessions.begin() as session:
row = session.get(
UserRow,
user_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None:
raise KeyError(f"user {user_id!r} not found")
row.password_hash = password_hash
row.must_change_password = 1 if must_change_password else 0
row.authentication_version += 1
row.updated_at = _iso(updated_at)
if revoke_sessions:
_revoke_user_session_rows(session, user_id, updated_at)
session.flush()
return _user_from_row(row)
def mark_user_login(self, user_id: str, *, now: datetime) -> Any:
with self._sessions.begin() as session:
row = session.get(UserRow, user_id)
if row is None:
raise KeyError(f"user {user_id!r} not found")
row.last_login_at = _iso(now)
row.updated_at = _iso(now)
session.flush()
return _user_from_row(row)
def create_user_session(self, user_session: Any) -> None:
with self._sessions.begin() as session:
session.add(
UserSessionRow(
id=user_session.id,
user_id=user_session.user_id,
token_digest=user_session.token_digest,
csrf_digest=user_session.csrf_digest,
authentication_version=user_session.authentication_version,
issued_at=_iso(user_session.issued_at),
last_seen_at=_iso(user_session.last_seen_at),
idle_expires_at=_iso(user_session.idle_expires_at),
absolute_expires_at=_iso(user_session.absolute_expires_at),
revoked_at=(
_iso(user_session.revoked_at)
if user_session.revoked_at is not None
else None
),
)
)
def get_authenticated_user_session(
self,
token_digest: str,
*,
now: datetime,
) -> Any | None:
from cloud.user_auth import AuthenticatedUserSession
with self._sessions() as session:
match = session.execute(
select(UserSessionRow, UserRow)
.join(UserRow, UserRow.id == UserSessionRow.user_id)
.where(
UserSessionRow.token_digest == token_digest,
UserSessionRow.revoked_at.is_(None),
UserRow.enabled == 1,
)
.limit(1)
).first()
if match is None:
return None
session_row, user_row = match
user = _user_from_row(user_row)
user_session = _user_session_from_row(session_row)
if (
user.authentication_version != user_session.authentication_version
or user_session.idle_expires_at <= now
or user_session.absolute_expires_at <= now
):
return None
return AuthenticatedUserSession(user=user, session=user_session)
def touch_user_session(
self,
session_id: str,
*,
last_seen_at: datetime,
idle_expires_at: datetime,
) -> Any:
with self._sessions.begin() as session:
row = session.get(UserSessionRow, session_id)
if row is None:
raise KeyError(f"session {session_id!r} not found")
row.last_seen_at = _iso(last_seen_at)
row.idle_expires_at = _iso(idle_expires_at)
session.flush()
return _user_session_from_row(row)
def revoke_user_session(self, session_id: str, *, revoked_at: datetime) -> bool:
with self._sessions.begin() as session:
row = session.get(UserSessionRow, session_id)
if row is None or row.revoked_at is not None:
return False
row.revoked_at = _iso(revoked_at)
return True
def revoke_user_sessions(self, user_id: str, *, revoked_at: datetime) -> int:
with self._sessions.begin() as session:
return _revoke_user_session_rows(session, user_id, revoked_at)
def get_login_throttle(
self,
username_normalized: str,
client_bucket: str,
) -> Any | None:
with self._sessions() as session:
row = session.get(LoginThrottleRow, (username_normalized, client_bucket))
return _login_throttle_from_row(row) if row else None
def record_login_failure(
self,
*,
username_normalized: str,
client_bucket: str,
now: datetime,
failure_limit: int,
failure_window: timedelta,
block_duration: timedelta,
) -> Any:
with self._sessions.begin() as session:
row = session.get(LoginThrottleRow, (username_normalized, client_bucket))
if row is None:
row = LoginThrottleRow(
username_normalized=username_normalized,
client_bucket=client_bucket,
failure_count=0,
window_started_at=_iso(now),
last_attempt_at=_iso(now),
blocked_until=None,
)
session.add(row)
window_started = _parse_dt(row.window_started_at) or now
if now - window_started > failure_window:
row.failure_count = 0
row.window_started_at = _iso(now)
row.blocked_until = None
row.failure_count += 1
row.last_attempt_at = _iso(now)
if row.failure_count >= failure_limit:
row.blocked_until = _iso(now + block_duration)
session.flush()
return _login_throttle_from_row(row)
def clear_login_throttle(self, username_normalized: str, client_bucket: str) -> None:
with self._sessions.begin() as session:
row = session.get(LoginThrottleRow, (username_normalized, client_bucket))
if row is not None:
session.delete(row)
def record_auth_audit(self, event: Any) -> None:
with self._sessions.begin() as session:
session.add(
AuthAuditRow(
id=event.id,
occurred_at=_iso(event.occurred_at),
actor_principal_id=event.actor_principal_id,
target_user_id=event.target_user_id,
action=event.action,
outcome=event.outcome,
correlation_id=event.correlation_id,
metadata_json=json.dumps(event.metadata, ensure_ascii=False),
)
)
def cleanup_auth_state(self, *, now: datetime, limit: int) -> int:
removed = 0
with self._sessions.begin() as session:
expired_sessions = session.scalars(
select(UserSessionRow)
.where(
(UserSessionRow.idle_expires_at <= _iso(now))
| (UserSessionRow.absolute_expires_at <= _iso(now))
)
.order_by(UserSessionRow.absolute_expires_at)
.limit(limit)
).all()
for row in expired_sessions:
session.delete(row)
removed += len(expired_sessions)
remaining = max(0, limit - removed)
if remaining:
stale_before = now - timedelta(days=1)
stale_throttles = session.scalars(
select(LoginThrottleRow)
.where(LoginThrottleRow.last_attempt_at <= _iso(stale_before))
.order_by(LoginThrottleRow.last_attempt_at)
.limit(remaining)
).all()
for row in stale_throttles:
session.delete(row)
removed += len(stale_throttles)
return removed
# -------------------------------------------------------------- governance
def get_user_submission_policy(self, user_id: str) -> Any | None:
with self._sessions() as session:
row = session.get(UserSubmissionPolicyRow, user_id)
return _user_submission_policy_from_row(row) if row is not None else None
def upsert_user_submission_policy(
self,
*,
user_id: str,
submission_enabled: bool,
allowed_host_ids: tuple[str, ...] | None,
allowed_device_targets: tuple[tuple[str, str], ...] | None,
updated_at: datetime,
expected_revision: int | None = None,
) -> Any:
from cloud.governance import GovernancePolicyConflictError
with self._sessions.begin() as session:
if session.get(UserRow, user_id) is None:
raise KeyError(f"unknown user {user_id!r}")
row = session.get(
UserSubmissionPolicyRow,
user_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None:
if expected_revision not in {None, 0}:
raise GovernancePolicyConflictError("submission policy revision changed")
row = UserSubmissionPolicyRow(
user_id=user_id,
revision=1,
submission_enabled=1 if submission_enabled else 0,
allowed_host_ids_json=_dump_optional_list(allowed_host_ids),
allowed_device_targets_json=_dump_optional_list(
allowed_device_targets
),
updated_at=_iso(updated_at),
)
session.add(row)
else:
if (
expected_revision is not None
and expected_revision != row.revision
):
raise GovernancePolicyConflictError("submission policy revision changed")
row.revision += 1
row.submission_enabled = 1 if submission_enabled else 0
row.allowed_host_ids_json = _dump_optional_list(allowed_host_ids)
row.allowed_device_targets_json = _dump_optional_list(
allowed_device_targets
)
row.updated_at = _iso(updated_at)
session.flush()
return _user_submission_policy_from_row(row)
def get_host_governance_policy(self, host_id: str) -> Any | None:
with self._sessions() as session:
row = session.get(HostGovernancePolicyRow, host_id)
return _host_governance_policy_from_row(row) if row is not None else None
def upsert_host_governance_policy(
self,
*,
host_id: str,
self_submission_enabled: bool,
max_active_tasks: int | None,
daily_token_budget: int | None,
updated_at: datetime,
expected_revision: int | None = None,
) -> Any:
from cloud.governance import GovernancePolicyConflictError
with self._sessions.begin() as session:
if session.get(HostRow, host_id) is None:
raise KeyError(f"unknown host {host_id!r}")
row = session.get(
HostGovernancePolicyRow,
host_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None:
if expected_revision not in {None, 0}:
raise GovernancePolicyConflictError("Host policy revision changed")
row = HostGovernancePolicyRow(
host_id=host_id,
revision=1,
self_submission_enabled=1 if self_submission_enabled else 0,
max_active_tasks=max_active_tasks,
daily_token_budget=daily_token_budget,
updated_at=_iso(updated_at),
)
session.add(row)
else:
if (
expected_revision is not None
and expected_revision != row.revision
):
raise GovernancePolicyConflictError("Host policy revision changed")
row.revision += 1
row.self_submission_enabled = 1 if self_submission_enabled else 0
row.max_active_tasks = max_active_tasks
row.daily_token_budget = daily_token_budget
row.updated_at = _iso(updated_at)
session.flush()
return _host_governance_policy_from_row(row)
def count_active_tasks_for_host(self, host_id: str) -> int:
with self._sessions() as session:
count = session.scalar(
select(func.count())
.select_from(ScheduledTaskRow)
.where(
ScheduledTaskRow.assigned_host_id == host_id,
ScheduledTaskRow.status.in_(("assigned", "dispatched")),
)
)
return int(count or 0)
def reserve_host_token_budget(
self,
*,
reservation_id: str,
host_id: str,
usage_day: str,
reserved_tokens: int,
task_id: str | None,
attempt: int | None,
created_at: datetime,
expires_at: datetime,
) -> Any | None:
from cloud.governance import TokenBudgetExceededError
with self._sessions.begin() as session:
statement = select(HostGovernancePolicyRow).where(
HostGovernancePolicyRow.host_id == host_id
)
if self.engine.dialect.name == "postgresql":
statement = statement.with_for_update()
policy = session.scalars(statement).first()
if policy is None or policy.daily_token_budget is None:
return None
used = session.scalar(
select(func.coalesce(func.sum(TokenUsageEventRow.total_tokens), 0)).where(
TokenUsageEventRow.host_id == host_id,
TokenUsageEventRow.usage_day == usage_day,
)
)
reserved = session.scalar(
select(func.coalesce(func.sum(TokenReservationRow.reserved_tokens), 0)).where(
TokenReservationRow.host_id == host_id,
TokenReservationRow.usage_day == usage_day,
TokenReservationRow.expires_at > _iso(created_at),
)
)
if int(used or 0) + int(reserved or 0) + reserved_tokens > policy.daily_token_budget:
raise TokenBudgetExceededError("Host daily token budget is exhausted")
row = TokenReservationRow(
id=reservation_id, host_id=host_id, usage_day=usage_day,
reserved_tokens=reserved_tokens, task_id=task_id, attempt=attempt,
created_at=_iso(created_at), expires_at=_iso(expires_at),
)
session.add(row)
session.flush()
return _token_reservation_from_row(row)
def settle_host_token_reservation(
self,
*,
reservation_id: str,
event_id: str,
provider: str,
model: str,
input_tokens: int | None,
output_tokens: int | None,
total_tokens: int,
occurred_at: datetime,
) -> Any | None:
with self._sessions.begin() as session:
row = session.get(
TokenReservationRow, reservation_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if row is None:
return None
event = TokenUsageEventRow(
id=event_id, host_id=row.host_id, usage_day=row.usage_day,
task_id=row.task_id, attempt=row.attempt, provider=provider, model=model,
input_tokens=input_tokens, output_tokens=output_tokens,
total_tokens=total_tokens, occurred_at=_iso(occurred_at),
)
session.add(event)
session.delete(row)
session.flush()
return _token_usage_event_from_row(event)
def cleanup_expired_token_reservations(self, *, now: datetime, limit: int) -> int:
with self._sessions.begin() as session:
rows = session.scalars(
select(TokenReservationRow)
.where(TokenReservationRow.expires_at <= _iso(now))
.order_by(TokenReservationRow.expires_at)
.limit(limit)
).all()
for row in rows:
session.delete(row)
return len(rows)
def get_host_token_usage_summary(
self, *, host_id: str, usage_day: str, now: datetime,
) -> Any:
from cloud.governance import TokenUsageSummary
with self._sessions() as session:
policy = session.get(HostGovernancePolicyRow, host_id)
used = session.scalar(
select(func.coalesce(func.sum(TokenUsageEventRow.total_tokens), 0)).where(
TokenUsageEventRow.host_id == host_id,
TokenUsageEventRow.usage_day == usage_day,
)
)
reserved = session.scalar(
select(func.coalesce(func.sum(TokenReservationRow.reserved_tokens), 0)).where(
TokenReservationRow.host_id == host_id,
TokenReservationRow.usage_day == usage_day,
TokenReservationRow.expires_at > _iso(now),
)
)
return TokenUsageSummary(
host_id=host_id, usage_day=usage_day,
daily_token_budget=(policy.daily_token_budget if policy else None),
used_tokens=int(used or 0), reserved_tokens=int(reserved or 0),
)
def list_reserved_device_ids(self, *, now: datetime) -> set[str]:
with self._sessions() as session:
device_ids = session.scalars(
select(ScheduledTaskRow.assigned_device_id).where(
ScheduledTaskRow.status.in_(("assigned", "dispatched")),
ScheduledTaskRow.assigned_device_id.is_not(None),
ScheduledTaskRow.lease_expires_at.is_not(None),
ScheduledTaskRow.lease_expires_at > _iso(now),
)
).all()
return {device_id for device_id in device_ids if device_id is not None}
def assign_task(
self,
*,
task_id: str,
host_id: str,
device_id: str,
lease_id: str,
lease_expires_at: datetime,
now: datetime,
) -> Any | None:
with self._sessions.begin() as session:
task_statement = select(ScheduledTaskRow).where(
ScheduledTaskRow.id == task_id
)
device_statement = select(PooledDeviceRow).where(
PooledDeviceRow.host_id == host_id,
PooledDeviceRow.device_id == device_id,
)
if self.engine.dialect.name == "postgresql":
task_statement = task_statement.with_for_update(skip_locked=True)
device_statement = device_statement.with_for_update()
task = session.scalars(task_statement).first()
if task is None or task.status != "queued":
return None
device = session.scalars(device_statement).first()
if device is None or device.status != "idle":
return None
policy_statement = select(HostGovernancePolicyRow).where(
HostGovernancePolicyRow.host_id == host_id
)
if self.engine.dialect.name == "postgresql":
policy_statement = policy_statement.with_for_update()
policy = session.scalars(policy_statement).first()
if policy is not None and policy.max_active_tasks is not None:
active_count = session.scalar(
select(func.count())
.select_from(ScheduledTaskRow)
.where(
ScheduledTaskRow.assigned_host_id == host_id,
ScheduledTaskRow.status.in_(("assigned", "dispatched")),
)
)
if int(active_count or 0) >= policy.max_active_tasks:
return None
active_reservation = session.scalar(
select(ScheduledTaskRow.id)
.where(
ScheduledTaskRow.id != task_id,
ScheduledTaskRow.assigned_device_id == device_id,
ScheduledTaskRow.status.in_(("assigned", "dispatched")),
ScheduledTaskRow.lease_expires_at.is_not(None),
ScheduledTaskRow.lease_expires_at > _iso(now),
)
.limit(1)
)
if active_reservation is not None:
return None
attempt = task.attempt_count + 1
task.status = "assigned"
task.assigned_host_id = host_id
task.assigned_device_id = device_id
task.attempt_count = attempt
task.lease_id = lease_id
task.lease_expires_at = _iso(lease_expires_at)
task.failure_reason = None
task.result_json = None
task.updated_at = _iso(now)
session.add(
TaskAttemptRow(
task_id=task.id,
attempt=attempt,
lease_id=lease_id,
host_id=host_id,
device_id=device_id,
status="assigned",
lease_expires_at=_iso(lease_expires_at),
created_at=_iso(now),
completed_at=None,
failure_reason=None,
result_json=None,
)
)
session.flush()
_log_task_lifecycle("assigned", task)
return _leased_assignment_from_row(task)
def claim_assignment(
self,
*,
host_id: str,
now: datetime,
) -> Any | None:
with self._sessions.begin() as session:
statement = (
select(ScheduledTaskRow)
.where(
ScheduledTaskRow.status == "assigned",
ScheduledTaskRow.assigned_host_id == host_id,
ScheduledTaskRow.lease_id.is_not(None),
ScheduledTaskRow.lease_expires_at.is_not(None),
ScheduledTaskRow.lease_expires_at > _iso(now),
)
.order_by(ScheduledTaskRow.created_at, ScheduledTaskRow.id)
.limit(1)
)
if self.engine.dialect.name == "postgresql":
statement = statement.with_for_update(skip_locked=True)
task = session.scalars(statement).first()
if task is None:
return None
attempt = session.get(
TaskAttemptRow,
(task.id, task.attempt_count),
with_for_update=self.engine.dialect.name == "postgresql",
)
if (
attempt is None
or attempt.status != "assigned"
or attempt.lease_id != task.lease_id
):
return None
task.status = "dispatched"
task.updated_at = _iso(now)
attempt.status = "dispatched"
session.flush()
_log_task_lifecycle("claimed", task)
return _leased_assignment_from_row(task)
def renew_lease(
self,
*,
task_id: str,
attempt: int,
lease_id: str,
host_id: str,
lease_expires_at: datetime,
now: datetime,
) -> str:
with self._sessions.begin() as session:
task = session.get(
ScheduledTaskRow,
task_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if task is None:
return "not_found"
if (
task.status not in {"assigned", "dispatched"}
or task.attempt_count != attempt
or task.lease_id != lease_id
or task.assigned_host_id != host_id
):
return "conflict"
current_expiry = _parse_dt(task.lease_expires_at)
if current_expiry is None or current_expiry <= now:
return "expired"
if lease_expires_at <= now:
return "conflict"
attempt_row = session.get(
TaskAttemptRow,
(task_id, attempt),
with_for_update=self.engine.dialect.name == "postgresql",
)
if (
attempt_row is None
or attempt_row.status not in {"assigned", "dispatched"}
or attempt_row.lease_id != lease_id
or attempt_row.host_id != host_id
):
return "conflict"
renewed_until = _iso(lease_expires_at)
task.lease_expires_at = renewed_until
task.updated_at = _iso(now)
attempt_row.lease_expires_at = renewed_until
_log_task_lifecycle("renewed", task)
return "renewed"
def record_task_result(
self,
*,
task_id: str,
attempt: int,
lease_id: str,
host_id: str,
status: str,
failure_reason: str | None,
terminal_result: dict[str, Any] | None,
completed_at: datetime,
) -> str:
with self._sessions.begin() as session:
task = session.get(
ScheduledTaskRow,
task_id,
with_for_update=self.engine.dialect.name == "postgresql",
)
if task is None:
return "conflict"
attempt_row = session.get(
TaskAttemptRow,
(task_id, attempt),
with_for_update=self.engine.dialect.name == "postgresql",
)
if (
attempt_row is None
or task.attempt_count != attempt
or task.lease_id != lease_id
or task.assigned_host_id != host_id
or attempt_row.lease_id != lease_id
or attempt_row.host_id != host_id
):
return "conflict"
if task.status in {"done", "failed"}:
if (
task.status == status
and task.failure_reason == failure_reason
and _parse_json_object(task.result_json) == terminal_result
and attempt_row.status == status
):
return "already_recorded"
return "conflict"
if task.status not in {"assigned", "dispatched"}:
return "conflict"
current_expiry = _parse_dt(task.lease_expires_at)
if current_expiry is None or current_expiry <= completed_at:
return "conflict"
if status not in {"done", "failed"}:
return "conflict"
result_json = (
json.dumps(terminal_result, ensure_ascii=False)
if terminal_result is not None
else None
)
completed_at_iso = _iso(completed_at)
task.status = status
task.failure_reason = failure_reason
task.result_json = result_json
task.updated_at = completed_at_iso
attempt_row.status = status
attempt_row.completed_at = completed_at_iso
attempt_row.failure_reason = failure_reason
attempt_row.result_json = result_json
_log_task_lifecycle("completed" if status == "done" else "failed", task)
return "recorded"
def reap_expired_leases(
self,
*,
now: datetime,
max_attempts: int,
) -> list[str]:
if max_attempts < 1:
raise ValueError("max_attempts must be at least 1")
with self._sessions.begin() as session:
statement = (
select(ScheduledTaskRow)
.where(
ScheduledTaskRow.status.in_(("assigned", "dispatched")),
ScheduledTaskRow.lease_expires_at.is_not(None),
ScheduledTaskRow.lease_expires_at <= _iso(now),
)
.order_by(ScheduledTaskRow.lease_expires_at, ScheduledTaskRow.id)
)
if self.engine.dialect.name == "postgresql":
statement = statement.with_for_update(skip_locked=True)
tasks = session.scalars(statement).all()
reaped_task_ids: list[str] = []
for task in tasks:
attempt_row = session.get(
TaskAttemptRow,
(task.id, task.attempt_count),
with_for_update=self.engine.dialect.name == "postgresql",
)
if attempt_row is None:
continue
attempt_row.status = "expired"
attempt_row.completed_at = _iso(now)
attempt_row.failure_reason = "lease expired"
task.updated_at = _iso(now)
task.lease_id = None
task.lease_expires_at = None
task.result_json = None
if task.attempt_count < max_attempts:
task.status = "queued"
task.assigned_host_id = None
task.assigned_device_id = None
task.failure_reason = None
else:
task.status = "failed"
task.failure_reason = (
f"lease expired after {task.attempt_count} attempts"
)
_log_task_lifecycle(
"retried" if task.status == "queued" else "failed",
task,
)
reaped_task_ids.append(task.id)
return reaped_task_ids
def list_task_attempts(self, task_id: str) -> list[Any]:
with self._sessions() as session:
rows = session.scalars(
select(TaskAttemptRow)
.where(TaskAttemptRow.task_id == task_id)
.order_by(TaskAttemptRow.attempt)
).all()
return [_task_attempt_from_row(row) for row in rows]
def health_check(self) -> None:
with self._sessions() as session:
session.execute(select(1))
def close(self) -> None:
self.engine.dispose()
def _iso(value: datetime) -> str:
return value.isoformat()
def _parse_dt(value: str | None) -> datetime | None:
if not value:
return None
try:
return datetime.fromisoformat(value)
except ValueError:
return None
def _host_from_row(row: HostRow) -> Any:
from cloud.pool import HostRegistration
return HostRegistration(
host_id=row.host_id,
address=row.address,
last_seen_at=_parse_dt(row.last_seen_at) or utc_now(),
)
def _host_enrollment_from_row(row: HostRow) -> Any:
from cloud.repository import HostEnrollment
enrolled_at = _parse_dt(row.enrolled_at) or _parse_dt(row.last_seen_at) or utc_now()
if row.agent_instance_id is None:
raise ValueError(f"Host {row.host_id!r} is not enrollment-managed")
return HostEnrollment(
host_id=row.host_id,
agent_instance_id=row.agent_instance_id,
display_name=row.display_name,
enrolled_at=enrolled_at,
revoked_at=_parse_dt(row.revoked_at),
)
def _device_enrollment_from_row(row: DeviceEnrollmentRow) -> Any:
from cloud.repository import DeviceEnrollment
try:
tags = list(json.loads(row.capability_tags_json))
except TypeError, ValueError:
tags = []
return DeviceEnrollment(
device_id=row.device_id,
host_id=row.host_id,
local_device_id=row.local_device_id,
driver_type=row.driver_type,
name=row.name,
capability_tags=tags,
enrolled_at=_parse_dt(row.enrolled_at) or utc_now(),
revoked_at=_parse_dt(row.revoked_at),
)
def _device_from_row(row: PooledDeviceRow) -> Any:
from cloud.pool import PooledDevice
try:
tags = list(json.loads(row.capability_tags_json))
except TypeError, ValueError:
tags = []
return PooledDevice(
device_id=row.device_id,
host_id=row.host_id,
driver_type=row.driver_type,
status=row.status,
capability_tags=tags,
synced_at=_parse_dt(row.synced_at),
)
def _task_from_row(row: ScheduledTaskRow) -> Any:
from cloud.scheduler import ScheduledTask, TaskConstraints
try:
constraints_data = json.loads(row.constraints_json)
except TypeError, ValueError:
constraints_data = {}
terminal_result = _parse_json_object(row.result_json)
return ScheduledTask(
id=row.id,
goal=row.goal,
workflow_definition_id=row.workflow_definition_id,
constraints=TaskConstraints(
driver_type=constraints_data.get("driver_type"),
capability_tags=list(constraints_data.get("capability_tags") or []),
target_host_id=constraints_data.get("target_host_id"),
target_device_id=constraints_data.get("target_device_id"),
),
status=row.status,
assigned_device_id=row.assigned_device_id,
assigned_host_id=row.assigned_host_id,
attempt_count=row.attempt_count,
lease_id=row.lease_id,
lease_expires_at=_parse_dt(row.lease_expires_at),
terminal_result=terminal_result,
failure_reason=row.failure_reason,
updated_at=_parse_dt(row.updated_at),
created_at=_parse_dt(row.created_at) or utc_now(),
)
def _dump_optional_list(value: tuple[Any, ...] | None) -> str | None:
return json.dumps(value, ensure_ascii=False) if value is not None else None
def _load_optional_string_tuple(value: str | None) -> tuple[str, ...] | None:
if value is None:
return None
try:
parsed = json.loads(value)
except (TypeError, ValueError):
parsed = []
return tuple(item for item in parsed if isinstance(item, str))
def _load_optional_device_targets(
value: str | None,
) -> tuple[tuple[str, str], ...] | None:
if value is None:
return None
try:
parsed = json.loads(value)
except (TypeError, ValueError):
parsed = []
return tuple(
(item[0], item[1])
for item in parsed
if isinstance(item, list | tuple)
and len(item) == 2
and all(isinstance(part, str) for part in item)
)
def _user_submission_policy_from_row(row: UserSubmissionPolicyRow) -> Any:
from cloud.governance import UserSubmissionPolicy
return UserSubmissionPolicy(
user_id=row.user_id,
revision=row.revision,
submission_enabled=bool(row.submission_enabled),
allowed_host_ids=_load_optional_string_tuple(row.allowed_host_ids_json),
allowed_device_targets=_load_optional_device_targets(
row.allowed_device_targets_json
),
updated_at=_parse_dt(row.updated_at) or utc_now(),
)
def _host_governance_policy_from_row(row: HostGovernancePolicyRow) -> Any:
from cloud.governance import HostGovernancePolicy
return HostGovernancePolicy(
host_id=row.host_id,
revision=row.revision,
self_submission_enabled=bool(row.self_submission_enabled),
max_active_tasks=row.max_active_tasks,
daily_token_budget=row.daily_token_budget,
updated_at=_parse_dt(row.updated_at) or utc_now(),
)
def _token_reservation_from_row(row: TokenReservationRow) -> Any:
from cloud.governance import TokenReservation
return TokenReservation(
id=row.id, host_id=row.host_id, usage_day=row.usage_day,
reserved_tokens=row.reserved_tokens, task_id=row.task_id, attempt=row.attempt,
created_at=_parse_dt(row.created_at) or utc_now(),
expires_at=_parse_dt(row.expires_at) or utc_now(),
)
def _token_usage_event_from_row(row: TokenUsageEventRow) -> Any:
from cloud.governance import TokenUsageEvent
return TokenUsageEvent(
id=row.id, host_id=row.host_id, usage_day=row.usage_day,
task_id=row.task_id, attempt=row.attempt, provider=row.provider, model=row.model,
input_tokens=row.input_tokens, output_tokens=row.output_tokens,
total_tokens=row.total_tokens, occurred_at=_parse_dt(row.occurred_at) or utc_now(),
)
def _task_attempt_from_row(row: TaskAttemptRow) -> Any:
from cloud.repository import TaskAttemptRecord
return TaskAttemptRecord(
task_id=row.task_id,
attempt=row.attempt,
lease_id=row.lease_id,
host_id=row.host_id,
device_id=row.device_id,
status=row.status,
lease_expires_at=_parse_dt(row.lease_expires_at) or utc_now(),
created_at=_parse_dt(row.created_at) or utc_now(),
completed_at=_parse_dt(row.completed_at),
failure_reason=row.failure_reason,
terminal_result=_parse_json_object(row.result_json),
)
def _leased_assignment_from_row(row: ScheduledTaskRow) -> Any:
from cloud.repository import LeasedAssignment
lease_expires_at = _parse_dt(row.lease_expires_at)
if (
row.lease_id is None
or lease_expires_at is None
or row.assigned_host_id is None
or row.assigned_device_id is None
):
raise ValueError(f"task {row.id!r} does not contain a complete lease")
return LeasedAssignment(
task_id=row.id,
attempt=row.attempt_count,
lease_id=row.lease_id,
lease_expires_at=lease_expires_at,
host_id=row.assigned_host_id,
device_id=row.assigned_device_id,
goal=row.goal,
workflow_definition_id=row.workflow_definition_id,
)
def _parse_json_object(value: str | None) -> dict[str, Any] | None:
if value is None:
return None
try:
parsed = json.loads(value)
except TypeError, ValueError:
return None
return parsed if isinstance(parsed, dict) else None
def _log_task_lifecycle(event: str, task: ScheduledTaskRow) -> None:
logger.info(
"cloud task lifecycle",
extra={
"event": event,
"correlation_id": current_correlation_id(),
"task_id": task.id,
"host_id": task.assigned_host_id,
"device_id": task.assigned_device_id,
"attempt": task.attempt_count,
"lease_id": task.lease_id,
},
)
def _revoke_user_session_rows(session: Any, user_id: str, revoked_at: datetime) -> int:
rows = session.scalars(
select(UserSessionRow).where(
UserSessionRow.user_id == user_id,
UserSessionRow.revoked_at.is_(None),
)
).all()
for row in rows:
row.revoked_at = _iso(revoked_at)
return len(rows)
def _user_from_row(row: UserRow) -> Any:
from cloud.user_auth import UserAccount
return UserAccount(
id=row.id,
username=row.username,
username_normalized=row.username_normalized,
display_name=row.display_name,
role=row.role,
enabled=bool(row.enabled),
must_change_password=bool(row.must_change_password),
authentication_version=row.authentication_version,
created_at=_parse_dt(row.created_at) or utc_now(),
updated_at=_parse_dt(row.updated_at) or utc_now(),
last_login_at=_parse_dt(row.last_login_at),
password_hash=row.password_hash,
)
def _user_session_from_row(row: UserSessionRow) -> Any:
from cloud.user_auth import UserSession
now = utc_now()
return UserSession(
id=row.id,
user_id=row.user_id,
token_digest=row.token_digest,
csrf_digest=row.csrf_digest,
authentication_version=row.authentication_version,
issued_at=_parse_dt(row.issued_at) or now,
last_seen_at=_parse_dt(row.last_seen_at) or now,
idle_expires_at=_parse_dt(row.idle_expires_at) or now,
absolute_expires_at=_parse_dt(row.absolute_expires_at) or now,
revoked_at=_parse_dt(row.revoked_at),
)
def _login_throttle_from_row(row: LoginThrottleRow) -> Any:
from cloud.user_auth import LoginThrottle
now = utc_now()
return LoginThrottle(
username_normalized=row.username_normalized,
client_bucket=row.client_bucket,
failure_count=row.failure_count,
window_started_at=_parse_dt(row.window_started_at) or now,
last_attempt_at=_parse_dt(row.last_attempt_at) or now,
blocked_until=_parse_dt(row.blocked_until),
)
def _plugin_from_row(row: PluginRow) -> tuple[Any, bool]:
from cloud.plugins import PluginManifest
return (
PluginManifest(
name=row.name,
version=row.version,
entry_point_kind=row.entry_point_kind,
target=row.target,
),
bool(row.wired),
)