import asyncio import os import tempfile import unittest from datetime import datetime, timezone from pathlib import Path, PurePosixPath from uuid import uuid4 from websockets.asyncio.server import serve from archive_clients.config import ClientConfig, ConnectionConfig, ServiceConfig from archive_clients.daemon import ArchiveClientDaemon from archive_clients.probes import FilesystemProbe from archive_clients.services import ServiceProbe from archive_clients.protocol import decode, encode, encode_message, new_envelope from archive_control.v1 import client_pb2, common_pb2, control_pb2, job_pb2 class DaemonTransportTests(unittest.IsolatedAsyncioTestCase): async def test_registration_heartbeat_and_duplicate_command(self): observed = {} job_id = str(uuid4()) async def control(websocket): registration = decode(await websocket.recv()) observed["token"] = registration.register_request.shared_token observed["root_names"] = [ item.root_name for item in registration.register_request.capabilities.filesystems ] observed["addresses"] = list( registration.register_request.capabilities .syncthing_advertised_addresses ) observed["device_id"] = ( registration.register_request.capabilities.syncthing_device_id ) observed["services"] = [ item.service for item in registration.register_request.capabilities.services ] response = new_envelope() response.correlation_id = registration.message_id response.register_response.status = client_pb2.REGISTRATION_STATUS_ACCEPTED response.register_response.negotiated_version.major = 1 await websocket.send(encode(response)) heartbeat = new_envelope() heartbeat.heartbeat.sequence = 7 await websocket.send(encode(heartbeat)) observed["heartbeat"] = decode(await websocket.recv()).heartbeat_ack.sequence command_id = str(uuid4()) command = new_envelope() command.command.command_id = command_id command.command.created_at.CopyFrom(command.sent_at) command.command.request_job_snapshot.job_ids.append(job_id) await websocket.send(encode(command)) observed["first"] = decode(await websocket.recv()).command_ack.status snapshot = decode(await websocket.recv()) observed["snapshot"] = snapshot.WhichOneof("payload") observed["snapshot_job_id"] = ( snapshot.job_snapshot.job.definition.job_id ) duplicate = new_envelope() duplicate.command.CopyFrom(command.command) await websocket.send(encode(duplicate)) observed["second"] = decode(await websocket.recv()).command_ack.status observed["duplicate_snapshot"] = ( decode(await websocket.recv()).WhichOneof("payload") ) unsupported = new_envelope() unsupported.command.command_id = str(uuid4()) unsupported.command.created_at.CopyFrom(unsupported.sent_at) unsupported.command.assign_job.SetInParent() await websocket.send(encode(unsupported)) rejected = decode(await websocket.recv()).command_ack observed["rejected"] = (rejected.status, rejected.error.code) unsupported_duplicate = new_envelope() unsupported_duplicate.command.CopyFrom(unsupported.command) await websocket.send(encode(unsupported_duplicate)) observed["rejected_duplicate"] = ( decode(await websocket.recv()).command_ack.status ) with tempfile.TemporaryDirectory() as directory: root = Path(directory) token = root / "token" token.write_text("shared-secret", encoding="utf-8") os.chmod(token, 0o600) async with serve(control, "127.0.0.1", 0, ping_interval=None) as server: port = server.sockets[0].getsockname()[1] service = ServiceConfig( "http://local", PurePosixPath("/api"), root, advertised_addresses=("dynamic",), ) config = ClientConfig( "cache-1", "Cache 1", "cache", f"ws://127.0.0.1:{port}", token, root / "state.db", root / "backups", service, service, ConnectionConfig(registration_timeout=2), ) probe = FilesystemProbe(root, True, True, True, True, True) service_probe = ServiceProbe( "syncthing", common_pb2.HEALTH_STATE_HEALTHY, datetime.now(timezone.utc), version="v2", device_id="DEVICE", ) daemon = ArchiveClientDaemon( config, [probe, probe], [service_probe] ) await asyncio.to_thread(daemon.store.initialize) definition = job_pb2.JobDefinition( job_id=job_id, operation=job_pb2.JOB_OPERATION_ARCHIVE, ) definition.created_at.GetCurrentTime() await asyncio.to_thread( daemon.store.save_job, job_id, encode_message(definition), "JOB_STATE_WAITING", 2, 3, False, ) await daemon._connection() self.assertEqual(observed["token"], "shared-secret") self.assertEqual(observed["root_names"], ["qbittorrent", "syncthing"]) self.assertEqual(observed["addresses"], ["dynamic"]) self.assertEqual(observed["device_id"], "DEVICE") self.assertEqual(observed["services"], ["syncthing"]) self.assertEqual(observed["heartbeat"], 7) self.assertEqual(observed["first"], control_pb2.COMMAND_ACK_STATUS_ACCEPTED) self.assertEqual(observed["second"], control_pb2.COMMAND_ACK_STATUS_DUPLICATE) self.assertEqual(observed["snapshot"], "job_snapshot") self.assertEqual(observed["snapshot_job_id"], job_id) self.assertEqual(observed["duplicate_snapshot"], "job_snapshot") self.assertEqual( observed["rejected"], ( control_pb2.COMMAND_ACK_STATUS_REJECTED, common_pb2.ERROR_CODE_UNSUPPORTED, ), ) self.assertEqual( observed["rejected_duplicate"], control_pb2.COMMAND_ACK_STATUS_REJECTED, ) if __name__ == "__main__": unittest.main()