415 lines
14 KiB
Python
415 lines
14 KiB
Python
"""SQLite-backed store for cloud runtime state.
|
|
|
|
Owns ``cloud/cloud.sqlite3`` with four tables: ``host_registrations``,
|
|
``pooled_devices``, ``scheduled_tasks``, and ``plugins``. Uses the same
|
|
connect-per-call ``sqlite3`` pattern as ``storage/task_metadata.py`` and
|
|
``workflow/store.py``.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import sqlite3
|
|
from dataclasses import asdict
|
|
from pathlib import Path
|
|
from typing import TYPE_CHECKING, Any
|
|
|
|
from core.models import utc_now
|
|
|
|
if TYPE_CHECKING:
|
|
from cloud.pool import HostRegistration, PooledDevice
|
|
from cloud.plugins import PluginManifest
|
|
from cloud.scheduler import ScheduledTask
|
|
|
|
|
|
class CloudStore:
|
|
"""Persisted state for the cloud runtime (hosts, devices, tasks, plugins)."""
|
|
|
|
def __init__(self, db_path: str | Path = "cloud/cloud.sqlite3") -> None:
|
|
self.db_path = Path(db_path)
|
|
self.db_path.parent.mkdir(parents=True, exist_ok=True)
|
|
self._ensure_schema()
|
|
|
|
# ------------------------------------------------------------------ hosts
|
|
|
|
def upsert_host(
|
|
self,
|
|
host_id: str,
|
|
*,
|
|
address: str | None,
|
|
last_seen_at: Any,
|
|
) -> None:
|
|
"""Insert or update a host row.
|
|
|
|
If ``address`` is None and the host already exists, the existing
|
|
address is preserved (a heartbeat sync should not blow away a
|
|
previously-registered address).
|
|
"""
|
|
with self._connect() as connection:
|
|
existing = connection.execute(
|
|
"select address from host_registrations where host_id = ?",
|
|
(host_id,),
|
|
).fetchone()
|
|
preserved_address = (
|
|
existing["address"] if (address is None and existing is not None) else address
|
|
)
|
|
connection.execute(
|
|
"""
|
|
insert into host_registrations (host_id, address, last_seen_at)
|
|
values (?, ?, ?)
|
|
on conflict(host_id) do update set
|
|
address = excluded.address,
|
|
last_seen_at = excluded.last_seen_at
|
|
""",
|
|
(host_id, preserved_address, _iso(last_seen_at)),
|
|
)
|
|
|
|
def replace_host_devices(
|
|
self,
|
|
host_id: str,
|
|
devices: "list[PooledDevice]",
|
|
) -> None:
|
|
"""Atomically replace one host's device rows with the given list."""
|
|
with self._connect() as connection:
|
|
connection.execute(
|
|
"delete from pooled_devices where host_id = ?",
|
|
(host_id,),
|
|
)
|
|
connection.executemany(
|
|
"""
|
|
insert into pooled_devices (
|
|
device_id, host_id, driver_type, status,
|
|
capability_tags_json, synced_at
|
|
) values (?, ?, ?, ?, ?, ?)
|
|
""",
|
|
[
|
|
(
|
|
d.device_id,
|
|
d.host_id,
|
|
d.driver_type,
|
|
d.status,
|
|
json.dumps(list(d.capability_tags), ensure_ascii=False),
|
|
_iso(d.synced_at),
|
|
)
|
|
for d in devices
|
|
],
|
|
)
|
|
|
|
def list_hosts(self) -> "list[HostRegistration]":
|
|
from cloud.pool import HostRegistration
|
|
|
|
with self._connect() as connection:
|
|
rows = connection.execute(
|
|
"select host_id, address, last_seen_at from host_registrations"
|
|
).fetchall()
|
|
return [_row_to_host(row) for row in rows]
|
|
|
|
def get_host(self, host_id: str) -> "HostRegistration | None":
|
|
from cloud.pool import HostRegistration
|
|
|
|
with self._connect() as connection:
|
|
row = connection.execute(
|
|
"select host_id, address, last_seen_at from host_registrations where host_id = ?",
|
|
(host_id,),
|
|
).fetchone()
|
|
return _row_to_host(row) if row else None
|
|
|
|
def list_devices(self) -> "list[PooledDevice]":
|
|
with self._connect() as connection:
|
|
rows = connection.execute(
|
|
"""
|
|
select device_id, host_id, driver_type, status,
|
|
capability_tags_json, synced_at
|
|
from pooled_devices
|
|
"""
|
|
).fetchall()
|
|
return [_row_to_device(row) for row in rows]
|
|
|
|
def get_device(self, device_id: str) -> "PooledDevice | None":
|
|
with self._connect() as connection:
|
|
row = connection.execute(
|
|
"""
|
|
select device_id, host_id, driver_type, status,
|
|
capability_tags_json, synced_at
|
|
from pooled_devices where device_id = ?
|
|
""",
|
|
(device_id,),
|
|
).fetchone()
|
|
return _row_to_device(row) if row else None
|
|
|
|
# ----------------------------------------------------------- scheduled tasks
|
|
|
|
def enqueue_task(self, task: "ScheduledTask") -> None:
|
|
with self._connect() as connection:
|
|
connection.execute(
|
|
"""
|
|
insert into scheduled_tasks (
|
|
id, goal, workflow_definition_id,
|
|
constraints_json, status,
|
|
assigned_device_id, assigned_host_id, created_at
|
|
) values (?, ?, ?, ?, ?, ?, ?, ?)
|
|
""",
|
|
(
|
|
task.id,
|
|
task.goal,
|
|
task.workflow_definition_id,
|
|
json.dumps(asdict(task.constraints), ensure_ascii=False),
|
|
task.status,
|
|
task.assigned_device_id,
|
|
task.assigned_host_id,
|
|
_iso(task.created_at),
|
|
),
|
|
)
|
|
|
|
def list_queued_tasks(self) -> "list[ScheduledTask]":
|
|
with self._connect() as connection:
|
|
rows = connection.execute(
|
|
"""
|
|
select id, goal, workflow_definition_id, constraints_json,
|
|
status, assigned_device_id, assigned_host_id, created_at
|
|
from scheduled_tasks
|
|
where status = 'queued'
|
|
order by created_at asc, id asc
|
|
"""
|
|
).fetchall()
|
|
return [_row_to_task(row) for row in rows]
|
|
|
|
def get_task(self, task_id: str) -> "ScheduledTask | None":
|
|
with self._connect() as connection:
|
|
row = connection.execute(
|
|
"""
|
|
select id, goal, workflow_definition_id, constraints_json,
|
|
status, assigned_device_id, assigned_host_id, created_at
|
|
from scheduled_tasks where id = ?
|
|
""",
|
|
(task_id,),
|
|
).fetchone()
|
|
return _row_to_task(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:
|
|
updates: dict[str, Any] = {}
|
|
if status is not None:
|
|
updates["status"] = status
|
|
if assigned_device_id is not None:
|
|
updates["assigned_device_id"] = assigned_device_id
|
|
if assigned_host_id is not None:
|
|
updates["assigned_host_id"] = assigned_host_id
|
|
if not updates:
|
|
return
|
|
assignments = ", ".join(f"{key} = ?" for key in updates)
|
|
values = [*updates.values(), task_id]
|
|
with self._connect() as connection:
|
|
connection.execute(
|
|
f"update scheduled_tasks set {assignments} where id = ?",
|
|
values,
|
|
)
|
|
|
|
def count_queued_tasks(self) -> int:
|
|
with self._connect() as connection:
|
|
row = connection.execute(
|
|
"select count(*) as count from scheduled_tasks where status = 'queued'"
|
|
).fetchone()
|
|
return int(row["count"])
|
|
|
|
# ----------------------------------------------------------------- plugins
|
|
|
|
def save_plugin(
|
|
self,
|
|
manifest: "PluginManifest",
|
|
*,
|
|
wired: bool,
|
|
) -> None:
|
|
with self._connect() as connection:
|
|
connection.execute(
|
|
"""
|
|
insert into plugins (
|
|
name, version, entry_point_kind, target, wired
|
|
) values (?, ?, ?, ?, ?)
|
|
on conflict(name) do update set
|
|
version = excluded.version,
|
|
entry_point_kind = excluded.entry_point_kind,
|
|
target = excluded.target,
|
|
wired = excluded.wired
|
|
""",
|
|
(
|
|
manifest.name,
|
|
manifest.version,
|
|
manifest.entry_point_kind,
|
|
manifest.target,
|
|
1 if wired else 0,
|
|
),
|
|
)
|
|
|
|
def list_plugins(self) -> "list[tuple[PluginManifest, bool]]":
|
|
from cloud.plugins import PluginManifest
|
|
|
|
with self._connect() as connection:
|
|
rows = connection.execute(
|
|
"select name, version, entry_point_kind, target, wired from plugins"
|
|
).fetchall()
|
|
return [
|
|
(
|
|
PluginManifest(
|
|
name=row["name"],
|
|
version=row["version"],
|
|
entry_point_kind=row["entry_point_kind"],
|
|
target=row["target"],
|
|
),
|
|
bool(row["wired"]),
|
|
)
|
|
for row in rows
|
|
]
|
|
|
|
def get_plugin(self, name: str) -> "tuple[PluginManifest, bool] | None":
|
|
from cloud.plugins import PluginManifest
|
|
|
|
with self._connect() as connection:
|
|
row = connection.execute(
|
|
"select name, version, entry_point_kind, target, wired from plugins where name = ?",
|
|
(name,),
|
|
).fetchone()
|
|
if row is None:
|
|
return None
|
|
manifest = PluginManifest(
|
|
name=row["name"],
|
|
version=row["version"],
|
|
entry_point_kind=row["entry_point_kind"],
|
|
target=row["target"],
|
|
)
|
|
return (manifest, bool(row["wired"]))
|
|
|
|
# ------------------------------------------------------------------- schema
|
|
|
|
def _ensure_schema(self) -> None:
|
|
with self._connect() as connection:
|
|
connection.execute(
|
|
"""
|
|
create table if not exists host_registrations (
|
|
host_id text primary key,
|
|
address text,
|
|
last_seen_at text not null
|
|
)
|
|
"""
|
|
)
|
|
connection.execute(
|
|
"""
|
|
create table if not exists pooled_devices (
|
|
device_id text primary key,
|
|
host_id text not null,
|
|
driver_type text not null,
|
|
status text not null,
|
|
capability_tags_json text not null,
|
|
synced_at text
|
|
)
|
|
"""
|
|
)
|
|
connection.execute(
|
|
"""
|
|
create table if not exists scheduled_tasks (
|
|
id text primary key,
|
|
goal text,
|
|
workflow_definition_id text,
|
|
constraints_json text not null,
|
|
status text not null,
|
|
assigned_device_id text,
|
|
assigned_host_id text,
|
|
created_at text not null
|
|
)
|
|
"""
|
|
)
|
|
connection.execute(
|
|
"""
|
|
create table if not exists plugins (
|
|
name text primary key,
|
|
version text not null,
|
|
entry_point_kind text not null,
|
|
target text not null,
|
|
wired integer not null
|
|
)
|
|
"""
|
|
)
|
|
|
|
def _connect(self) -> sqlite3.Connection:
|
|
connection = sqlite3.connect(self.db_path)
|
|
connection.row_factory = sqlite3.Row
|
|
return connection
|
|
|
|
|
|
def _iso(value: Any) -> str:
|
|
if value is None:
|
|
return ""
|
|
if isinstance(value, str):
|
|
return value
|
|
return value.isoformat()
|
|
|
|
|
|
def _parse_dt(value: Any):
|
|
if not value:
|
|
return None
|
|
if isinstance(value, str):
|
|
from datetime import datetime
|
|
|
|
try:
|
|
return datetime.fromisoformat(value)
|
|
except ValueError:
|
|
return None
|
|
return value
|
|
|
|
|
|
def _row_to_host(row: sqlite3.Row):
|
|
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 _row_to_device(row: sqlite3.Row):
|
|
from cloud.pool import PooledDevice
|
|
|
|
tags_raw = row["capability_tags_json"]
|
|
try:
|
|
tags = list(json.loads(tags_raw)) if tags_raw else []
|
|
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 _row_to_task(row: sqlite3.Row):
|
|
from cloud.scheduler import ScheduledTask, TaskConstraints
|
|
|
|
try:
|
|
constraints_data = json.loads(row["constraints_json"]) if row["constraints_json"] else {}
|
|
except (TypeError, ValueError):
|
|
constraints_data = {}
|
|
constraints = TaskConstraints(
|
|
driver_type=constraints_data.get("driver_type"),
|
|
capability_tags=list(constraints_data.get("capability_tags") or []),
|
|
)
|
|
return ScheduledTask(
|
|
id=row["id"],
|
|
goal=row["goal"],
|
|
workflow_definition_id=row["workflow_definition_id"],
|
|
constraints=constraints,
|
|
status=row["status"],
|
|
assigned_device_id=row["assigned_device_id"],
|
|
assigned_host_id=row["assigned_host_id"],
|
|
created_at=_parse_dt(row["created_at"]) or utc_now(),
|
|
)
|