diff --git a/README.md b/README.md index c741600..02fab75 100644 --- a/README.md +++ b/README.md @@ -22,6 +22,10 @@ so later job preflight can reject them without hiding the resource. The qBittorrent read adapter uses cookie authentication, bounded responses, one reauthentication attempt on session expiry, hash-scoped file/metainfo fetches, and never logs credentials, cookies, response bodies, or endpoints. +Durable inventory commands now stream bounded atomic summary, lookup, and +content-tree chunks. Page tokens are guarded by timestamp-independent snapshot +revisions, stale trees fail explicitly, and slow scans run outside the socket +reader so heartbeat acknowledgements remain responsive. ```bash archive-client --config /etc/archive-control/client.toml --check-config diff --git a/src/archive_clients/cli.py b/src/archive_clients/cli.py index d5020db..12cec10 100644 --- a/src/archive_clients/cli.py +++ b/src/archive_clients/cli.py @@ -13,6 +13,7 @@ from archive_clients.config import ClientConfig from archive_clients.daemon import ArchiveClientDaemon from archive_clients.logging_config import configure_logging from archive_clients.probes import probe_root, probe_writable_directory +from archive_clients.qbittorrent import QBittorrentReader from archive_clients.services import probe_qbittorrent, probe_syncthing @@ -66,7 +67,10 @@ def main(argv: Sequence[str] | None = None) -> int: "health": service.state, }, ) - asyncio.run(ArchiveClientDaemon(config, probes, service_probes).run()) + asyncio.run(ArchiveClientDaemon( + config, probes, service_probes, + resource_reader=QBittorrentReader(config.qbittorrent), + ).run()) except KeyboardInterrupt: logger.info( "client_stopped", diff --git a/src/archive_clients/daemon.py b/src/archive_clients/daemon.py index 56d57dc..d5958e1 100644 --- a/src/archive_clients/daemon.py +++ b/src/archive_clients/daemon.py @@ -13,6 +13,7 @@ from websockets.asyncio.client import connect from archive_clients.backup import SQLiteBackupManager from archive_clients.config import ClientConfig +from archive_clients.inventory import InventoryService from archive_clients.locking import DatabaseLease from archive_clients.probes import FilesystemProbe from archive_clients.protocol import ( @@ -22,9 +23,12 @@ from archive_clients.protocol import ( encode_message, new_envelope, ) +from archive_clients.qbittorrent import QBittorrentReader 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 +from archive_control.v1 import ( + client_pb2, common_pb2, control_pb2, inventory_pb2, job_pb2, +) logger = logging.getLogger(__name__) @@ -36,6 +40,7 @@ class ArchiveClientDaemon: config: ClientConfig, probes: list[FilesystemProbe], service_probes: list[ServiceProbe], + resource_reader: QBittorrentReader | None = None, ): if len(probes) != 2: raise ValueError( @@ -44,6 +49,10 @@ class ArchiveClientDaemon: self.config = config self.probes = probes self.service_probes = service_probes + self.inventory = ( + InventoryService(resource_reader, config.client_id) + if resource_reader is not None else None + ) self.store = ClientStore(config.state_db) self.backups = SQLiteBackupManager( config.state_db, config.backup_dir, config.backup @@ -137,12 +146,17 @@ class ArchiveClientDaemon: logger.info("control_connection_registered") outbound: asyncio.Queue[str] = asyncio.Queue(maxsize=100) writer = asyncio.create_task(self._writer(websocket, outbound)) + command_tasks: set[asyncio.Task[None]] = set() try: async for frame in websocket: - await self._handle(decode(frame), outbound) + await self._handle(decode(frame), outbound, command_tasks) finally: writer.cancel() - await asyncio.gather(writer, return_exceptions=True) + for task in command_tasks: + task.cancel() + await asyncio.gather( + writer, *command_tasks, return_exceptions=True + ) def _registration(self): envelope = new_envelope() @@ -195,6 +209,11 @@ class ArchiveClientDaemon: 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) + if self.inventory is not None: + request.capabilities.features.extend(( + client_pb2.CLIENT_FEATURE_INVENTORY_CHUNKS, + client_pb2.CLIENT_FEATURE_CONTENT_TREE, + )) for cursor in self.store.list_active_job_cursors(): active = request.active_jobs.add() active.job_id = str(cursor["job_id"]) @@ -211,7 +230,10 @@ class ArchiveClientDaemon: await websocket.send(await outbound.get()) async def _handle( - self, envelope: Any, outbound: asyncio.Queue[str] + self, + envelope: Any, + outbound: asyncio.Queue[str], + command_tasks: set[asyncio.Task[None]], ) -> None: payload = envelope.WhichOneof("payload") if payload == "heartbeat": @@ -220,7 +242,7 @@ class ArchiveClientDaemon: response.heartbeat_ack.sequence = envelope.heartbeat.sequence await outbound.put(encode(response)) elif payload == "command": - await self._accept_command(envelope, outbound) + await self._accept_command(envelope, outbound, command_tasks) elif payload == "protocol_error": logger.warning( "control_reported_protocol_error", @@ -228,7 +250,10 @@ class ArchiveClientDaemon: ) async def _accept_command( - self, envelope: Any, outbound: asyncio.Queue[str] + self, + envelope: Any, + outbound: asyncio.Queue[str], + command_tasks: set[asyncio.Task[None]], ) -> None: command = envelope.command snapshot_rows: list[dict[str, object]] = [] @@ -301,9 +326,55 @@ class ArchiveClientDaemon: row["last_event_sequence"] ) await outbound.put(encode(snapshot)) + elif ( + accepted is not None + and accepted_for_execution + and command.WhichOneof("payload") == "inventory_query" + and self.inventory is not None + ): + task = asyncio.create_task( + self._send_inventory( + command.inventory_query, envelope.message_id, outbound + ), + name=f"inventory-{command.inventory_query.query_id}", + ) + command_tasks.add(task) + task.add_done_callback( + lambda completed: self._command_finished( + completed, command_tasks + ) + ) + + async def _send_inventory( + self, + query: Any, + correlation_id: str, + outbound: asyncio.Queue[str], + ) -> None: + assert self.inventory is not None + chunks = await asyncio.to_thread(self.inventory.execute, query) + for chunk in chunks: + response = new_envelope() + response.correlation_id = correlation_id + response.inventory_chunk.CopyFrom(chunk) + await outbound.put(encode(response)) @staticmethod + def _command_finished( + task: asyncio.Task[None], command_tasks: set[asyncio.Task[None]] + ) -> None: + command_tasks.discard(task) + if task.cancelled(): + return + error = task.exception() + if error is not None: + logger.error( + "background_command_failed", + extra={"error_type": type(error).__name__}, + ) + def _initial_acknowledgement( + self, command: Any, missing_snapshot_jobs: set[str], has_snapshot_jobs: bool, @@ -324,6 +395,22 @@ class ArchiveClientDaemon: acknowledgement.error.message = "requested client job is not found" else: acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_ACCEPTED + elif command.WhichOneof("payload") == "inventory_query": + scope = command.inventory_query.scope + if self.inventory is None: + acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_REJECTED + acknowledgement.error.code = common_pb2.ERROR_CODE_UNAVAILABLE + acknowledgement.error.message = "inventory adapter is unavailable" + elif scope not in { + inventory_pb2.INVENTORY_SCOPE_RESOURCE_SUMMARIES, + inventory_pb2.INVENTORY_SCOPE_RESOURCE_LOOKUP, + inventory_pb2.INVENTORY_SCOPE_CONTENT_TREE, + }: + acknowledgement.status = control_pb2.COMMAND_ACK_STATUS_REJECTED + acknowledgement.error.code = common_pb2.ERROR_CODE_UNSUPPORTED + acknowledgement.error.message = "inventory scope is unsupported" + 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 diff --git a/src/archive_clients/inventory.py b/src/archive_clients/inventory.py new file mode 100644 index 0000000..9ca334e --- /dev/null +++ b/src/archive_clients/inventory.py @@ -0,0 +1,266 @@ +"""On-demand, revisioned inventory query execution and bounded chunking.""" + +from __future__ import annotations + +import hashlib +import uuid +from typing import Iterable + +from archive_clients.qbittorrent import QBittorrentError, QBittorrentReader +from archive_clients.resources import NormalizedResource, build_content_tree +from archive_control.v1 import common_pb2, inventory_pb2, resource_pb2 + + +class InventoryRequestError(ValueError): + pass + + +class InventoryService: + def __init__( + self, reader: QBittorrentReader, client_id: str, + chunk_target_bytes: int = 300 * 1024, + ): + self.reader = reader + self.client_id = client_id + self.chunk_target_bytes = chunk_target_bytes + + def execute( + self, query: inventory_pb2.InventoryQuery + ) -> list[inventory_pb2.InventoryChunk]: + try: + try: + parsed_query_id = uuid.UUID(query.query_id) + except (ValueError, AttributeError) as exc: + raise InventoryRequestError( + "query ID must be a canonical UUID" + ) from exc + if str(parsed_query_id) != query.query_id: + raise InventoryRequestError( + "query ID must be a canonical UUID" + ) + if query.scope == inventory_pb2.INVENTORY_SCOPE_RESOURCE_SUMMARIES: + return self._summaries(query, lookup=False) + if query.scope == inventory_pb2.INVENTORY_SCOPE_RESOURCE_LOOKUP: + return self._summaries(query, lookup=True) + if query.scope == inventory_pb2.INVENTORY_SCOPE_CONTENT_TREE: + return self._content_tree(query) + return [self._error( + query.query_id, common_pb2.ERROR_CODE_UNSUPPORTED, + "inventory scope is not supported", False, + )] + except InventoryRequestError as exc: + return [self._error( + query.query_id, common_pb2.ERROR_CODE_INVALID_ARGUMENT, + str(exc), False, + )] + except ValueError: + return [self._error( + query.query_id, common_pb2.ERROR_CODE_INTEGRITY_CHECK_FAILED, + "qBittorrent resource metadata is inconsistent", False, + )] + except QBittorrentError: + return [self._error( + query.query_id, common_pb2.ERROR_CODE_UNAVAILABLE, + "qBittorrent inventory is unavailable", True, + )] + + def _summaries( + self, query: inventory_pb2.InventoryQuery, lookup: bool, + ) -> list[inventory_pb2.InventoryChunk]: + if lookup: + if not query.resource_ids: + raise InventoryRequestError("resource lookup requires IDs") + resources = self._lookup_all(query.resource_ids) + else: + resources = self.reader.list_resources(query.page_filter) + if query.selected_complete_only: + resources = [ + item for item in resources + if item.summary.selected_complete_files.ranges + ] + resources.sort(key=lambda item: ( + item.summary.display_name.casefold(), + item.summary.resource_id.info_hash_v1_hex, + item.summary.resource_id.info_hash_v2_hex, + )) + revision = _revision(item.summary for item in resources) + if query.expected_revision and query.expected_revision != revision: + return [self._error( + query.query_id, common_pb2.ERROR_CODE_STALE_STATE, + "inventory revision changed", False, + )] + offset, page_size = _page(query) + selected = resources[offset:offset + page_size] + next_page = ( + str(offset + page_size) + if offset + page_size < len(resources) else "" + ) + return self._summary_chunks( + query.query_id, revision, selected, next_page + ) + + def _content_tree( + self, query: inventory_pb2.InventoryQuery + ) -> list[inventory_pb2.InventoryChunk]: + if len(query.resource_ids) != 1: + raise InventoryRequestError("content tree requires exactly one ID") + resource = self._lookup(query.resource_ids[0]) + if resource is None: + return [self._error( + query.query_id, common_pb2.ERROR_CODE_NOT_FOUND, + "resource is not present", False, + )] + revision = resource.summary.content_revision + if query.expected_revision and query.expected_revision != revision: + return [self._error( + query.query_id, common_pb2.ERROR_CODE_STALE_STATE, + "resource content revision changed", False, + )] + available = { + item.file_index for item in resource.files + if item.selected and item.completed_bytes == item.logical_bytes + } + entries = build_content_tree(resource.files, available) + return self._tree_chunks(query.query_id, revision, resource, entries) + + def _lookup_all( + self, resource_ids: Iterable[resource_pb2.ResourceId] + ) -> list[NormalizedResource]: + found: dict[tuple[str, str], NormalizedResource] = {} + for resource_id in resource_ids: + resource = self._lookup(resource_id) + if resource is not None: + identity = resource.summary.resource_id + found[( + identity.info_hash_v1_hex, identity.info_hash_v2_hex, + )] = resource + return list(found.values()) + + def _lookup( + self, resource_id: resource_pb2.ResourceId + ) -> NormalizedResource | None: + hashes = [ + value for value in ( + resource_id.info_hash_v1_hex, + resource_id.info_hash_v2_hex, + ) if value + ] + if not hashes: + raise InventoryRequestError("resource ID has no info hash") + for info_hash in hashes: + resource = self.reader.get_resource(info_hash) + if resource is None: + continue + actual = resource.summary.resource_id + if ( + resource_id.info_hash_v1_hex + and actual.info_hash_v1_hex != resource_id.info_hash_v1_hex + ) or ( + resource_id.info_hash_v2_hex + and actual.info_hash_v2_hex != resource_id.info_hash_v2_hex + ): + raise InventoryRequestError("resource identity hashes conflict") + return resource + return None + + def _summary_chunks( + self, query_id: str, revision: str, + resources: list[NormalizedResource], next_page: str, + ) -> list[inventory_pb2.InventoryChunk]: + groups = _bounded_groups( + [item.summary for item in resources], self.chunk_target_bytes + ) + snapshot_id = str(uuid.uuid4()) + return [self._chunk( + query_id, snapshot_id, revision, index, index == len(groups) - 1, + next_page if index == len(groups) - 1 else "", + summaries=group, + ) for index, group in enumerate(groups)] + + def _tree_chunks( + self, query_id: str, revision: str, resource: NormalizedResource, + entries: list[resource_pb2.ContentTreeEntry], + ) -> list[inventory_pb2.InventoryChunk]: + groups = _bounded_groups(entries, self.chunk_target_bytes) + snapshot_id = str(uuid.uuid4()) + return [self._chunk( + query_id, snapshot_id, revision, index, + index == len(groups) - 1, "", + tree=(resource.summary.resource_id, group), + ) for index, group in enumerate(groups)] + + @staticmethod + def _chunk( + query_id: str, snapshot_id: str, revision: str, + index: int, last: bool, + next_page: str, + summaries: list[resource_pb2.ResourceSummary] | None = None, + tree: tuple[ + resource_pb2.ResourceId, + list[resource_pb2.ContentTreeEntry], + ] | None = None, + ) -> inventory_pb2.InventoryChunk: + chunk = inventory_pb2.InventoryChunk( + query_id=query_id, snapshot_id=snapshot_id, + revision=revision, chunk_index=index, last_chunk=last, + next_page_token=next_page, + ) + if summaries is not None: + chunk.resource_summaries.resources.extend(summaries) + elif tree is not None: + chunk.content_tree.resource_id.CopyFrom(tree[0]) + chunk.content_tree.entries.extend(tree[1]) + chunk.content_tree.resource_content_revision = revision + return chunk + + @staticmethod + def _error( + query_id: str, code: int, message: str, retryable: bool, + ) -> inventory_pb2.InventoryChunk: + chunk = inventory_pb2.InventoryChunk( + query_id=query_id, snapshot_id=str(uuid.uuid4()), + revision="error", last_chunk=True, + ) + chunk.error.code = code + chunk.error.message = message + chunk.error.retryable = retryable + return chunk + + +def _bounded_groups(items: list, target: int) -> list[list]: + if not items: + return [[]] + groups: list[list] = [] + current: list = [] + size = 0 + for item in items: + item_size = item.ByteSize() + if current and size + item_size > target: + groups.append(current) + current = [] + size = 0 + current.append(item) + size += item_size + groups.append(current) + return groups + + +def _revision(messages: Iterable) -> str: + digest = hashlib.sha256() + for message in messages: + stable = message.__class__() + stable.CopyFrom(message) + if "observed_at" in stable.DESCRIPTOR.fields_by_name: + stable.ClearField("observed_at") + digest.update(stable.SerializeToString(deterministic=True)) + return digest.hexdigest() + + +def _page(query: inventory_pb2.InventoryQuery) -> tuple[int, int]: + size = query.page.page_size or 20 + if size > 100: + raise InventoryRequestError("page size cannot exceed 100") + token = query.page.page_token + if token and (not token.isdigit() or len(token) > 10): + raise InventoryRequestError("page token is invalid") + return int(token or 0), size diff --git a/tests/test_daemon.py b/tests/test_daemon.py index f50d875..6a07e6f 100644 --- a/tests/test_daemon.py +++ b/tests/test_daemon.py @@ -1,9 +1,11 @@ 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 @@ -13,10 +15,67 @@ 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 +from archive_control.v1 import ( + client_pb2, common_pb2, control_pb2, inventory_pb2, job_pb2, +) class DaemonTransportTests(unittest.IsolatedAsyncioTestCase): + 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()) @@ -39,6 +98,9 @@ class DaemonTransportTests(unittest.IsolatedAsyncioTestCase): 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 @@ -80,6 +142,21 @@ class DaemonTransportTests(unittest.IsolatedAsyncioTestCase): 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) @@ -104,7 +181,11 @@ class DaemonTransportTests(unittest.IsolatedAsyncioTestCase): datetime.now(timezone.utc), version="v2", device_id="DEVICE", ) daemon = ArchiveClientDaemon( - config, [probe, probe], [service_probe] + 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( @@ -128,6 +209,10 @@ class DaemonTransportTests(unittest.IsolatedAsyncioTestCase): 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.assertEqual(observed["heartbeat"], 7) self.assertEqual(observed["first"], control_pb2.COMMAND_ACK_STATUS_ACCEPTED) self.assertEqual(observed["second"], control_pb2.COMMAND_ACK_STATUS_DUPLICATE) @@ -145,6 +230,13 @@ class DaemonTransportTests(unittest.IsolatedAsyncioTestCase): 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__": diff --git a/tests/test_inventory.py b/tests/test_inventory.py new file mode 100644 index 0000000..ad5cdc4 --- /dev/null +++ b/tests/test_inventory.py @@ -0,0 +1,106 @@ +import unittest +from uuid import uuid4 + +from archive_clients.bencode import Metainfo +from archive_clients.inventory import InventoryService +from archive_clients.resources import NormalizedResource +from archive_control.v1 import common_pb2, inventory_pb2, resource_pb2 + + +class _Reader: + def __init__(self, resources): + self.resources = resources + + def list_resources(self, name_filter=""): + return [ + item for item in self.resources + if name_filter.casefold() in item.summary.display_name.casefold() + ] + + def get_resource(self, torrent_hash): + for item in self.resources: + identity = item.summary.resource_id + if torrent_hash in { + identity.info_hash_v1_hex, identity.info_hash_v2_hex, + }: + return item + return None + + +def _resource(index: int) -> NormalizedResource: + info_hash = f"{index + 1:040x}" + summary = resource_pb2.ResourceSummary( + qb_torrent_id=info_hash, + display_name=f"Resource {index}", + content_revision=f"revision-{index}", + canonical_paths=True, + ) + summary.resource_id.info_hash_v1_hex = info_hash + summary.selected_files.ranges.add(first=0, last=0) + summary.selected_complete_files.ranges.add(first=0, last=0) + file = resource_pb2.TorrentFile( + file_index=0, canonical_path=f"resource-{index}/file.bin", + logical_bytes=10, completed_bytes=10, selected=True, + ) + return NormalizedResource( + summary, (file,), Metainfo(info_hash, "", ()), + ) + + +class InventoryTests(unittest.TestCase): + def test_summary_chunks_share_atomic_snapshot_and_page(self): + service = InventoryService( + _Reader([_resource(index) for index in range(4)]), + "cache-1", chunk_target_bytes=1, + ) + query = inventory_pb2.InventoryQuery( + query_id=str(uuid4()), + scope=inventory_pb2.INVENTORY_SCOPE_RESOURCE_SUMMARIES, + ) + query.page.page_size = 3 + chunks = service.execute(query) + self.assertEqual(len(chunks), 3) + self.assertEqual([chunk.chunk_index for chunk in chunks], [0, 1, 2]) + self.assertEqual(len({chunk.snapshot_id for chunk in chunks}), 1) + self.assertFalse(chunks[0].last_chunk) + self.assertTrue(chunks[-1].last_chunk) + self.assertEqual(chunks[-1].next_page_token, "3") + query.expected_revision = chunks[-1].revision + query.page.page_token = chunks[-1].next_page_token + next_page = service.execute(query) + self.assertFalse(next_page[0].HasField("error")) + self.assertEqual( + next_page[0].resource_summaries.resources[0].display_name, + "Resource 3", + ) + + def test_content_tree_requires_matching_revision(self): + resource = _resource(0) + service = InventoryService(_Reader([resource]), "cache-1") + query = inventory_pb2.InventoryQuery( + query_id=str(uuid4()), + scope=inventory_pb2.INVENTORY_SCOPE_CONTENT_TREE, + expected_revision="stale", + ) + query.resource_ids.add().CopyFrom(resource.summary.resource_id) + stale = service.execute(query) + self.assertEqual(stale[0].error.code, common_pb2.ERROR_CODE_STALE_STATE) + query.expected_revision = resource.summary.content_revision + chunks = service.execute(query) + self.assertEqual(chunks[0].WhichOneof("payload"), "content_tree") + self.assertEqual(chunks[0].content_tree.entries[0].available_file_count, 1) + + def test_invalid_query_is_a_terminal_error_chunk(self): + query = inventory_pb2.InventoryQuery( + query_id="invalid", + scope=inventory_pb2.INVENTORY_SCOPE_RESOURCE_SUMMARIES, + ) + result = InventoryService(_Reader([]), "cache-1").execute(query) + self.assertTrue(result[0].last_chunk) + self.assertEqual( + result[0].error.code, common_pb2.ERROR_CODE_INVALID_ARGUMENT + ) + + +if __name__ == "__main__": + unittest.main()