From 5ae177033353728123287123ae893ba6f03c3047 Mon Sep 17 00:00:00 2001 From: Robert Washbourne Date: Sat, 3 Oct 2026 11:58:26 +0000 Subject: [PATCH] Track reused indexes as allocated across controller checkpoints --- tests/test_controller.py | 61 +++++++++++++++++++++++++++++++++++- transfer_queue/controller.py | 2 +- 2 files changed, 61 insertions(+), 2 deletions(-) diff --git a/tests/test_controller.py b/tests/test_controller.py index 32c1603a..434d46c7 100644 --- a/tests/test_controller.py +++ b/tests/test_controller.py @@ -14,6 +14,7 @@ # limitations under the License. import logging +import pickle from uuid import uuid4 import pytest @@ -21,7 +22,7 @@ import torch import zmq -from transfer_queue.controller import TQ_CONTROLLER_ZMQ_MAX_SOCKETS, TransferQueueController +from transfer_queue.controller import TQ_CONTROLLER_ZMQ_MAX_SOCKETS, PartitionIndexManager, TransferQueueController from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType, create_zmq_socket # Set up logging @@ -1362,6 +1363,64 @@ def test_controller_kv_retrieve_keys_multiple_partitions(self, ray_setup): class TestTransferQueueControllerCheckpoint: """Tests for TransferQueueController save/load checkpoint functionality.""" + @pytest.mark.parametrize("destination", ["old", "new"]) + @pytest.mark.parametrize("count", [2, 3, 5]) + def test_reused_indexes_remain_allocated(self, destination, count): + manager = PartitionIndexManager() + released = manager.allocate_indexes("old", 3) + retained = manager.allocate_indexes("retained", 2) + manager.release_partition("old") + + allocated = manager.allocate_indexes(destination, count) + assert set(allocated) & set(released) + assert not set(allocated) & set(retained) + assert manager.allocated_indexes == set(allocated) | set(retained) + assert not manager.allocated_indexes & set(manager.reusable_indexes) + + restored = pickle.loads(pickle.dumps(manager)) + assert restored.allocated_indexes == set().union(*restored.partition_to_indexes.values()) + assert not restored.allocated_indexes & set(restored.reusable_indexes) + + restored.release_partition(destination) + assert restored.allocated_indexes == set(retained) + + @pytest.mark.parametrize("destination", ["old", "new"]) + def test_controller_checkpoint_preserves_reused_index_ownership(self, ray_setup, tmp_path, destination): + controller = TransferQueueController.remote() + for partition, count in [("old", 3), ("retained", 2)]: + ray.get( + controller.get_metadata.remote( + data_fields=["value"], batch_size=count, partition_id=partition, mode="insert" + ) + ) + ray.get(controller.clear_partition.remote("old")) + reused = ray.get( + controller.get_metadata.remote( + data_fields=["value"], batch_size=4, partition_id=destination, mode="insert" + ) + ) + assert set(reused.global_indexes) == {0, 1, 2, 5} + + checkpoint = tmp_path / "controller.pkl" + ray.get(controller.save_checkpoint.remote(str(checkpoint))) + with checkpoint.open("rb") as saved: + state = pickle.load(saved) + manager = state["index_manager"] + assert manager["allocated_indexes"] == {0, 1, 2, 3, 4, 5} + assert manager["allocated_indexes"] == set().union(*manager["partition_to_indexes"].values()) + assert not manager["allocated_indexes"] & set(manager["reusable_indexes"]) + + restored = TransferQueueController.remote() + ray.get(restored.load_checkpoint.remote(str(checkpoint))) + assert set(ray.get(restored.get_partition_index_range.remote(destination))) == set(reused.global_indexes) + retained = ray.get(restored.get_partition_index_range.remote("retained")) + assert set(retained) == {3, 4} + ray.get(restored.clear_partition.remote(destination)) + after = tmp_path / "after-clear.pkl" + ray.get(restored.save_checkpoint.remote(str(after))) + with after.open("rb") as saved: + assert pickle.load(saved)["index_manager"]["allocated_indexes"] == set(retained) + def test_controller_checkpoint_round_trip(self, ray_setup, tmp_path): """Save and load controller state; verify partition data is preserved.""" gbs = 4 diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index b45a85a6..faa5dc95 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -121,12 +121,12 @@ def allocate_indexes(self, partition_id, count=1) -> list[int]: new_indexes = list(range(start_index, end_index)) # Batch update status - self.allocated_indexes.update(new_indexes) self.global_index_counter = end_index indexes.extend(new_indexes) # Record partition-index relationship + self.allocated_indexes.update(indexes) self.partition_to_indexes[partition_id].update(indexes) return indexes