import asyncio import os import tempfile import threading import unittest from datetime import datetime, timezone from pathlib import Path, PurePosixPath from unittest.mock import Mock 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.syncthing import ConfiguredRoute from archive_clients.protocol import decode, encode, encode_message, new_envelope from archive_control.v1 import ( client_pb2, common_pb2, control_pb2, inventory_pb2, job_pb2, route_pb2, ) class DaemonTransportTests(unittest.IsolatedAsyncioTestCase): async def test_ensure_route_is_durable_and_duplicate_replays_updates(self): with tempfile.TemporaryDirectory() as directory: root = Path(directory) token = root / "token" token.write_text("shared-secret", encoding="utf-8") os.chmod(token, 0o600) service = ServiceConfig( "http://local", PurePosixPath("/sync"), root ) config = ClientConfig( "cache-1", "Cache 1", "cache", "ws://control", token, root / "state.db", root / "backups", service, service, ) local_route = route_pb2.LocalRoute( route_id="route-1", local_relative_path="routes/route-1", folder_type=route_pb2.SYNCTHING_FOLDER_TYPE_SEND_RECEIVE, local_syncthing_device_id="LOCAL", peer_syncthing_device_ids=["PEER"], state=route_pb2.ROUTE_STATE_PROVISIONING, writable=True, ) local_route.observed_at.GetCurrentTime() manager = Mock() manager.configure.return_value = ConfiguredRoute( "LOCAL", root / "routes/route-1", local_route, True, True ) manager.verify.return_value = (True, True) probe = FilesystemProbe(root, True, True, True, True, True) syncthing_probe = ServiceProbe( "syncthing", common_pb2.HEALTH_STATE_HEALTHY, datetime.now(timezone.utc), device_id="LOCAL", ) daemon = ArchiveClientDaemon( config, [probe, probe], [syncthing_probe], route_manager=manager, ) await asyncio.to_thread(daemon.store.initialize) command = new_envelope() command.command.command_id = str(uuid4()) command.command.created_at.CopyFrom(command.sent_at) spec = command.command.ensure_route.route spec.route_id = "route-1" spec.peer_client_id = "archive-1" spec.peer_syncthing_device_id = "PEER" spec.peer_addresses.append("dynamic") spec.local_relative_path = "routes/route-1" spec.setup_timeout_seconds = 1800 outbound = asyncio.Queue() tasks = set() await daemon._handle(command, outbound, tasks) self.assertEqual( decode(await outbound.get()).command_ack.status, control_pb2.COMMAND_ACK_STATUS_ACCEPTED, ) await next(iter(tasks)) first_updates = [decode(await outbound.get()).route_update for _ in range(3)] self.assertEqual( [update.sequence for update in first_updates], [1, 2, 3] ) self.assertEqual( first_updates[-1].verification.state, route_pb2.ROUTE_STATE_READY, ) self.assertTrue(first_updates[-1].route.archive_control_created) duplicate = new_envelope() duplicate.command.CopyFrom(command.command) tasks = set() await daemon._handle(duplicate, outbound, tasks) self.assertEqual( decode(await outbound.get()).command_ack.status, control_pb2.COMMAND_ACK_STATUS_DUPLICATE, ) await next(iter(tasks)) replay = [decode(await outbound.get()).route_update for _ in range(3)] self.assertEqual( [update.update_id for update in replay], [update.update_id for update in first_updates], ) self.assertEqual(manager.configure.call_count, 1) accepted_before_crash = new_envelope() accepted_before_crash.command.CopyFrom(command.command) accepted_before_crash.command.command_id = str(uuid4()) acknowledgement = daemon._initial_acknowledgement( accepted_before_crash.command, set(), False ) await asyncio.to_thread( daemon.store.accept_command, accepted_before_crash.command.command_id, encode_message(accepted_before_crash.command), encode_message(acknowledgement), ) restarted = ArchiveClientDaemon( config, [probe, probe], [syncthing_probe], route_manager=manager, ) resumed_outbound = asyncio.Queue() resumed_tasks = set() await restarted._resume_route_commands( resumed_outbound, resumed_tasks ) await next(iter(resumed_tasks)) resumed = [ decode(await resumed_outbound.get()).route_update for _ in range(3) ] self.assertEqual([item.sequence for item in resumed], [1, 2, 3]) async def test_slow_inventory_does_not_block_heartbeat(self): with tempfile.TemporaryDirectory() as directory: root = Path(directory) token = root / "token" token.write_text("shared-secret", encoding="utf-8") os.chmod(token, 0o600) service = ServiceConfig( "http://local", PurePosixPath("/api"), root ) config = ClientConfig( "cache-1", "Cache 1", "cache", "ws://control", token, root / "state.db", root / "backups", service, service, ) gate = threading.Event() def slow_inventory(_filter=""): gate.wait(2) return [] reader = Mock( list_resources=Mock(side_effect=slow_inventory), get_resource=Mock(return_value=None), ) probe = FilesystemProbe(root, True, True, True, True, True) daemon = ArchiveClientDaemon( config, [probe, probe], [], resource_reader=reader ) await asyncio.to_thread(daemon.store.initialize) outbound = asyncio.Queue() tasks = set() command = new_envelope() command.command.command_id = str(uuid4()) command.command.created_at.CopyFrom(command.sent_at) command.command.inventory_query.query_id = str(uuid4()) command.command.inventory_query.scope = ( inventory_pb2.INVENTORY_SCOPE_RESOURCE_SUMMARIES ) await daemon._handle(command, outbound, tasks) self.assertEqual( decode(await outbound.get()).command_ack.status, control_pb2.COMMAND_ACK_STATUS_ACCEPTED, ) inventory_task = next(iter(tasks)) heartbeat = new_envelope() heartbeat.heartbeat.sequence = 9 await daemon._handle(heartbeat, outbound, tasks) self.assertEqual( decode(await outbound.get()).heartbeat_ack.sequence, 9 ) gate.set() await inventory_task self.assertTrue( decode(await outbound.get()).inventory_chunk.last_chunk ) 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 ] observed["features"] = list( registration.register_request.capabilities.features ) 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 ) inventory = new_envelope() inventory.command.command_id = str(uuid4()) inventory.command.created_at.CopyFrom(inventory.sent_at) inventory.command.inventory_query.query_id = str(uuid4()) inventory.command.inventory_query.scope = ( inventory_pb2.INVENTORY_SCOPE_RESOURCE_SUMMARIES ) await websocket.send(encode(inventory)) observed["inventory_ack"] = ( decode(await websocket.recv()).command_ack.status ) result = decode(await websocket.recv()).inventory_chunk observed["inventory_result"] = ( result.last_chunk, result.WhichOneof("payload") ) 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], resource_reader=Mock( list_resources=Mock(return_value=[]), get_resource=Mock(return_value=None), ), ) 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.assertIn( client_pb2.CLIENT_FEATURE_INVENTORY_CHUNKS, observed["features"], ) self.assertIn( client_pb2.CLIENT_FEATURE_ROUTE_PROVISIONING, observed["features"], ) 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, ) self.assertEqual( observed["inventory_ack"], control_pb2.COMMAND_ACK_STATUS_ACCEPTED, ) self.assertEqual( observed["inventory_result"], (True, "resource_summaries") ) if __name__ == "__main__": unittest.main()