FluxIP/app/fleet/repository.py

1068 lines
41 KiB
Python

from __future__ import annotations
import sqlite3
import uuid
from collections.abc import Mapping, Sequence
from dataclasses import asdict, dataclass
from datetime import timedelta
from typing import Any
from app.core.time import to_iso, utc_now
from app.database.database import Database
@dataclass(slots=True)
class ManagedInstanceRecord:
id: str
aws_account_id: str | None
cloudflare_account_id: str | None
display_name: str
aws_region: str
lightsail_instance_name: str
cloudflare_zone_name: str
cloudflare_zone_id: str
cloudflare_record_name: str
socks_port: int
proxy_health_check: bool
health_timeout_seconds: int
release_grace_seconds: int
enabled: bool
config_version: int
last_known_ip: str | None
last_checked_at: str | None
created_at: str
updated_at: str
archived_at: str | None
def to_dict(self) -> dict[str, Any]:
return asdict(self)
@dataclass(slots=True)
class InstanceGroupRecord:
id: str
name: str
enabled: bool
interval_minutes: int
next_run_at: str | None
last_run_at: str | None
config_version: int
created_at: str
updated_at: str
archived_at: str | None
member_ids: list[str]
def to_dict(self) -> dict[str, Any]:
return asdict(self)
@dataclass(slots=True)
class FleetRunRecord:
id: str
target_type: str
target_id: str
target_name: str
trigger: str
status: str
active_slot: int | None
current_item_id: str | None
lease_owner: str | None
lease_until: str | None
error_code: str | None
error_message: str | None
total_items: int
succeeded_items: int
started_at: str
updated_at: str
finished_at: str | None
@dataclass(slots=True)
class FleetRunItemRecord:
id: str
run_id: str
instance_id: str
position: int
status: str
stage: str
stage_started_at: str
attempt_count: int
config_version: int
instance_display_name: str
aws_account_id: str | None
cloudflare_account_id: str | None
aws_region: str
lightsail_instance_name: str
cloudflare_zone_name: str
cloudflare_zone_id: str
cloudflare_record_name: str
socks_port: int
proxy_health_check: bool
health_timeout_seconds: int
release_grace_seconds: int
old_static_ip_name: str | None
old_ip: str | None
new_static_ip_name: str | None
new_ip: str | None
dns_ip_before: str | None
dns_ip_after: str | None
grace_until: str | None
aws_operation_id: str | None
error_code: str | None
error_message: str | None
rollback_from_stage: str | None
rollback_reason_code: str | None
rollback_reason_message: str | None
started_at: str
updated_at: str
finished_at: str | None
INSTANCE_COLUMNS = """
id, aws_account_id, cloudflare_account_id, display_name, aws_region,
lightsail_instance_name, cloudflare_zone_name,
cloudflare_zone_id, cloudflare_record_name, socks_port, proxy_health_check,
health_timeout_seconds, release_grace_seconds, enabled, config_version,
last_known_ip, last_checked_at, created_at, updated_at, archived_at
"""
GROUP_COLUMNS = """
id, name, enabled, interval_minutes, next_run_at, last_run_at, config_version,
created_at, updated_at, archived_at
"""
FLEET_RUN_COLUMNS = """
id, target_type, target_id, target_name, trigger, status, active_slot,
current_item_id, lease_owner, lease_until, error_code, error_message,
total_items, succeeded_items, started_at, updated_at, finished_at
"""
FLEET_RUN_ITEM_COLUMNS = """
id, run_id, instance_id, position, status, stage, stage_started_at, attempt_count,
config_version, instance_display_name, aws_region, lightsail_instance_name,
aws_account_id, cloudflare_account_id, cloudflare_zone_name, cloudflare_zone_id,
cloudflare_record_name, socks_port,
proxy_health_check, health_timeout_seconds, release_grace_seconds,
old_static_ip_name, old_ip, new_static_ip_name, new_ip, dns_ip_before,
dns_ip_after, grace_until, aws_operation_id, error_code, error_message,
rollback_from_stage, rollback_reason_code, rollback_reason_message,
started_at, updated_at, finished_at
"""
INSTANCE_WRITABLE_FIELDS = (
"aws_account_id",
"cloudflare_account_id",
"display_name",
"aws_region",
"lightsail_instance_name",
"cloudflare_zone_name",
"cloudflare_zone_id",
"cloudflare_record_name",
"socks_port",
"proxy_health_check",
"health_timeout_seconds",
"release_grace_seconds",
"enabled",
)
GROUP_WRITABLE_FIELDS = ("name", "enabled", "interval_minutes")
class FleetRepository:
def __init__(self, database: Database) -> None:
self.database = database
def list_instances(self, *, include_archived: bool = False) -> list[ManagedInstanceRecord]:
where = "" if include_archived else "WHERE archived_at IS NULL"
with self.database.connect() as connection:
rows = connection.execute(
f"SELECT {INSTANCE_COLUMNS} FROM managed_instances "
f"{where} ORDER BY display_name COLLATE NOCASE, id"
).fetchall()
return [self._instance_from_row(row) for row in rows]
def get_instance(
self, instance_id: str, *, include_archived: bool = False
) -> ManagedInstanceRecord | None:
archived_clause = "" if include_archived else " AND archived_at IS NULL"
with self.database.connect() as connection:
row = connection.execute(
f"SELECT {INSTANCE_COLUMNS} FROM managed_instances "
f"WHERE id = ?{archived_clause}",
(instance_id,),
).fetchone()
return self._instance_from_row(row) if row else None
def create_instance(self, values: Mapping[str, Any]) -> ManagedInstanceRecord:
instance_id = str(values.get("id") or uuid.uuid4())
group_id = self._normalize_group_id(values.get("group_id"))
now = to_iso()
normalized = self._instance_values(values)
columns = ", ".join(INSTANCE_WRITABLE_FIELDS)
placeholders = ", ".join("?" for _ in INSTANCE_WRITABLE_FIELDS)
parameters = [normalized[field] for field in INSTANCE_WRITABLE_FIELDS]
with self.database.connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
self._ensure_writes_allowed(connection)
if group_id is not None and self._get_group_row(connection, group_id) is None:
raise RuntimeError("GROUP_NOT_FOUND")
self._validate_account_bindings(
connection,
normalized["aws_account_id"],
normalized["cloudflare_account_id"],
)
connection.execute(
f"""
INSERT INTO managed_instances(
id, {columns}, config_version, created_at, updated_at
) VALUES (?, {placeholders}, 1, ?, ?)
""",
(instance_id, *parameters, now, now),
)
if group_id is not None:
self._append_group_member(connection, group_id, instance_id)
self._touch_group_versions(connection, [group_id], now)
connection.commit()
except sqlite3.IntegrityError as exc:
connection.rollback()
raise RuntimeError(self._instance_integrity_code(exc)) from exc
except Exception:
connection.rollback()
raise
created = self.get_instance(instance_id)
if created is None: # pragma: no cover - insert and read share the same database
raise RuntimeError("INSTANCE_CREATION_FAILED")
return created
def update_instance(
self,
instance_id: str,
values: Mapping[str, Any],
*,
operation_owner_id: str | None = None,
) -> ManagedInstanceRecord:
updates = {
field: self._database_value(field, values[field])
for field in INSTANCE_WRITABLE_FIELDS
if field in values
}
group_supplied = "group_id" in values
requested_group_id = (
self._normalize_group_id(values["group_id"]) if group_supplied else None
)
with self.database.connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
self._ensure_writes_allowed(connection, operation_owner_id)
current = self._get_instance_row(connection, instance_id)
if current is None:
raise RuntimeError("INSTANCE_NOT_FOUND")
if (
"config_version" in values
and int(values["config_version"]) != int(current["config_version"])
):
raise RuntimeError("INSTANCE_CONFIG_VERSION_CONFLICT")
current_group_id = self._instance_group_id(connection, instance_id)
target_group_id = requested_group_id if group_supplied else current_group_id
if (
group_supplied
and target_group_id is not None
and self._get_group_row(connection, target_group_id) is None
):
raise RuntimeError("GROUP_NOT_FOUND")
group_changed = target_group_id != current_group_id
if not updates and not group_changed:
connection.commit()
return self._instance_from_row(current)
merged = {**dict(current), **updates}
self._validate_account_bindings(
connection,
merged["aws_account_id"],
merged["cloudflare_account_id"],
)
self._validate_dns_scope(
str(merged["cloudflare_zone_name"]),
str(merged["cloudflare_record_name"]),
)
final_enabled = bool(merged["enabled"])
if bool(current["enabled"]) and (
not final_enabled or group_changed
):
self._ensure_group_keeps_enabled_member(connection, instance_id)
updated_at = to_iso()
if updates:
assignments = ", ".join(f"{field} = ?" for field in updates)
connection.execute(
f"""
UPDATE managed_instances
SET {assignments}, config_version = config_version + 1,
updated_at = ?
WHERE id = ? AND archived_at IS NULL
""",
(*updates.values(), updated_at, instance_id),
)
else:
connection.execute(
"""
UPDATE managed_instances
SET config_version = config_version + 1, updated_at = ?
WHERE id = ? AND archived_at IS NULL
""",
(updated_at, instance_id),
)
if group_changed:
if current_group_id is not None:
connection.execute(
"DELETE FROM instance_group_members WHERE instance_id = ?",
(instance_id,),
)
if target_group_id is not None:
self._append_group_member(
connection, target_group_id, instance_id
)
self._touch_group_versions(
connection,
[
group_id
for group_id in (current_group_id, target_group_id)
if group_id is not None
],
updated_at,
)
connection.commit()
except sqlite3.IntegrityError as exc:
connection.rollback()
raise RuntimeError(self._instance_integrity_code(exc)) from exc
except Exception:
connection.rollback()
raise
updated = self.get_instance(instance_id)
if updated is None: # pragma: no cover
raise RuntimeError("INSTANCE_NOT_FOUND")
return updated
def update_instance_status(
self,
instance_id: str,
values: Mapping[str, str | None],
*,
operation_owner_id: str | None = None,
) -> ManagedInstanceRecord:
updates = {
field: values[field]
for field in ("last_known_ip", "last_checked_at")
if field in values
}
if not updates:
current = self.get_instance(instance_id)
if current is None:
raise RuntimeError("INSTANCE_NOT_FOUND")
return current
with self.database.connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
self._ensure_writes_allowed(connection, operation_owner_id)
if self._get_instance_row(connection, instance_id) is None:
raise RuntimeError("INSTANCE_NOT_FOUND")
assignments = ", ".join(f"{field} = ?" for field in updates)
connection.execute(
f"UPDATE managed_instances SET {assignments}, updated_at = ? "
"WHERE id = ? AND archived_at IS NULL",
(*updates.values(), to_iso(), instance_id),
)
connection.commit()
except Exception:
connection.rollback()
raise
updated = self.get_instance(instance_id)
if updated is None: # pragma: no cover
raise RuntimeError("INSTANCE_NOT_FOUND")
return updated
def archive_instance(self, instance_id: str) -> ManagedInstanceRecord:
now = to_iso()
with self.database.connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
self._ensure_writes_allowed(connection)
current = self._get_instance_row(connection, instance_id)
if current is None:
raise RuntimeError("INSTANCE_NOT_FOUND")
group_rows = connection.execute(
"""
SELECT groups.id, groups.enabled
FROM instance_group_members AS members
JOIN instance_groups AS groups ON groups.id = members.group_id
WHERE members.instance_id = ? AND groups.archived_at IS NULL
""",
(instance_id,),
).fetchall()
connection.execute(
"DELETE FROM instance_group_members WHERE instance_id = ?",
(instance_id,),
)
connection.execute(
"""
UPDATE managed_instances
SET enabled = 0, config_version = config_version + 1,
updated_at = ?, archived_at = ?
WHERE id = ? AND archived_at IS NULL
""",
(now, now, instance_id),
)
for group in group_rows:
if bool(group["enabled"]) and not self._group_has_enabled_member(
connection, str(group["id"])
):
connection.execute(
"""
UPDATE instance_groups
SET enabled = 0, next_run_at = NULL,
config_version = config_version + 1, updated_at = ?
WHERE id = ?
""",
(now, group["id"]),
)
else:
connection.execute(
"""
UPDATE instance_groups
SET config_version = config_version + 1, updated_at = ?
WHERE id = ?
""",
(now, group["id"]),
)
connection.commit()
except Exception:
connection.rollback()
raise
archived = self.get_instance(instance_id, include_archived=True)
if archived is None: # pragma: no cover
raise RuntimeError("INSTANCE_NOT_FOUND")
return archived
def list_groups(self, *, include_archived: bool = False) -> list[InstanceGroupRecord]:
where = "" if include_archived else "WHERE archived_at IS NULL"
with self.database.connect() as connection:
rows = connection.execute(
f"SELECT {GROUP_COLUMNS} FROM instance_groups "
f"{where} ORDER BY name COLLATE NOCASE, id"
).fetchall()
memberships = connection.execute(
"""
SELECT group_id, instance_id FROM instance_group_members
ORDER BY group_id, position, instance_id
"""
).fetchall()
members_by_group: dict[str, list[str]] = {}
for membership in memberships:
members_by_group.setdefault(str(membership["group_id"]), []).append(
str(membership["instance_id"])
)
return [
self._group_from_row(row, members_by_group.get(str(row["id"]), []))
for row in rows
]
def get_group(
self, group_id: str, *, include_archived: bool = False
) -> InstanceGroupRecord | None:
archived_clause = "" if include_archived else " AND archived_at IS NULL"
with self.database.connect() as connection:
row = connection.execute(
f"SELECT {GROUP_COLUMNS} FROM instance_groups "
f"WHERE id = ?{archived_clause}",
(group_id,),
).fetchone()
members = (
connection.execute(
"""
SELECT instance_id FROM instance_group_members
WHERE group_id = ? ORDER BY position, instance_id
""",
(group_id,),
).fetchall()
if row
else []
)
return (
self._group_from_row(row, [str(item["instance_id"]) for item in members])
if row
else None
)
def create_group(self, values: Mapping[str, Any]) -> InstanceGroupRecord:
group_id = str(values.get("id") or uuid.uuid4())
member_ids = self._normalize_member_ids(values.get("member_ids", []))
enabled = bool(values.get("enabled", False))
interval_minutes = int(values.get("interval_minutes", 60))
now = utc_now()
now_iso = to_iso(now)
next_run_at = to_iso(now + timedelta(minutes=interval_minutes)) if enabled else None
with self.database.connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
self._ensure_writes_allowed(connection)
self._validate_members(connection, member_ids, require_enabled=enabled)
connection.execute(
"""
INSERT INTO instance_groups(
id, name, enabled, interval_minutes, next_run_at,
config_version, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, 1, ?, ?)
""",
(
group_id,
str(values["name"]),
int(enabled),
interval_minutes,
next_run_at,
now_iso,
now_iso,
),
)
self._insert_members(connection, group_id, member_ids)
self._touch_instance_versions(connection, member_ids, now_iso)
connection.commit()
except sqlite3.IntegrityError as exc:
connection.rollback()
raise RuntimeError(self._group_integrity_code(exc)) from exc
except Exception:
connection.rollback()
raise
created = self.get_group(group_id)
if created is None: # pragma: no cover
raise RuntimeError("GROUP_CREATION_FAILED")
return created
def update_group(self, group_id: str, values: Mapping[str, Any]) -> InstanceGroupRecord:
updates = {
field: self._database_value(field, values[field])
for field in GROUP_WRITABLE_FIELDS
if field in values
}
member_ids = (
self._normalize_member_ids(values["member_ids"])
if "member_ids" in values
else None
)
with self.database.connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
self._ensure_writes_allowed(connection)
current = self._get_group_row(connection, group_id)
if current is None:
raise RuntimeError("GROUP_NOT_FOUND")
if (
"config_version" in values
and int(values["config_version"]) != int(current["config_version"])
):
raise RuntimeError("GROUP_CONFIG_VERSION_CONFLICT")
current_member_ids = (
self._member_ids(connection, group_id)
if member_ids is not None
else None
)
if not updates and member_ids is None:
connection.commit()
return self._group_from_row(
current, self._member_ids(connection, group_id)
)
enabled = bool(updates.get("enabled", current["enabled"]))
interval = int(updates.get("interval_minutes", current["interval_minutes"]))
if member_ids is not None:
self._validate_members(
connection,
member_ids,
require_enabled=enabled,
)
elif enabled and not self._group_has_enabled_member(connection, group_id):
raise RuntimeError("GROUP_REQUIRES_ENABLED_MEMBER")
enabled_changed = "enabled" in updates and enabled != bool(current["enabled"])
interval_changed = (
"interval_minutes" in updates
and interval != int(current["interval_minutes"])
)
next_run_at = current["next_run_at"]
if enabled_changed:
next_run_at = (
to_iso(utc_now() + timedelta(minutes=interval)) if enabled else None
)
elif enabled and interval_changed:
next_run_at = to_iso(utc_now() + timedelta(minutes=interval))
updates["next_run_at"] = next_run_at
assignments = ", ".join(f"{field} = ?" for field in updates)
updated_at = to_iso()
connection.execute(
f"""
UPDATE instance_groups
SET {assignments}, config_version = config_version + 1, updated_at = ?
WHERE id = ? AND archived_at IS NULL
""",
(*updates.values(), updated_at, group_id),
)
if member_ids is not None:
connection.execute(
"DELETE FROM instance_group_members WHERE group_id = ?", (group_id,)
)
self._insert_members(connection, group_id, member_ids)
changed_member_ids = sorted(
set(current_member_ids or []) ^ set(member_ids)
)
self._touch_instance_versions(
connection, changed_member_ids, updated_at
)
connection.commit()
except sqlite3.IntegrityError as exc:
connection.rollback()
raise RuntimeError(self._group_integrity_code(exc)) from exc
except Exception:
connection.rollback()
raise
updated = self.get_group(group_id)
if updated is None: # pragma: no cover
raise RuntimeError("GROUP_NOT_FOUND")
return updated
def save_group_members(
self, group_id: str, member_ids: Sequence[str]
) -> InstanceGroupRecord:
normalized_ids = self._normalize_member_ids(member_ids)
with self.database.connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
self._ensure_writes_allowed(connection)
group = self._get_group_row(connection, group_id)
if group is None:
raise RuntimeError("GROUP_NOT_FOUND")
current_member_ids = self._member_ids(connection, group_id)
self._validate_members(
connection,
normalized_ids,
require_enabled=bool(group["enabled"]),
)
connection.execute(
"DELETE FROM instance_group_members WHERE group_id = ?", (group_id,)
)
self._insert_members(connection, group_id, normalized_ids)
updated_at = to_iso()
changed_member_ids = sorted(
set(current_member_ids) ^ set(normalized_ids)
)
self._touch_instance_versions(
connection, changed_member_ids, updated_at
)
connection.execute(
"""
UPDATE instance_groups
SET config_version = config_version + 1, updated_at = ?
WHERE id = ? AND archived_at IS NULL
""",
(updated_at, group_id),
)
connection.commit()
except sqlite3.IntegrityError as exc:
connection.rollback()
raise RuntimeError(self._group_integrity_code(exc)) from exc
except Exception:
connection.rollback()
raise
updated = self.get_group(group_id)
if updated is None: # pragma: no cover
raise RuntimeError("GROUP_NOT_FOUND")
return updated
def archive_group(self, group_id: str) -> InstanceGroupRecord:
now = to_iso()
with self.database.connect() as connection:
try:
connection.execute("BEGIN IMMEDIATE")
self._ensure_writes_allowed(connection)
if self._get_group_row(connection, group_id) is None:
raise RuntimeError("GROUP_NOT_FOUND")
member_ids = self._member_ids(connection, group_id)
connection.execute(
"DELETE FROM instance_group_members WHERE group_id = ?", (group_id,)
)
self._touch_instance_versions(connection, member_ids, now)
connection.execute(
"""
UPDATE instance_groups
SET enabled = 0, next_run_at = NULL,
config_version = config_version + 1,
updated_at = ?, archived_at = ?
WHERE id = ? AND archived_at IS NULL
""",
(now, now, group_id),
)
connection.commit()
except Exception:
connection.rollback()
raise
archived = self.get_group(group_id, include_archived=True)
if archived is None: # pragma: no cover
raise RuntimeError("GROUP_NOT_FOUND")
return archived
def get_fleet_run(self, run_id: str) -> FleetRunRecord | None:
with self.database.connect() as connection:
row = connection.execute(
f"SELECT {FLEET_RUN_COLUMNS} FROM fleet_runs WHERE id = ?", (run_id,)
).fetchone()
return FleetRunRecord(**dict(row)) if row else None
def get_fleet_run_item(self, item_id: str) -> FleetRunItemRecord | None:
with self.database.connect() as connection:
row = connection.execute(
f"SELECT {FLEET_RUN_ITEM_COLUMNS} FROM fleet_run_items WHERE id = ?",
(item_id,),
).fetchone()
return self._run_item_from_row(row) if row else None
def list_fleet_run_items(self, run_id: str) -> list[FleetRunItemRecord]:
with self.database.connect() as connection:
rows = connection.execute(
f"SELECT {FLEET_RUN_ITEM_COLUMNS} FROM fleet_run_items "
"WHERE run_id = ? ORDER BY position, id",
(run_id,),
).fetchall()
return [self._run_item_from_row(row) for row in rows]
@staticmethod
def _ensure_writes_allowed(
connection: sqlite3.Connection,
operation_owner_id: str | None = None,
) -> None:
connection.execute(
"""
DELETE FROM fleet_operation_locks
WHERE kind = 'dns_sync' AND lease_until IS NOT NULL AND lease_until < ?
""",
(to_iso(),),
)
if connection.execute(
"SELECT 1 FROM fleet_runs WHERE active_slot = 1 LIMIT 1"
).fetchone():
raise RuntimeError("FLEET_RUN_ACTIVE")
if operation_owner_id is None:
lock = connection.execute(
"SELECT 1 FROM fleet_operation_locks LIMIT 1"
).fetchone()
else:
lock = connection.execute(
"""
SELECT 1 FROM fleet_operation_locks
WHERE owner_id != ? LIMIT 1
""",
(operation_owner_id,),
).fetchone()
if lock:
raise RuntimeError("FLEET_RUN_ACTIVE")
@staticmethod
def _instance_values(values: Mapping[str, Any]) -> dict[str, Any]:
defaults: dict[str, Any] = {
"aws_account_id": None,
"cloudflare_account_id": None,
"cloudflare_zone_id": "",
"socks_port": 1080,
"proxy_health_check": True,
"health_timeout_seconds": 120,
"release_grace_seconds": 75,
"enabled": True,
}
normalized = {**defaults, **values}
missing = [field for field in INSTANCE_WRITABLE_FIELDS if field not in normalized]
if missing:
raise RuntimeError("INSTANCE_DATA_INCOMPLETE")
FleetRepository._validate_dns_scope(
str(normalized["cloudflare_zone_name"]),
str(normalized["cloudflare_record_name"]),
)
return {
field: FleetRepository._database_value(field, normalized[field])
for field in INSTANCE_WRITABLE_FIELDS
}
@staticmethod
def _database_value(field: str, value: Any) -> Any:
if field in {"enabled", "proxy_health_check"}:
return int(bool(value))
return value
@staticmethod
def _validate_dns_scope(zone_name: str, record_name: str) -> None:
zone = zone_name.rstrip(".").lower()
record = record_name.rstrip(".").lower()
if record != zone and not record.endswith(f".{zone}"):
raise RuntimeError("DNS_RECORD_OUTSIDE_ZONE")
@staticmethod
def _validate_account_bindings(
connection: sqlite3.Connection,
aws_account_id: object,
cloudflare_account_id: object,
) -> None:
for account_id, provider in (
(aws_account_id, "aws"),
(cloudflare_account_id, "cloudflare"),
):
if account_id is None:
continue
row = connection.execute(
"""
SELECT provider FROM credential_accounts
WHERE id = ? AND archived_at IS NULL
""",
(str(account_id),),
).fetchone()
if row is None:
raise RuntimeError(f"{provider.upper()}_ACCOUNT_NOT_FOUND")
if str(row["provider"]) != provider:
raise RuntimeError(f"{provider.upper()}_ACCOUNT_PROVIDER_MISMATCH")
@staticmethod
def _normalize_member_ids(member_ids: object) -> list[str]:
if isinstance(member_ids, (str, bytes)) or not isinstance(member_ids, Sequence):
raise RuntimeError("INVALID_GROUP_MEMBERS")
normalized = [str(instance_id).strip() for instance_id in member_ids]
if any(not instance_id for instance_id in normalized):
raise RuntimeError("INVALID_GROUP_MEMBERS")
if len(set(normalized)) != len(normalized):
raise RuntimeError("DUPLICATE_GROUP_MEMBER")
return normalized
@staticmethod
def _validate_members(
connection: sqlite3.Connection,
member_ids: Sequence[str],
*,
require_enabled: bool,
) -> None:
if not member_ids:
if require_enabled:
raise RuntimeError("GROUP_REQUIRES_ENABLED_MEMBER")
return
placeholders = ", ".join("?" for _ in member_ids)
rows = connection.execute(
f"""
SELECT id, enabled FROM managed_instances
WHERE id IN ({placeholders}) AND archived_at IS NULL
""",
tuple(member_ids),
).fetchall()
if len(rows) != len(member_ids):
raise RuntimeError("INSTANCE_NOT_FOUND")
if require_enabled and not any(bool(row["enabled"]) for row in rows):
raise RuntimeError("GROUP_REQUIRES_ENABLED_MEMBER")
@staticmethod
def _insert_members(
connection: sqlite3.Connection, group_id: str, member_ids: Sequence[str]
) -> None:
connection.executemany(
"""
INSERT INTO instance_group_members(group_id, instance_id, position)
VALUES (?, ?, ?)
""",
((group_id, instance_id, position) for position, instance_id in enumerate(member_ids)),
)
@staticmethod
def _append_group_member(
connection: sqlite3.Connection, group_id: str, instance_id: str
) -> None:
connection.execute(
"""
INSERT INTO instance_group_members(group_id, instance_id, position)
SELECT ?, ?, COALESCE(MAX(position), -1) + 1
FROM instance_group_members
WHERE group_id = ?
""",
(group_id, instance_id, group_id),
)
@staticmethod
def _touch_group_versions(
connection: sqlite3.Connection,
group_ids: Sequence[str],
updated_at: str,
) -> None:
connection.executemany(
"""
UPDATE instance_groups
SET config_version = config_version + 1, updated_at = ?
WHERE id = ? AND archived_at IS NULL
""",
((updated_at, group_id) for group_id in dict.fromkeys(group_ids)),
)
@staticmethod
def _touch_instance_versions(
connection: sqlite3.Connection,
instance_ids: Sequence[str],
updated_at: str,
) -> None:
connection.executemany(
"""
UPDATE managed_instances
SET config_version = config_version + 1, updated_at = ?
WHERE id = ? AND archived_at IS NULL
""",
((updated_at, instance_id) for instance_id in dict.fromkeys(instance_ids)),
)
@staticmethod
def _member_ids(connection: sqlite3.Connection, group_id: str) -> list[str]:
rows = connection.execute(
"""
SELECT instance_id FROM instance_group_members
WHERE group_id = ? ORDER BY position, instance_id
""",
(group_id,),
).fetchall()
return [str(row["instance_id"]) for row in rows]
@staticmethod
def _instance_group_id(
connection: sqlite3.Connection, instance_id: str
) -> str | None:
row = connection.execute(
"SELECT group_id FROM instance_group_members WHERE instance_id = ?",
(instance_id,),
).fetchone()
return str(row["group_id"]) if row else None
@staticmethod
def _normalize_group_id(value: object) -> str | None:
if value is None:
return None
normalized = str(value).strip()
return normalized or None
@staticmethod
def _group_has_enabled_member(connection: sqlite3.Connection, group_id: str) -> bool:
return (
connection.execute(
"""
SELECT 1
FROM instance_group_members AS members
JOIN managed_instances AS instances ON instances.id = members.instance_id
WHERE members.group_id = ?
AND instances.archived_at IS NULL
AND instances.enabled = 1
LIMIT 1
""",
(group_id,),
).fetchone()
is not None
)
@staticmethod
def _ensure_group_keeps_enabled_member(
connection: sqlite3.Connection, instance_id: str
) -> None:
group = connection.execute(
"""
SELECT groups.id
FROM instance_group_members AS members
JOIN instance_groups AS groups ON groups.id = members.group_id
WHERE members.instance_id = ?
AND groups.archived_at IS NULL
AND groups.enabled = 1
""",
(instance_id,),
).fetchone()
if group is None:
return
other = connection.execute(
"""
SELECT 1
FROM instance_group_members AS members
JOIN managed_instances AS instances ON instances.id = members.instance_id
WHERE members.group_id = ?
AND members.instance_id <> ?
AND instances.archived_at IS NULL
AND instances.enabled = 1
LIMIT 1
""",
(group["id"], instance_id),
).fetchone()
if other is None:
raise RuntimeError("GROUP_REQUIRES_ENABLED_MEMBER")
@staticmethod
def _get_instance_row(
connection: sqlite3.Connection, instance_id: str
) -> sqlite3.Row | None:
return connection.execute(
f"SELECT {INSTANCE_COLUMNS} FROM managed_instances "
"WHERE id = ? AND archived_at IS NULL",
(instance_id,),
).fetchone()
@staticmethod
def _get_group_row(
connection: sqlite3.Connection, group_id: str
) -> sqlite3.Row | None:
return connection.execute(
f"SELECT {GROUP_COLUMNS} FROM instance_groups "
"WHERE id = ? AND archived_at IS NULL",
(group_id,),
).fetchone()
@staticmethod
def _instance_from_row(row: sqlite3.Row) -> ManagedInstanceRecord:
values = dict(row)
values["proxy_health_check"] = bool(values["proxy_health_check"])
values["enabled"] = bool(values["enabled"])
return ManagedInstanceRecord(**values)
@staticmethod
def _group_from_row(row: sqlite3.Row, member_ids: list[str]) -> InstanceGroupRecord:
values = dict(row)
values["enabled"] = bool(values["enabled"])
return InstanceGroupRecord(**values, member_ids=member_ids)
@staticmethod
def _run_item_from_row(row: sqlite3.Row) -> FleetRunItemRecord:
values = dict(row)
values["proxy_health_check"] = bool(values["proxy_health_check"])
return FleetRunItemRecord(**values)
@staticmethod
def _instance_integrity_code(exc: sqlite3.IntegrityError) -> str:
message = str(exc)
if "instance_group_members.instance_id" in message:
return "INSTANCE_ALREADY_GROUPED"
if "managed_instances.display_name" in message:
return "INSTANCE_DISPLAY_NAME_CONFLICT"
if (
"managed_instances.aws_region" in message
and "managed_instances.lightsail_instance_name" in message
) or "uq_managed_instances_aws_target_active" in message:
return "INSTANCE_AWS_TARGET_CONFLICT"
if "uq_managed_instances_dns_record_active" in message:
return "INSTANCE_DNS_RECORD_CONFLICT"
if "managed_instances.id" in message:
return "INSTANCE_ID_CONFLICT"
if "CHECK constraint failed" in message or "NOT NULL constraint failed" in message:
return "INVALID_INSTANCE_DATA"
return "INSTANCE_CONSTRAINT_CONFLICT"
@staticmethod
def _group_integrity_code(exc: sqlite3.IntegrityError) -> str:
message = str(exc)
if "instance_groups.name" in message:
return "GROUP_NAME_CONFLICT"
if "instance_groups.id" in message:
return "GROUP_ID_CONFLICT"
if "instance_group_members.instance_id" in message:
return "INSTANCE_ALREADY_GROUPED"
if "FOREIGN KEY constraint failed" in message:
return "GROUP_MEMBER_NOT_FOUND"
if "CHECK constraint failed" in message or "NOT NULL constraint failed" in message:
return "INVALID_GROUP_DATA"
return "GROUP_CONSTRAINT_CONFLICT"