feat: stream on-demand inventory queries

This commit is contained in:
2026-07-23 01:46:58 +00:00
parent 720fa67202
commit 4058b6d3c8
6 changed files with 568 additions and 9 deletions
+4
View File
@@ -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
+5 -1
View File
@@ -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",
+93 -6
View File
@@ -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
+266
View File
@@ -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
+94 -2
View File
@@ -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__":
+106
View File
@@ -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()