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
+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()