feat: complete client runtime foundation
This commit is contained in:
+151
-17
@@ -11,7 +11,9 @@ from typing import Any
|
||||
|
||||
from websockets.asyncio.client import connect
|
||||
|
||||
from archive_clients.backup import SQLiteBackupManager
|
||||
from archive_clients.config import ClientConfig
|
||||
from archive_clients.locking import DatabaseLease
|
||||
from archive_clients.probes import FilesystemProbe
|
||||
from archive_clients.protocol import (
|
||||
decode,
|
||||
@@ -20,6 +22,7 @@ from archive_clients.protocol import (
|
||||
encode_message,
|
||||
new_envelope,
|
||||
)
|
||||
from archive_clients.services import ServiceProbe
|
||||
from archive_clients.state import ClientStore, CommandConflict
|
||||
from archive_control.v1 import client_pb2, common_pb2, control_pb2, job_pb2
|
||||
|
||||
@@ -28,13 +31,49 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ArchiveClientDaemon:
|
||||
def __init__(self, config: ClientConfig, probes: list[FilesystemProbe]):
|
||||
def __init__(
|
||||
self,
|
||||
config: ClientConfig,
|
||||
probes: list[FilesystemProbe],
|
||||
service_probes: list[ServiceProbe],
|
||||
):
|
||||
if len(probes) != 2:
|
||||
raise ValueError(
|
||||
"qBittorrent and Syncthing filesystem probes are required"
|
||||
)
|
||||
self.config = config
|
||||
self.probes = probes
|
||||
self.service_probes = service_probes
|
||||
self.store = ClientStore(config.state_db)
|
||||
self.backups = SQLiteBackupManager(
|
||||
config.state_db, config.backup_dir, config.backup
|
||||
)
|
||||
self._lease = DatabaseLease(config.state_db)
|
||||
|
||||
async def run(self) -> None:
|
||||
await asyncio.to_thread(self.store.initialize)
|
||||
await asyncio.to_thread(self._lease.acquire)
|
||||
backup_task: asyncio.Task[None] | None = None
|
||||
try:
|
||||
await asyncio.to_thread(self.store.initialize)
|
||||
await asyncio.to_thread(self._warn_if_backup_shares_filesystem)
|
||||
backup_task = asyncio.create_task(
|
||||
self._backup_loop(), name="archive-client-backups"
|
||||
)
|
||||
await self._connection_loop()
|
||||
finally:
|
||||
if backup_task is not None:
|
||||
backup_task.cancel()
|
||||
await asyncio.gather(backup_task, return_exceptions=True)
|
||||
await asyncio.to_thread(self._lease.release)
|
||||
|
||||
def _warn_if_backup_shares_filesystem(self) -> None:
|
||||
if (
|
||||
self.config.state_db.parent.stat().st_dev
|
||||
== self.config.backup_dir.stat().st_dev
|
||||
):
|
||||
logger.warning("database_and_backup_share_filesystem")
|
||||
|
||||
async def _connection_loop(self) -> None:
|
||||
delay = self.config.connection.reconnect_initial
|
||||
while True:
|
||||
started = time.monotonic()
|
||||
@@ -43,13 +82,39 @@ class ArchiveClientDaemon:
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.warning("control connection ended: %s", type(exc).__name__)
|
||||
if time.monotonic() - started >= self.config.connection.reconnect_reset_after:
|
||||
logger.warning(
|
||||
"control_connection_ended",
|
||||
extra={"error_type": type(exc).__name__},
|
||||
)
|
||||
if (
|
||||
time.monotonic() - started
|
||||
>= self.config.connection.reconnect_reset_after
|
||||
):
|
||||
delay = self.config.connection.reconnect_initial
|
||||
wait = random.uniform(0, delay) if self.config.connection.reconnect_jitter else delay
|
||||
wait = (
|
||||
random.uniform(0, delay)
|
||||
if self.config.connection.reconnect_jitter else delay
|
||||
)
|
||||
await asyncio.sleep(wait)
|
||||
delay = min(delay * 2, self.config.connection.reconnect_max)
|
||||
|
||||
async def _backup_loop(self) -> None:
|
||||
while True:
|
||||
await asyncio.sleep(self.config.backup.interval)
|
||||
try:
|
||||
record = await asyncio.to_thread(self.backups.create, "scheduled")
|
||||
logger.info(
|
||||
"database_backup_created",
|
||||
extra={"backup_name": record.database.name},
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
"database_backup_failed",
|
||||
extra={"error_type": type(exc).__name__},
|
||||
)
|
||||
|
||||
async def _connection(self) -> None:
|
||||
async with connect(
|
||||
self.config.control_endpoint, ping_interval=None, compression=None,
|
||||
@@ -69,6 +134,7 @@ class ArchiveClientDaemon:
|
||||
raise RuntimeError("control rejected registration")
|
||||
if response.register_response.negotiated_version.major != 1:
|
||||
raise RuntimeError("control negotiated an unsupported protocol version")
|
||||
logger.info("control_connection_registered")
|
||||
outbound: asyncio.Queue[str] = asyncio.Queue(maxsize=100)
|
||||
writer = asyncio.create_task(self._writer(websocket, outbound))
|
||||
try:
|
||||
@@ -94,6 +160,23 @@ class ArchiveClientDaemon:
|
||||
request.capabilities.syncthing_advertised_addresses.extend(
|
||||
self.config.syncthing.advertised_addresses
|
||||
)
|
||||
for probe in self.service_probes:
|
||||
health = request.capabilities.services.add()
|
||||
health.service = probe.service
|
||||
health.state = probe.state
|
||||
health.version = probe.version
|
||||
health.api_version = probe.api_version
|
||||
health.detail = probe.detail
|
||||
health.checked_at.FromDatetime(probe.checked_at)
|
||||
if probe.service == "qbittorrent":
|
||||
request.capabilities.qbittorrent_version = probe.version
|
||||
request.capabilities.qbittorrent_web_api_version = (
|
||||
probe.api_version
|
||||
)
|
||||
request.capabilities.libtorrent_version = probe.libtorrent_version
|
||||
elif probe.service == "syncthing":
|
||||
request.capabilities.syncthing_version = probe.version
|
||||
request.capabilities.syncthing_device_id = probe.device_id
|
||||
for root_name, probe in zip(
|
||||
("qbittorrent", "syncthing"), self.probes, strict=True
|
||||
):
|
||||
@@ -102,9 +185,15 @@ class ArchiveClientDaemon:
|
||||
filesystem.readable = probe.readable
|
||||
filesystem.writable = probe.writable
|
||||
filesystem.hard_link = probe.hard_link
|
||||
filesystem.reflink = probe.reflink
|
||||
filesystem.sparse_files = probe.sparse_files
|
||||
if all(probe.hard_link for probe in self.probes):
|
||||
request.capabilities.features.append(client_pb2.CLIENT_FEATURE_HARD_LINK)
|
||||
if all(probe.reflink for probe in self.probes):
|
||||
request.capabilities.features.append(client_pb2.CLIENT_FEATURE_REFLINK)
|
||||
if all(probe.sparse_files for probe in self.probes):
|
||||
request.capabilities.features.append(client_pb2.CLIENT_FEATURE_SPARSE_FILES)
|
||||
request.capabilities.features.append(client_pb2.CLIENT_FEATURE_DB_BACKUP)
|
||||
for cursor in self.store.list_active_job_cursors():
|
||||
active = request.active_jobs.add()
|
||||
active.job_id = str(cursor["job_id"])
|
||||
@@ -132,14 +221,32 @@ class ArchiveClientDaemon:
|
||||
elif payload == "command":
|
||||
await self._accept_command(envelope, outbound)
|
||||
elif payload == "protocol_error":
|
||||
logger.warning("control reported protocol error code=%s", envelope.protocol_error.error.code)
|
||||
logger.warning(
|
||||
"control_reported_protocol_error",
|
||||
extra={"error_code": envelope.protocol_error.error.code},
|
||||
)
|
||||
|
||||
async def _accept_command(
|
||||
self, envelope: Any, outbound: asyncio.Queue[str]
|
||||
) -> None:
|
||||
command = envelope.command
|
||||
acknowledgement = self._initial_acknowledgement(command)
|
||||
snapshot_rows: list[dict[str, object]] = []
|
||||
if command.WhichOneof("payload") == "request_job_snapshot":
|
||||
requested = list(command.request_job_snapshot.job_ids)
|
||||
snapshot_rows = await asyncio.to_thread(
|
||||
self.store.job_snapshot_rows, requested
|
||||
)
|
||||
missing = set(requested) - {
|
||||
str(row["job_id"]) for row in snapshot_rows
|
||||
}
|
||||
else:
|
||||
requested = []
|
||||
missing = set()
|
||||
acknowledgement = self._initial_acknowledgement(
|
||||
command, missing, bool(requested)
|
||||
)
|
||||
accepted = None
|
||||
accepted_for_execution = False
|
||||
try:
|
||||
accepted = await asyncio.to_thread(
|
||||
self.store.accept_command,
|
||||
@@ -155,9 +262,15 @@ class ArchiveClientDaemon:
|
||||
acknowledgement.status
|
||||
== control_pb2.COMMAND_ACK_STATUS_ACCEPTED
|
||||
):
|
||||
accepted_for_execution = True
|
||||
acknowledgement.status = (
|
||||
control_pb2.COMMAND_ACK_STATUS_DUPLICATE
|
||||
)
|
||||
else:
|
||||
accepted_for_execution = (
|
||||
acknowledgement.status
|
||||
== control_pb2.COMMAND_ACK_STATUS_ACCEPTED
|
||||
)
|
||||
except CommandConflict:
|
||||
acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_REJECTED
|
||||
acknowledgement.error.code = common_pb2.ERROR_CODE_CONFLICT
|
||||
@@ -169,26 +282,47 @@ class ArchiveClientDaemon:
|
||||
|
||||
if (
|
||||
accepted is not None
|
||||
and not accepted.duplicate
|
||||
and acknowledgement.status
|
||||
== control_pb2.COMMAND_ACK_STATUS_ACCEPTED
|
||||
and accepted_for_execution
|
||||
and command.WhichOneof("payload") == "request_job_snapshot"
|
||||
):
|
||||
snapshot = new_envelope()
|
||||
snapshot.correlation_id = envelope.message_id
|
||||
snapshot.client_state_snapshot.snapshot_id = str(uuid.uuid4())
|
||||
snapshot.client_state_snapshot.observed_at.CopyFrom(snapshot.sent_at)
|
||||
await outbound.put(encode(snapshot))
|
||||
for row in snapshot_rows:
|
||||
snapshot = new_envelope()
|
||||
snapshot.correlation_id = envelope.message_id
|
||||
job_snapshot = snapshot.job_snapshot
|
||||
job_snapshot.job.definition.CopyFrom(decode_message(
|
||||
str(row["definition_json"]), job_pb2.JobDefinition()
|
||||
))
|
||||
job_snapshot.job.state = job_pb2.JobState.Value(str(row["state"]))
|
||||
job_snapshot.job.revision = int(row["revision"])
|
||||
job_snapshot.job.committed = bool(row["committed"])
|
||||
job_snapshot.job.updated_at.CopyFrom(snapshot.sent_at)
|
||||
job_snapshot.last_event_sequence = int(
|
||||
row["last_event_sequence"]
|
||||
)
|
||||
await outbound.put(encode(snapshot))
|
||||
|
||||
@staticmethod
|
||||
def _initial_acknowledgement(command: Any) -> control_pb2.CommandAck:
|
||||
def _initial_acknowledgement(
|
||||
command: Any,
|
||||
missing_snapshot_jobs: set[str],
|
||||
has_snapshot_jobs: bool,
|
||||
) -> control_pb2.CommandAck:
|
||||
acknowledgement = control_pb2.CommandAck(command_id=command.command_id)
|
||||
if not command.command_id or command.WhichOneof("payload") is None:
|
||||
acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_REJECTED
|
||||
acknowledgement.error.code = common_pb2.ERROR_CODE_INVALID_ARGUMENT
|
||||
acknowledgement.error.message = "command ID and payload are required"
|
||||
elif command.WhichOneof("payload") == "request_job_snapshot":
|
||||
acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_ACCEPTED
|
||||
if not has_snapshot_jobs:
|
||||
acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_REJECTED
|
||||
acknowledgement.error.code = common_pb2.ERROR_CODE_INVALID_ARGUMENT
|
||||
acknowledgement.error.message = "snapshot job IDs are required"
|
||||
elif missing_snapshot_jobs:
|
||||
acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_REJECTED
|
||||
acknowledgement.error.code = common_pb2.ERROR_CODE_NOT_FOUND
|
||||
acknowledgement.error.message = "requested client job is not found"
|
||||
else:
|
||||
acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_ACCEPTED
|
||||
else:
|
||||
acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_REJECTED
|
||||
acknowledgement.error.code = common_pb2.ERROR_CODE_UNSUPPORTED
|
||||
|
||||
Reference in New Issue
Block a user