feat: complete client runtime foundation

This commit is contained in:
2026-07-22 16:25:38 +00:00
parent c11a7b5b5b
commit 1219117403
20 changed files with 1201 additions and 48 deletions
+151 -17
View File
@@ -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